diff --git a/src/plugins/intel_cpu/src/nodes/kernels/scaled_attn/common.hpp b/src/plugins/intel_cpu/src/nodes/kernels/scaled_attn/common.hpp index f7afc9641bb..bd05801c139 100644 --- a/src/plugins/intel_cpu/src/nodes/kernels/scaled_attn/common.hpp +++ b/src/plugins/intel_cpu/src/nodes/kernels/scaled_attn/common.hpp @@ -82,6 +82,11 @@ static constexpr size_t vec_len_f32_avx2 = vec_len_avx2 / sizeof(float); auto vec_f16 = _mm256_loadu_si256(reinterpret_cast(a)); return _mm512_cvtph_ps(vec_f16); } + inline __m512 mm512_uni_loadu_tail_ps(const ov::float16* a, size_t count) { + auto mask = (1 << count) - 1; + auto f16_vec = _mm256_maskz_loadu_epi16(mask, a); + return _mm512_cvtph_ps(f16_vec); + } inline void mm512_uni_storeu_ps(ov::float16* addr, __m512 v) { __m256i vec_f16 = _mm512_cvtps_ph(v, 0); _mm256_storeu_si256(reinterpret_cast<__m256i *>(addr), vec_f16); @@ -149,6 +154,11 @@ static constexpr size_t vec_len_f32_avx2 = vec_len_avx2 / sizeof(float); auto o = _mm256_cvtph_ps(vec_f16); return o; } + inline __m256 mm256_uni_loadu_tail_ps(const ov::float16* a, const size_t count) { + ov::float16 tmp_values[8] = {0}; + std::memcpy(tmp_values, a, count * sizeof(ov::float16)); + return mm256_uni_loadu_ps(tmp_values); + } inline void mm256_uni_storeu_ps(ov::float16* a, __m256 v) { __m128i vec_f16 = _mm256_cvtps_ph(v, 0); _mm_storeu_si128(reinterpret_cast<__m128i *>(a), vec_f16); diff --git a/src/plugins/intel_cpu/src/nodes/kernels/scaled_attn/executor_pa.cpp b/src/plugins/intel_cpu/src/nodes/kernels/scaled_attn/executor_pa.cpp index cd46be61746..d07f7490f1b 100644 --- a/src/plugins/intel_cpu/src/nodes/kernels/scaled_attn/executor_pa.cpp +++ b/src/plugins/intel_cpu/src/nodes/kernels/scaled_attn/executor_pa.cpp @@ -599,10 +599,11 @@ static void attn_reduce(T* dst, float* temp, size_t M, size_t S, size_t temp_str } } -// N and K must be multiple of 16 +// N must be multiple of 16 template -void transpose_16Nx16K(TDST* dst, TSRC* src, TDST* tmp, size_t N, size_t K, size_t dst_stride, size_t src_stride) { - for (size_t k = 0; k < K; k += 16) { +void transpose_16NxK(TDST* dst, TSRC* src, TDST* tmp, size_t N, size_t K, size_t dst_stride, size_t src_stride) { + size_t k = 0; + for (; k + 16 <= K; k += 16) { for (size_t n = 0; n < N; n += 16) { transpose_16x16_kernel(dst + n, src + n * src_stride, dst_stride, src_stride); } @@ -610,19 +611,24 @@ void transpose_16Nx16K(TDST* dst, TSRC* src, TDST* tmp, size_t N, size_t K, size dst += 16 * dst_stride; src += 16; } + if (k < K) { + for (size_t n = 0; n < N; n += 16) { + transpose_16xK_kernel(dst + n, src + n * src_stride, K - k, dst_stride, src_stride); + } + } } #if defined(HAVE_AVX512F) -static void transpose_16Nx16K(ov::bfloat16* dst, ov::bfloat16* src, ov::bfloat16* tmp, size_t N, size_t K, size_t dst_stride, size_t src_stride) { +static void transpose_16NxK(ov::bfloat16* dst, ov::bfloat16* src, ov::bfloat16* tmp, size_t N, size_t K, size_t dst_stride, size_t src_stride) { // will treat as uint32_t transpose auto s = reinterpret_cast(src); auto d = reinterpret_cast(dst); - transpose_16Nx16K(d, s, reinterpret_cast(0), N, K >> 1, dst_stride, src_stride >> 1); + transpose_16NxK(d, s, reinterpret_cast(0), N, K >> 1, dst_stride, src_stride >> 1); } #endif template -void transpose_16Nx16K(TDST* dst, uint8_t* src, TDST* tmp, size_t N, size_t K, size_t dst_stride, size_t src_stride) { +void transpose_16NxK(TDST* dst, uint8_t* src, TDST* tmp, size_t N, size_t K, size_t dst_stride, size_t src_stride) { // The layout for per token per head: // |scale(f32)|zeropoint(f32)|quantized feature(u8,idx_1)|quantized feature(u8,idx_2)|...|quantized feature(u8,idx_S)| // The quantized feature will start from 8bytes=sizeof(float)+sizeof(float) @@ -634,7 +640,7 @@ void transpose_16Nx16K(TDST* dst, uint8_t* src, TDST* tmp, size_t N, size_t K, s s += src_stride + 2 * sizeof(float); t += src_stride; } - transpose_16Nx16K(dst, tmp, reinterpret_cast(0), N, K, dst_stride, src_stride); + transpose_16NxK(dst, tmp, reinterpret_cast(0), N, K, dst_stride, src_stride); } // dequant f16/u8 to float @@ -664,55 +670,55 @@ void dequant(TDST* dst, uint8_t* src, size_t N, size_t K) { #if defined(HAVE_AVX512F) // pack bf16/u8 to bf16 -static void pack_32x32_kernel(ov::bfloat16* dst, ov::bfloat16* src, size_t stride) { +static void pack_32x32_kernel(ov::bfloat16* dst, ov::bfloat16* src, size_t dst_stride, size_t src_stride) { static const uint64_t idx[8] = {0, 4, 1, 5, 2, 6, 3, 7}; auto midx = _mm512_loadu_si512(idx); for (size_t i = 0; i < 16; i++) { auto a = _mm512_loadu_si512(src); // [a1 a2 a3 a4 | a5 a6 a7 a8] total 512-bits in 8 64bits unit - auto b = _mm512_loadu_si512(src + stride); // [b1 b2 b3 b4 | b5 b6 b7 b8] total 512-bits + auto b = _mm512_loadu_si512(src + src_stride); // [b1 b2 b3 b4 | b5 b6 b7 b8] total 512-bits a = _mm512_permutexvar_epi64(midx, a); // [a1 a5 | a2 a6 | a3 a7 | a4 a8] b = _mm512_permutexvar_epi64(midx, b); // [b1 b5 | b2 b6 | b3 b7 | b4 b8] auto B0 = _mm512_unpacklo_epi16(a, b); // [ a1&b1 a2&b2 a3&b3 a4&b4] for each 128-bits lane, interleave word in low 64 bits auto B1 = _mm512_unpackhi_epi16(a, b); // [ a5&b5 a6&b6 a7&b7 a8&b8] for each 128-bits lane, interleave word in high 64 bits _mm512_storeu_si512(dst, B0); _mm512_storeu_si512(dst + 32, B1); - src += 2 * stride; - dst += 2 * stride; + src += 2 * src_stride; + dst += 2 * dst_stride; } } -static void pack_32x16_kernel(ov::bfloat16* dst, ov::bfloat16* src, size_t stride) { +static void pack_32x16_kernel(ov::bfloat16* dst, ov::bfloat16* src, size_t dst_stride, size_t src_stride) { static const uint64_t idx[8] = {0, 4, 1, 5, 2, 6, 3, 7}; auto midx = _mm512_loadu_si512(idx); for (size_t i = 0; i < 16; i++) { auto x = _mm256_loadu_si256(reinterpret_cast<__m256i*>(src)); // [a1 a2 a3 a4] total 256-bits in 4 64bits unit - auto y = _mm256_loadu_si256(reinterpret_cast<__m256i*>(src + stride)); // [b1 b2 b3 b4] total 256-bits + auto y = _mm256_loadu_si256(reinterpret_cast<__m256i*>(src + src_stride)); // [b1 b2 b3 b4] total 256-bits auto a = _mm512_castsi256_si512(x); auto b = _mm512_castsi256_si512(y); a = _mm512_permutexvar_epi64(midx, a); // [a1 x | a2 x | a3 x | a4 x] b = _mm512_permutexvar_epi64(midx, b); // [b1 x | b2 x | b3 x | b4 x] auto B0 = _mm512_unpacklo_epi16(a, b); _mm512_storeu_si512(dst, B0); - src += 2 * stride; - dst += 2 * stride; + src += 2 * src_stride; + dst += 2 * dst_stride; } } -static void pack_32Nx16K(ov::bfloat16* dst, ov::bfloat16* src, ov::bfloat16* tmp, size_t N, size_t K, size_t stride) { +static void pack_32Nx16K(ov::bfloat16* dst, ov::bfloat16* src, ov::bfloat16* tmp, size_t N, size_t K, size_t dst_stride, size_t src_stride) { for (size_t n = 0; n < N; n += 32) { size_t k = 0; for (; k + 32 <= K; k += 32) { - pack_32x32_kernel(dst + k * 2, src + k, stride); + pack_32x32_kernel(dst + k * 2, src + k, dst_stride, src_stride); } if (k < K) - pack_32x16_kernel(dst + k * 2, src + k, stride); + pack_32x16_kernel(dst + k * 2, src + k, dst_stride, src_stride); - dst += 32 * stride; - src += 32 * stride; + dst += 32 * dst_stride; + src += 32 * src_stride; } } -static void pack_32Nx16K(ov::bfloat16* dst, uint8_t* src, ov::bfloat16* tmp, size_t N, size_t K, size_t stride) { +static void pack_32Nx16K(ov::bfloat16* dst, uint8_t* src, ov::bfloat16* tmp, size_t N, size_t K, size_t dst_stride, size_t src_stride) { // The layout for per token per head: // |scale(f32)|zeropoint(f32)|quantized feature(u8,idx_1)|quantized feature(u8,idx_2)|...|quantized feature(u8,idx_S)| // The quantized feature will start from 8bytes=sizeof(float)+sizeof(float) @@ -721,15 +727,15 @@ static void pack_32Nx16K(ov::bfloat16* dst, uint8_t* src, ov::bfloat16* tmp, siz for (size_t n = 0; n < N; n ++) { auto f = reinterpret_cast(s); attn_dequant_u8_kernel(s + 2 * sizeof(float), t, K, f[0], f[1]); - s += stride + 2 * sizeof(float); - t += stride; + s += src_stride + 2 * sizeof(float); + t += src_stride; } - pack_32Nx16K(dst, tmp, reinterpret_cast(0), N, K, stride); + pack_32Nx16K(dst, tmp, reinterpret_cast(0), N, K, dst_stride, src_stride); } #endif template -static void pack_32Nx16K(float* dst, T* src, float* tmp, size_t N, size_t K, size_t stride) { +static void pack_32Nx16K(float* dst, T* src, float* tmp, size_t N, size_t K, size_t dst_stride, size_t src_stride) { // never called OPENVINO_THROW("pack_32Nx16K: should not be called."); } @@ -858,7 +864,7 @@ struct MHAHelper { void init_reorder_buffers(size_t batch, size_t kv_len_in_blocks) { _qk_scratch_b.resize({batch, kv_len_in_blocks, _Hk, _block_size * _S}); - _wv_scratch_b.resize({batch, kv_len_in_blocks, _Hk, _block_size * _S}); + _wv_scratch_b.resize({batch, kv_len_in_blocks, _Hk, _block_size * rnd_up(_S, _block_size)}); } // compute one block(such as 32 tokens) of query in M dimension: softmax(q_block*k')*v @@ -1307,7 +1313,7 @@ struct MHA { auto ithr = parallel_get_thread_num(); auto* k_ptr = k_cache.ptr(block_number, hk); auto* v_ptr = v_cache.ptr(block_number, hk); - transpose_16Nx16K(_helper._qk_scratch_b.template ptr(batch_in_reorder, kv_block, hk), + transpose_16NxK(_helper._qk_scratch_b.template ptr(batch_in_reorder, kv_block, hk), k_ptr, _helper._output.template ptr(ithr), _helper._block_size, @@ -1318,6 +1324,7 @@ struct MHA { _helper._output.template ptr(ithr), _helper._block_size, _helper._S, + rnd_up(_helper._S, _helper._block_size), _helper._S); } else { // need to decompress diff --git a/src/plugins/intel_cpu/src/nodes/kernels/scaled_attn/transpose_kernel.hpp b/src/plugins/intel_cpu/src/nodes/kernels/scaled_attn/transpose_kernel.hpp index b39028792ee..b719246e497 100644 --- a/src/plugins/intel_cpu/src/nodes/kernels/scaled_attn/transpose_kernel.hpp +++ b/src/plugins/intel_cpu/src/nodes/kernels/scaled_attn/transpose_kernel.hpp @@ -133,6 +133,50 @@ inline void transpose_16x16_kernel(float* _dst, T* src, size_t dst_stride, size_ _mm512_storeu_si512(dst + 15 * dst_stride, rf); } +template +inline void transpose_16xK_kernel(float* _dst, T* src, size_t K, size_t dst_stride, size_t src_stride) { + auto* dst = reinterpret_cast(_dst); + __m512i r0, r1, r2, r3, r4, r5, r6, r7, r8, r9, ra, rb, rc, rd, re, rf; + r0 = _mm512_castps_si512(mm512_uni_loadu_tail_ps(src, K)); + r1 = _mm512_castps_si512(mm512_uni_loadu_tail_ps(src + src_stride, K)); + r2 = _mm512_castps_si512(mm512_uni_loadu_tail_ps(src + 2 * src_stride, K)); + r3 = _mm512_castps_si512(mm512_uni_loadu_tail_ps(src + 3 * src_stride, K)); + r4 = _mm512_castps_si512(mm512_uni_loadu_tail_ps(src + 4 * src_stride, K)); + r5 = _mm512_castps_si512(mm512_uni_loadu_tail_ps(src + 5 * src_stride, K)); + r6 = _mm512_castps_si512(mm512_uni_loadu_tail_ps(src + 6 * src_stride, K)); + r7 = _mm512_castps_si512(mm512_uni_loadu_tail_ps(src + 7 * src_stride, K)); + r8 = _mm512_castps_si512(mm512_uni_loadu_tail_ps(src + 8 * src_stride, K)); + r9 = _mm512_castps_si512(mm512_uni_loadu_tail_ps(src + 9 * src_stride, K)); + ra = _mm512_castps_si512(mm512_uni_loadu_tail_ps(src + 10 * src_stride, K)); + rb = _mm512_castps_si512(mm512_uni_loadu_tail_ps(src + 11 * src_stride, K)); + rc = _mm512_castps_si512(mm512_uni_loadu_tail_ps(src + 12 * src_stride, K)); + rd = _mm512_castps_si512(mm512_uni_loadu_tail_ps(src + 13 * src_stride, K)); + re = _mm512_castps_si512(mm512_uni_loadu_tail_ps(src + 14 * src_stride, K)); + rf = _mm512_castps_si512(mm512_uni_loadu_tail_ps(src + 15 * src_stride, K)); + + transpose_m512i_16x16(r0, r1, r2, r3, r4, r5, r6, r7, r8, r9, ra, rb, rc, rd, re, rf); + +#define S(m) _mm512_storeu_si512(dst + 0x##m * dst_stride, r##m) +#define S8() S(0); S(1); S(2); S(3); S(4); S(5); S(6); S(7); + switch (K) { + case 8: S8(); break; + case 9: S8() S(8); break; + case 10: S8(); S(8); S(9); break; + case 11: S8(); S(8); S(9); S(a); break; + case 12: S8(); S(8); S(9); S(a); S(b); break; + case 13: S8(); S(8); S(9); S(a); S(b); S(c); break; + case 14: S8(); S(8); S(9); S(a); S(b); S(c); S(d); break; + case 15: S8(); S(8); S(9); S(a); S(b); S(c); S(d); S(e); break; + case 1: S(0); break; + case 2: S(0); S(1); break; + case 3: S(0); S(1); S(2); break; + case 4: S(0); S(1); S(2); S(3); break; + case 5: S(0); S(1); S(2); S(3); S(4); break; + case 6: S(0); S(1); S(2); S(3); S(4); S(5); break; + case 7: S(0); S(1); S(2); S(3); S(4); S(5); S(6); break; + } +} + inline void transpose_16x16_kernel(uint32_t* dst, uint32_t* src, size_t dst_stride, size_t src_stride) { __m512i r0, r1, r2, r3, r4, r5, r6, r7, r8, r9, ra, rb, rc, rd, re, rf; r0 = _mm512_loadu_si512(src); @@ -172,6 +216,50 @@ inline void transpose_16x16_kernel(uint32_t* dst, uint32_t* src, size_t dst_stri _mm512_storeu_si512(dst + 15 * dst_stride, rf); } +inline void transpose_16xK_kernel(uint32_t* dst, uint32_t* src, size_t K, size_t dst_stride, size_t src_stride) { + __m512i r0, r1, r2, r3, r4, r5, r6, r7, r8, r9, ra, rb, rc, rd, re, rf; + __mmask16 k = 0xffff >> (16 - K); + + r0 = _mm512_maskz_loadu_epi32(k, src); + r1 = _mm512_maskz_loadu_epi32(k, src + src_stride); + r2 = _mm512_maskz_loadu_epi32(k, src + 2 * src_stride); + r3 = _mm512_maskz_loadu_epi32(k, src + 3 * src_stride); + r4 = _mm512_maskz_loadu_epi32(k, src + 4 * src_stride); + r5 = _mm512_maskz_loadu_epi32(k, src + 5 * src_stride); + r6 = _mm512_maskz_loadu_epi32(k, src + 6 * src_stride); + r7 = _mm512_maskz_loadu_epi32(k, src + 7 * src_stride); + r8 = _mm512_maskz_loadu_epi32(k, src + 8 * src_stride); + r9 = _mm512_maskz_loadu_epi32(k, src + 9 * src_stride); + ra = _mm512_maskz_loadu_epi32(k, src + 10 * src_stride); + rb = _mm512_maskz_loadu_epi32(k, src + 11 * src_stride); + rc = _mm512_maskz_loadu_epi32(k, src + 12 * src_stride); + rd = _mm512_maskz_loadu_epi32(k, src + 13 * src_stride); + re = _mm512_maskz_loadu_epi32(k, src + 14 * src_stride); + rf = _mm512_maskz_loadu_epi32(k, src + 15 * src_stride); + + transpose_m512i_16x16(r0, r1, r2, r3, r4, r5, r6, r7, r8, r9, ra, rb, rc, rd, re, rf); + + switch (K) { + case 8: S8(); break; + case 9: S8() S(8); break; + case 10: S8(); S(8); S(9); break; + case 11: S8(); S(8); S(9); S(a); break; + case 12: S8(); S(8); S(9); S(a); S(b); break; + case 13: S8(); S(8); S(9); S(a); S(b); S(c); break; + case 14: S8(); S(8); S(9); S(a); S(b); S(c); S(d); break; + case 15: S8(); S(8); S(9); S(a); S(b); S(c); S(d); S(e); break; + case 1: S(0); break; + case 2: S(0); S(1); break; + case 3: S(0); S(1); S(2); break; + case 4: S(0); S(1); S(2); S(3); break; + case 5: S(0); S(1); S(2); S(3); S(4); break; + case 6: S(0); S(1); S(2); S(3); S(4); S(5); break; + case 7: S(0); S(1); S(2); S(3); S(4); S(5); S(6); break; + } +#undef S +#undef S8 +} + #elif defined(HAVE_AVX2) // https://stackoverflow.com/questions/25622745/transpose-an-8x8-float-using-avx-avx2 @@ -235,6 +323,64 @@ inline void transpose_16x16_kernel(float* dst, T* src, size_t dst_stride, size_t } } +template +inline void transpose_16xK_kernel(float* dst, T* src, size_t K, size_t dst_stride, size_t src_stride) { + __m256 r0, r1, r2, r3, r4, r5, r6, r7; + + if (K >= 8) { + for (int j = 0; j < 16; j += 8) { + r0 = mm256_uni_loadu_ps(src + src_stride * j); + r1 = mm256_uni_loadu_ps(src + src_stride * (1 + j)); + r2 = mm256_uni_loadu_ps(src + src_stride * (2 + j)); + r3 = mm256_uni_loadu_ps(src + src_stride * (3 + j)); + r4 = mm256_uni_loadu_ps(src + src_stride * (4 + j)); + r5 = mm256_uni_loadu_ps(src + src_stride * (5 + j)); + r6 = mm256_uni_loadu_ps(src + src_stride * (6 + j)); + r7 = mm256_uni_loadu_ps(src + src_stride * (7 + j)); + + transpose_8x8(r0, r1, r2, r3, r4, r5, r6, r7); + + _mm256_storeu_ps(dst + j, r0); + _mm256_storeu_ps(dst + j + dst_stride, r1); + _mm256_storeu_ps(dst + j + dst_stride * 2, r2); + _mm256_storeu_ps(dst + j + dst_stride * 3, r3); + _mm256_storeu_ps(dst + j + dst_stride * 4, r4); + _mm256_storeu_ps(dst + j + dst_stride * 5, r5); + _mm256_storeu_ps(dst + j + dst_stride * 6, r6); + _mm256_storeu_ps(dst + j + dst_stride * 7, r7); + } + src += 8; + dst += 8 * dst_stride; + K -= 8; + } + if (K > 0) { + for (int j = 0; j < 16; j += 8) { + r0 = mm256_uni_loadu_tail_ps(src + src_stride * j, K); + r1 = mm256_uni_loadu_tail_ps(src + src_stride * (1 + j), K); + r2 = mm256_uni_loadu_tail_ps(src + src_stride * (2 + j), K); + r3 = mm256_uni_loadu_tail_ps(src + src_stride * (3 + j), K); + r4 = mm256_uni_loadu_tail_ps(src + src_stride * (4 + j), K); + r5 = mm256_uni_loadu_tail_ps(src + src_stride * (5 + j), K); + r6 = mm256_uni_loadu_tail_ps(src + src_stride * (6 + j), K); + r7 = mm256_uni_loadu_tail_ps(src + src_stride * (7 + j), K); + + transpose_8x8(r0, r1, r2, r3, r4, r5, r6, r7); + +#define S(m) _mm256_storeu_ps(dst + j + m * dst_stride, r##m) + switch (K) { + case 1: S(0); break; + case 2: S(0); S(1); break; + case 3: S(0); S(1); S(2); break; + case 4: S(0); S(1); S(2); S(3); break; + case 5: S(0); S(1); S(2); S(3); S(4); break; + case 6: S(0); S(1); S(2); S(3); S(4); S(5); break; + case 7: S(0); S(1); S(2); S(3); S(4); S(5); S(6); break; + } +#undef S + } + } +} + #else template @@ -246,6 +392,15 @@ inline void transpose_16x16_kernel(TDST* dst, TSRC* src, size_t dst_stride, size } } +template +inline void transpose_16xK_kernel(TDST* dst, TSRC* src, size_t K, size_t dst_stride, size_t src_stride) { + for (size_t i = 0; i < K; i++) { + for (size_t j = 0; j < 16; j++) { + dst[i * dst_stride + j] = static_cast(src[i + j * src_stride]); + } + } +} + #endif } // namespace XARCH