[CPU] Fix segment fault for PagedAttention when head size is multiple of 16 (#24627)
### Details: - *Fix segment fault for PagedAttention when head size is multiple of 16* - *...* ### Tickets: - *ticket-id*
This commit is contained in:
parent
87f25aa14b
commit
f400fe5d87
|
|
@ -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<const __m256i*>(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);
|
||||
|
|
|
|||
|
|
@ -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<typename TDST, typename TSRC>
|
||||
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<uint32_t*>(src);
|
||||
auto d = reinterpret_cast<uint32_t*>(dst);
|
||||
transpose_16Nx16K(d, s, reinterpret_cast<uint32_t*>(0), N, K >> 1, dst_stride, src_stride >> 1);
|
||||
transpose_16NxK(d, s, reinterpret_cast<uint32_t*>(0), N, K >> 1, dst_stride, src_stride >> 1);
|
||||
}
|
||||
#endif
|
||||
|
||||
template<typename TDST>
|
||||
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<TDST*>(0), N, K, dst_stride, src_stride);
|
||||
transpose_16NxK(dst, tmp, reinterpret_cast<TDST*>(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<float*>(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<ov::bfloat16*>(0), N, K, stride);
|
||||
pack_32Nx16K(dst, tmp, reinterpret_cast<ov::bfloat16*>(0), N, K, dst_stride, src_stride);
|
||||
}
|
||||
#endif
|
||||
|
||||
template<typename T>
|
||||
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<DATA_TYPE>({batch, kv_len_in_blocks, _Hk, _block_size * _S});
|
||||
_wv_scratch_b.resize<DATA_TYPE>({batch, kv_len_in_blocks, _Hk, _block_size * _S});
|
||||
_wv_scratch_b.resize<DATA_TYPE>({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<KVCACHE_TYPE>(block_number, hk);
|
||||
auto* v_ptr = v_cache.ptr<KVCACHE_TYPE>(block_number, hk);
|
||||
transpose_16Nx16K(_helper._qk_scratch_b.template ptr<DATA_TYPE>(batch_in_reorder, kv_block, hk),
|
||||
transpose_16NxK(_helper._qk_scratch_b.template ptr<DATA_TYPE>(batch_in_reorder, kv_block, hk),
|
||||
k_ptr,
|
||||
_helper._output.template ptr<DATA_TYPE>(ithr),
|
||||
_helper._block_size,
|
||||
|
|
@ -1318,6 +1324,7 @@ struct MHA {
|
|||
_helper._output.template ptr<DATA_TYPE>(ithr),
|
||||
_helper._block_size,
|
||||
_helper._S,
|
||||
rnd_up(_helper._S, _helper._block_size),
|
||||
_helper._S);
|
||||
} else {
|
||||
// need to decompress
|
||||
|
|
|
|||
|
|
@ -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<typename T>
|
||||
inline void transpose_16xK_kernel(float* _dst, T* src, size_t K, size_t dst_stride, size_t src_stride) {
|
||||
auto* dst = reinterpret_cast<uint32_t*>(_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<typename T>
|
||||
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<typename TSRC, typename TDST>
|
||||
|
|
@ -246,6 +392,15 @@ inline void transpose_16x16_kernel(TDST* dst, TSRC* src, size_t dst_stride, size
|
|||
}
|
||||
}
|
||||
|
||||
template<typename TSRC, typename TDST>
|
||||
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<TDST>(src[i + j * src_stride]);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#endif
|
||||
|
||||
} // namespace XARCH
|
||||
|
|
|
|||
Loading…
Reference in New Issue