37#if defined(__AVX512F__) || defined(__AVX2__) || defined(__AVX__) || defined(__SSE4_1__)
40#if defined(__ARM_NEON) || defined(__aarch64__)
48 static int cached = -1;
50 const char *env = getenv(
"CK_DEBUG_Q8_0_REF");
51 cached = (env && env[0] && env[0] !=
'0') ? 1 : 0;
58 static int cached = -1;
60 const char *env = getenv(
"CK_DEBUG_Q8_0_Q8_0_REF");
61 cached = (env && env[0] && env[0] !=
'0') ? 1 : 0;
68 static int cached = -1;
77 float val = fval + 12582912.f;
79 memcpy(&i, &val,
sizeof(
int));
80 return (i & 0x007fffff) - 0x00400000;
83#if defined(__INTEL_LLVM_COMPILER)
85#define CK_Q80_NOINLINE_OPTNONE __attribute__((noinline, optnone))
86#elif defined(__GNUC__)
87#define CK_Q80_NOINLINE_OPTNONE __attribute__((noinline, optimize("O0")))
89#define CK_Q80_NOINLINE_OPTNONE
92static CK_Q80_NOINLINE_OPTNONE
float
93ck_q8_0_div_rounded_f32(
float numerator,
float denominator)
100 volatile float n = numerator;
101 volatile float d = denominator;
102 volatile float result = n / d;
128 const int nb = k /
QK8_0;
131 const __m256 sign_bit = _mm256_set1_ps(-0.0f);
133 for (
int i = 0; i < nb; i++) {
134 __m256 v0 = _mm256_loadu_ps(x + 0);
135 __m256 v1 = _mm256_loadu_ps(x + 8);
136 __m256 v2 = _mm256_loadu_ps(x + 16);
137 __m256 v3 = _mm256_loadu_ps(x + 24);
140 __m256 max_abs = _mm256_andnot_ps(sign_bit, v0);
141 max_abs = _mm256_max_ps(max_abs, _mm256_andnot_ps(sign_bit, v1));
142 max_abs = _mm256_max_ps(max_abs, _mm256_andnot_ps(sign_bit, v2));
143 max_abs = _mm256_max_ps(max_abs, _mm256_andnot_ps(sign_bit, v3));
145 __m128 max4 = _mm_max_ps(_mm256_extractf128_ps(max_abs, 1),
146 _mm256_castps256_ps128(max_abs));
147 max4 = _mm_max_ps(max4, _mm_movehl_ps(max4, max4));
148 max4 = _mm_max_ss(max4, _mm_movehdup_ps(max4));
149 const float max_scalar = _mm_cvtss_f32(max4);
151#if defined(__INTEL_LLVM_COMPILER)
152 const float d = ck_q8_0_div_rounded_f32(max_scalar, 127.0f);
153 const float id = max_scalar != 0.0f
154 ? ck_q8_0_div_rounded_f32(127.0f, max_scalar)
157 const float d = max_scalar / 127.0f;
158 const float id = max_scalar != 0.0f ? 127.0f / max_scalar : 0.0f;
162 const __m256 mul = _mm256_set1_ps(
id);
163 v0 = _mm256_mul_ps(v0, mul);
164 v1 = _mm256_mul_ps(v1, mul);
165 v2 = _mm256_mul_ps(v2, mul);
166 v3 = _mm256_mul_ps(v3, mul);
169 v0 = _mm256_round_ps(v0, _MM_ROUND_NEAREST);
170 v1 = _mm256_round_ps(v1, _MM_ROUND_NEAREST);
171 v2 = _mm256_round_ps(v2, _MM_ROUND_NEAREST);
172 v3 = _mm256_round_ps(v3, _MM_ROUND_NEAREST);
174 __m256i i0 = _mm256_cvtps_epi32(v0);
175 __m256i i1 = _mm256_cvtps_epi32(v1);
176 __m256i i2 = _mm256_cvtps_epi32(v2);
177 __m256i i3 = _mm256_cvtps_epi32(v3);
180 i0 = _mm256_packs_epi32(i0, i1);
181 i2 = _mm256_packs_epi32(i2, i3);
182 i0 = _mm256_packs_epi16(i0, i2);
184 const __m256i perm = _mm256_setr_epi32(0, 4, 1, 5, 2, 6, 3, 7);
185 i0 = _mm256_permutevar8x32_epi32(i0, perm);
186 _mm256_storeu_si256((__m256i *)y[i].qs, i0);
188 __m128i ni0 = _mm256_castsi256_si128(i0);
189 __m128i ni1 = _mm256_extractf128_si256(i0, 1);
190 __m128i ni2 = _mm256_castsi256_si128(i1);
191 __m128i ni3 = _mm256_extractf128_si256(i1, 1);
192 __m128i ni4 = _mm256_castsi256_si128(i2);
193 __m128i ni5 = _mm256_extractf128_si256(i2, 1);
194 __m128i ni6 = _mm256_castsi256_si128(i3);
195 __m128i ni7 = _mm256_extractf128_si256(i3, 1);
197 ni0 = _mm_packs_epi32(ni0, ni1);
198 ni2 = _mm_packs_epi32(ni2, ni3);
199 ni4 = _mm_packs_epi32(ni4, ni5);
200 ni6 = _mm_packs_epi32(ni6, ni7);
202 ni0 = _mm_packs_epi16(ni0, ni2);
203 ni4 = _mm_packs_epi16(ni4, ni6);
205 _mm_storeu_si128((__m128i *)(y[i].qs + 0), ni0);
206 _mm_storeu_si128((__m128i *)(y[i].qs + 16), ni4);
210 for (
int i = 0; i < nb; i++) {
211 const float *xb = x + i *
QK8_0;
215 for (
int j = 0; j <
QK8_0; j++) {
216 float av = xb[j] >= 0 ? xb[j] : -xb[j];
217 if (av > amax) amax = av;
221 float d = amax / 127.0f;
222 float id = d != 0.0f ? 127.0f / amax : 0.0f;
228 for (
int j = 0; j <
QK8_0; j++) {
229 float v = xb[j] *
id;
231 if (q > 127) q = 127;
232 if (q < -127) q = -127;
233 y[i].
qs[j] = (int8_t)q;
258 const size_t row_bytes_in = (size_t)k *
sizeof(
float);
261 uint8_t *out = (uint8_t *)vy;
262 const uint8_t *in = (
const uint8_t *)x;
264 for (
int row = 0; row < num_rows; ++row) {
266 (
const float *)(in + row * row_bytes_in),
267 (
void *)(out + row * row_bytes_out),
286 const size_t row_bytes_in = (size_t)k *
sizeof(
float);
289 const size_t row_bytes_out = (size_t)(k / 256) *
sizeof(
block_q8_K);
291 uint8_t *out = (uint8_t *)vy;
292 const uint8_t *in = (
const uint8_t *)x;
294 for (
int row = 0; row < num_rows; ++row) {
296 (
const float *)(in + row * row_bytes_in),
297 (
void *)(out + row * row_bytes_out),
322 const int blocks_per_row = K /
QK8_0;
324 for (
int row = 0; row < M; row++) {
327 for (
int b = 0; b < blocks_per_row; b++) {
328 const block_q8_0 *block = &blocks[row * blocks_per_row + b];
330 const float *xp = &x[b *
QK8_0];
332 for (
int i = 0; i <
QK8_0; i++) {
333 sum += d * (float)block->
qs[i] * xp[i];
345void gemv_q8_0_avx512(
float *y,
351 const int blocks_per_row = K /
QK8_0;
353 for (
int row = 0; row < M; row++) {
354 __m512 acc = _mm512_setzero_ps();
356 for (
int b = 0; b < blocks_per_row; b++) {
357 const block_q8_0 *block = &blocks[row * blocks_per_row + b];
359 const float *xp = &x[b *
QK8_0];
362 for (
int chunk = 0; chunk < 2; chunk++) {
364 __m128i q8 = _mm_loadu_si128((
const __m128i *)&block->
qs[chunk * 16]);
367 __m512i q32 = _mm512_cvtepi8_epi32(q8);
370 __m512 w = _mm512_mul_ps(_mm512_cvtepi32_ps(q32), vscale);
373 __m512 x_vec = _mm512_loadu_ps(&xp[chunk * 16]);
376 acc = _mm512_fmadd_ps(w, x_vec, acc);
380 y[row] = _mm512_reduce_add_ps(acc);
397#if defined(__AVX2__) && !defined(__AVX512F__)
400static inline float hsum_avx2_q8(__m256 v) {
401 __m128 lo = _mm256_castps256_ps128(v);
402 __m128 hi = _mm256_extractf128_ps(v, 1);
403 lo = _mm_add_ps(lo, hi);
404 __m128 shuf = _mm_shuffle_ps(lo, lo, _MM_SHUFFLE(2, 3, 0, 1));
405 __m128 sums = _mm_add_ps(lo, shuf);
406 shuf = _mm_movehl_ps(shuf, sums);
407 sums = _mm_add_ss(sums, shuf);
408 return _mm_cvtss_f32(sums);
417void gemv_q8_0_avx2(
float *y,
423 const int blocks_per_row = K /
QK8_0;
425 for (
int row = 0; row < M; row++) {
426 __m256 acc = _mm256_setzero_ps();
428 for (
int b = 0; b < blocks_per_row; b++) {
429 const block_q8_0 *block = &blocks[row * blocks_per_row + b];
431 const __m256 vscale = _mm256_set1_ps(d);
432 const float *xp = &x[b *
QK8_0];
438 __m128i q8 = _mm_loadl_epi64((
const __m128i *)&block->
qs[0]);
439 __m256i q32 = _mm256_cvtepi8_epi32(q8);
440 __m256 wf = _mm256_mul_ps(_mm256_cvtepi32_ps(q32), vscale);
441 __m256 xv = _mm256_loadu_ps(&xp[0]);
442 acc = _mm256_fmadd_ps(wf, xv, acc);
447 __m128i q8 = _mm_loadl_epi64((
const __m128i *)&block->
qs[8]);
448 __m256i q32 = _mm256_cvtepi8_epi32(q8);
449 __m256 wf = _mm256_mul_ps(_mm256_cvtepi32_ps(q32), vscale);
450 __m256 xv = _mm256_loadu_ps(&xp[8]);
451 acc = _mm256_fmadd_ps(wf, xv, acc);
456 __m128i q8 = _mm_loadl_epi64((
const __m128i *)&block->
qs[16]);
457 __m256i q32 = _mm256_cvtepi8_epi32(q8);
458 __m256 wf = _mm256_mul_ps(_mm256_cvtepi32_ps(q32), vscale);
459 __m256 xv = _mm256_loadu_ps(&xp[16]);
460 acc = _mm256_fmadd_ps(wf, xv, acc);
465 __m128i q8 = _mm_loadl_epi64((
const __m128i *)&block->
qs[24]);
466 __m256i q32 = _mm256_cvtepi8_epi32(q8);
467 __m256 wf = _mm256_mul_ps(_mm256_cvtepi32_ps(q32), vscale);
468 __m256 xv = _mm256_loadu_ps(&xp[24]);
469 acc = _mm256_fmadd_ps(wf, xv, acc);
473 y[row] = hsum_avx2_q8(acc);
490#if defined(__AVX__) && !defined(__AVX2__) && !defined(__AVX512F__)
493static inline float hsum_sse_q8(__m128 v) {
494 __m128 shuf = _mm_shuffle_ps(v, v, _MM_SHUFFLE(2, 3, 0, 1));
495 __m128 sums = _mm_add_ps(v, shuf);
496 shuf = _mm_movehl_ps(shuf, sums);
497 sums = _mm_add_ss(sums, shuf);
498 return _mm_cvtss_f32(sums);
507void gemv_q8_0_avx(
float *y,
513 const int blocks_per_row = K /
QK8_0;
515 for (
int row = 0; row < M; row++) {
517 __m128 acc0 = _mm_setzero_ps();
518 __m128 acc1 = _mm_setzero_ps();
519 __m128 acc2 = _mm_setzero_ps();
520 __m128 acc3 = _mm_setzero_ps();
522 for (
int b = 0; b < blocks_per_row; b++) {
523 const block_q8_0 *block = &blocks[row * blocks_per_row + b];
525 const float *xp = &x[b *
QK8_0];
526 const __m128 vscale = _mm_set1_ps(d);
529 __m128i q8_0 = _mm_loadu_si128((
const __m128i *)&block->
qs[0]);
530 __m128i q8_1 = _mm_loadu_si128((
const __m128i *)&block->
qs[16]);
535 __m128i q16 = _mm_cvtepi8_epi16(q8_0);
536 __m128i q32 = _mm_cvtepi16_epi32(q16);
537 __m128 w = _mm_mul_ps(_mm_cvtepi32_ps(q32), vscale);
538 __m128 vx = _mm_loadu_ps(&xp[0]);
539 acc0 = _mm_add_ps(acc0, _mm_mul_ps(w, vx));
544 __m128i q16 = _mm_cvtepi8_epi16(q8_0);
545 __m128i q32 = _mm_cvtepi16_epi32(_mm_srli_si128(q16, 8));
546 __m128 w = _mm_mul_ps(_mm_cvtepi32_ps(q32), vscale);
547 __m128 vx = _mm_loadu_ps(&xp[4]);
548 acc1 = _mm_add_ps(acc1, _mm_mul_ps(w, vx));
553 __m128i q8_shifted = _mm_srli_si128(q8_0, 8);
554 __m128i q16 = _mm_cvtepi8_epi16(q8_shifted);
555 __m128i q32 = _mm_cvtepi16_epi32(q16);
556 __m128 w = _mm_mul_ps(_mm_cvtepi32_ps(q32), vscale);
557 __m128 vx = _mm_loadu_ps(&xp[8]);
558 acc2 = _mm_add_ps(acc2, _mm_mul_ps(w, vx));
563 __m128i q8_shifted = _mm_srli_si128(q8_0, 8);
564 __m128i q16 = _mm_cvtepi8_epi16(q8_shifted);
565 __m128i q32 = _mm_cvtepi16_epi32(_mm_srli_si128(q16, 8));
566 __m128 w = _mm_mul_ps(_mm_cvtepi32_ps(q32), vscale);
567 __m128 vx = _mm_loadu_ps(&xp[12]);
568 acc3 = _mm_add_ps(acc3, _mm_mul_ps(w, vx));
574 __m128i q16 = _mm_cvtepi8_epi16(q8_1);
575 __m128i q32 = _mm_cvtepi16_epi32(q16);
576 __m128 w = _mm_mul_ps(_mm_cvtepi32_ps(q32), vscale);
577 __m128 vx = _mm_loadu_ps(&xp[16]);
578 acc0 = _mm_add_ps(acc0, _mm_mul_ps(w, vx));
583 __m128i q16 = _mm_cvtepi8_epi16(q8_1);
584 __m128i q32 = _mm_cvtepi16_epi32(_mm_srli_si128(q16, 8));
585 __m128 w = _mm_mul_ps(_mm_cvtepi32_ps(q32), vscale);
586 __m128 vx = _mm_loadu_ps(&xp[20]);
587 acc1 = _mm_add_ps(acc1, _mm_mul_ps(w, vx));
592 __m128i q8_shifted = _mm_srli_si128(q8_1, 8);
593 __m128i q16 = _mm_cvtepi8_epi16(q8_shifted);
594 __m128i q32 = _mm_cvtepi16_epi32(q16);
595 __m128 w = _mm_mul_ps(_mm_cvtepi32_ps(q32), vscale);
596 __m128 vx = _mm_loadu_ps(&xp[24]);
597 acc2 = _mm_add_ps(acc2, _mm_mul_ps(w, vx));
602 __m128i q8_shifted = _mm_srli_si128(q8_1, 8);
603 __m128i q16 = _mm_cvtepi8_epi16(q8_shifted);
604 __m128i q32 = _mm_cvtepi16_epi32(_mm_srli_si128(q16, 8));
605 __m128 w = _mm_mul_ps(_mm_cvtepi32_ps(q32), vscale);
606 __m128 vx = _mm_loadu_ps(&xp[28]);
607 acc3 = _mm_add_ps(acc3, _mm_mul_ps(w, vx));
612 __m128 sum01 = _mm_add_ps(acc0, acc1);
613 __m128 sum23 = _mm_add_ps(acc2, acc3);
614 __m128 sum = _mm_add_ps(sum01, sum23);
616 y[row] = hsum_sse_q8(sum);
621#if defined(__SSE4_1__)
622#include <immintrin.h>
625#define SSE_Q8_BLOCK(q8_reg, offset, xp, d_val, acc) do { \
626 __m128 vx = _mm_loadu_ps(&(xp)[offset]); \
627 __m128i qw = _mm_cvtepi8_epi32(_mm_srli_si128(q8_reg, offset)); \
628 __m128 vw = _mm_cvtepi32_ps(qw); \
629 acc = _mm_add_ps(acc, _mm_mul_ps(_mm_mul_ps(vw, vx), _mm_set1_ps(d_val))); \
632void gemv_q8_0_sse(
float *y,
638 const int blocks_per_row = K /
QK8_0;
640 for (
int row = 0; row < M; row++) {
641 __m128 acc = _mm_setzero_ps();
643 for (
int b = 0; b < blocks_per_row; b++) {
644 const block_q8_0 *block = &blocks[row * blocks_per_row + b];
646 const float *xp = &x[b *
QK8_0];
649 __m128i q8_0 = _mm_loadu_si128((
const __m128i *)&block->
qs[0]);
650 __m128i q8_1 = _mm_loadu_si128((
const __m128i *)&block->
qs[16]);
653 SSE_Q8_BLOCK(q8_0, 0, xp, d_val, acc);
654 SSE_Q8_BLOCK(q8_0, 4, xp, d_val, acc);
655 SSE_Q8_BLOCK(q8_0, 8, xp, d_val, acc);
656 SSE_Q8_BLOCK(q8_0, 12, xp, d_val, acc);
659 const float *xp1 = xp + 16;
660 SSE_Q8_BLOCK(q8_1, 0, xp1, d_val, acc);
661 SSE_Q8_BLOCK(q8_1, 4, xp1, d_val, acc);
662 SSE_Q8_BLOCK(q8_1, 8, xp1, d_val, acc);
663 SSE_Q8_BLOCK(q8_1, 12, xp1, d_val, acc);
667 acc = _mm_add_ps(acc, _mm_shuffle_ps(acc, acc, _MM_SHUFFLE(1, 0, 3, 2)));
668 acc = _mm_add_ps(acc, _mm_shuffle_ps(acc, acc, _MM_SHUFFLE(0, 1, 0, 1)));
669 _mm_store_ss(&y[row], acc);
705#if defined(__AVX512F__)
706 gemv_q8_0_avx512(y, W, x, M, K);
707#elif defined(__AVX2__)
708 gemv_q8_0_avx2(y, W, x, M, K);
709#elif defined(__AVX__)
710 gemv_q8_0_avx(y, W, x, M, K);
711#elif defined(__SSE4_1__)
712 gemv_q8_0_sse(y, W, x, M, K);
730 for (
int n = 0; n < N; n++) {
731 gemv_q8_0(&Y[n * M], W, &X[n * K], M, K);
756 for (
int m = 0; m < M; m++) {
757 gemv_q8_0(&
C[(
size_t)m * N], B, &A[(
size_t)m * K], N, K);
759 for (
int n = 0; n < N; n++)
C[(
size_t)m * N + n] += bias[n];
764#if defined(__AVX512F__)
765static inline float hsum512_ps_q80(__m512 v)
767 return _mm512_reduce_add_ps(v);
770static void gemm_nt_q8_0_m4n4_avx512(
const float *A,
777 const int blocks_per_row = K /
QK8_0;
778 const int M4 = M & ~3;
779 const int N4 = N & ~3;
781 for (
int m = 0; m < M4; m += 4) {
782 for (
int n = 0; n < N4; n += 4) {
783 __m512 acc00 = _mm512_setzero_ps(), acc01 = _mm512_setzero_ps();
784 __m512 acc02 = _mm512_setzero_ps(), acc03 = _mm512_setzero_ps();
785 __m512 acc10 = _mm512_setzero_ps(), acc11 = _mm512_setzero_ps();
786 __m512 acc12 = _mm512_setzero_ps(), acc13 = _mm512_setzero_ps();
787 __m512 acc20 = _mm512_setzero_ps(), acc21 = _mm512_setzero_ps();
788 __m512 acc22 = _mm512_setzero_ps(), acc23 = _mm512_setzero_ps();
789 __m512 acc30 = _mm512_setzero_ps(), acc31 = _mm512_setzero_ps();
790 __m512 acc32 = _mm512_setzero_ps(), acc33 = _mm512_setzero_ps();
792 const block_q8_0 *b0 = blocks + (size_t)(n + 0) * blocks_per_row;
793 const block_q8_0 *b1 = blocks + (size_t)(n + 1) * blocks_per_row;
794 const block_q8_0 *b2 = blocks + (size_t)(n + 2) * blocks_per_row;
795 const block_q8_0 *b3 = blocks + (size_t)(n + 3) * blocks_per_row;
796 const float *a0 = A + (size_t)(m + 0) * K;
797 const float *a1 = A + (size_t)(m + 1) * K;
798 const float *a2 = A + (size_t)(m + 2) * K;
799 const float *a3 = A + (size_t)(m + 3) * K;
801 for (
int ib = 0; ib < blocks_per_row; ++ib) {
802 const int k0 = ib *
QK8_0;
803 for (
int chunk = 0; chunk < 2; ++chunk) {
804 const int off = chunk * 16;
805 const __m512 x0 = _mm512_loadu_ps(a0 + k0 + off);
806 const __m512 x1 = _mm512_loadu_ps(a1 + k0 + off);
807 const __m512 x2 = _mm512_loadu_ps(a2 + k0 + off);
808 const __m512 x3 = _mm512_loadu_ps(a3 + k0 + off);
810 __m128i q0 = _mm_loadu_si128((
const __m128i *)&b0[ib].qs[off]);
811 __m128i q1 = _mm_loadu_si128((
const __m128i *)&b1[ib].qs[off]);
812 __m128i q2 = _mm_loadu_si128((
const __m128i *)&b2[ib].qs[off]);
813 __m128i q3 = _mm_loadu_si128((
const __m128i *)&b3[ib].qs[off]);
815 __m512 w0 = _mm512_mul_ps(_mm512_cvtepi32_ps(_mm512_cvtepi8_epi32(q0)),
817 __m512 w1 = _mm512_mul_ps(_mm512_cvtepi32_ps(_mm512_cvtepi8_epi32(q1)),
819 __m512 w2 = _mm512_mul_ps(_mm512_cvtepi32_ps(_mm512_cvtepi8_epi32(q2)),
821 __m512 w3 = _mm512_mul_ps(_mm512_cvtepi32_ps(_mm512_cvtepi8_epi32(q3)),
824 acc00 = _mm512_fmadd_ps(w0, x0, acc00);
825 acc01 = _mm512_fmadd_ps(w1, x0, acc01);
826 acc02 = _mm512_fmadd_ps(w2, x0, acc02);
827 acc03 = _mm512_fmadd_ps(w3, x0, acc03);
828 acc10 = _mm512_fmadd_ps(w0, x1, acc10);
829 acc11 = _mm512_fmadd_ps(w1, x1, acc11);
830 acc12 = _mm512_fmadd_ps(w2, x1, acc12);
831 acc13 = _mm512_fmadd_ps(w3, x1, acc13);
832 acc20 = _mm512_fmadd_ps(w0, x2, acc20);
833 acc21 = _mm512_fmadd_ps(w1, x2, acc21);
834 acc22 = _mm512_fmadd_ps(w2, x2, acc22);
835 acc23 = _mm512_fmadd_ps(w3, x2, acc23);
836 acc30 = _mm512_fmadd_ps(w0, x3, acc30);
837 acc31 = _mm512_fmadd_ps(w1, x3, acc31);
838 acc32 = _mm512_fmadd_ps(w2, x3, acc32);
839 acc33 = _mm512_fmadd_ps(w3, x3, acc33);
843 const float b00 = bias ? bias[n + 0] : 0.0f;
844 const float b01 = bias ? bias[n + 1] : 0.0f;
845 const float b02 = bias ? bias[n + 2] : 0.0f;
846 const float b03 = bias ? bias[n + 3] : 0.0f;
847 C[(size_t)(m + 0) * N + n + 0] = hsum512_ps_q80(acc00) + b00;
848 C[(size_t)(m + 0) * N + n + 1] = hsum512_ps_q80(acc01) + b01;
849 C[(size_t)(m + 0) * N + n + 2] = hsum512_ps_q80(acc02) + b02;
850 C[(size_t)(m + 0) * N + n + 3] = hsum512_ps_q80(acc03) + b03;
851 C[(size_t)(m + 1) * N + n + 0] = hsum512_ps_q80(acc10) + b00;
852 C[(size_t)(m + 1) * N + n + 1] = hsum512_ps_q80(acc11) + b01;
853 C[(size_t)(m + 1) * N + n + 2] = hsum512_ps_q80(acc12) + b02;
854 C[(size_t)(m + 1) * N + n + 3] = hsum512_ps_q80(acc13) + b03;
855 C[(size_t)(m + 2) * N + n + 0] = hsum512_ps_q80(acc20) + b00;
856 C[(size_t)(m + 2) * N + n + 1] = hsum512_ps_q80(acc21) + b01;
857 C[(size_t)(m + 2) * N + n + 2] = hsum512_ps_q80(acc22) + b02;
858 C[(size_t)(m + 2) * N + n + 3] = hsum512_ps_q80(acc23) + b03;
859 C[(size_t)(m + 3) * N + n + 0] = hsum512_ps_q80(acc30) + b00;
860 C[(size_t)(m + 3) * N + n + 1] = hsum512_ps_q80(acc31) + b01;
861 C[(size_t)(m + 3) * N + n + 2] = hsum512_ps_q80(acc32) + b02;
862 C[(size_t)(m + 3) * N + n + 3] = hsum512_ps_q80(acc33) + b03;
870 for (
int m = M4; m < M; ++m) {
871 gemv_q8_0(&
C[(
size_t)m * N], B, &A[(
size_t)m * K], N, K);
873 for (
int n = 0; n < N; ++n)
C[(
size_t)m * N + n] += bias[n];
885#if defined(__AVX512F__)
887 gemm_nt_q8_0_m4n4_avx512(A, B, bias,
C, M, N, K);
913 const int blocks_per_row = K /
QK8_0;
916 memset(dX, 0, K *
sizeof(
float));
919 for (
int row = 0; row < M; row++) {
920 const float dy = dY[row];
922 for (
int b = 0; b < blocks_per_row; b++) {
923 const block_q8_0 *block = &blocks[row * blocks_per_row + b];
925 float *dxp = &dX[b *
QK8_0];
927 for (
int i = 0; i <
QK8_0; i++) {
928 dxp[i] += d * (float)block->
qs[i] * dy;
938void gemv_q8_0_backward_avx512(
float *dX,
944 const int blocks_per_row = K /
QK8_0;
947 memset(dX, 0, K *
sizeof(
float));
949 for (
int row = 0; row < M; row++) {
950 const __m512 vdy = _mm512_set1_ps(dY[row]);
952 for (
int b = 0; b < blocks_per_row; b++) {
953 const block_q8_0 *block = &blocks[row * blocks_per_row + b];
955 float *dxp = &dX[b *
QK8_0];
958 for (
int chunk = 0; chunk < 2; chunk++) {
960 __m128i q8 = _mm_loadu_si128((
const __m128i *)&block->
qs[chunk * 16]);
961 __m512i q32 = _mm512_cvtepi8_epi32(q8);
962 __m512 w = _mm512_mul_ps(_mm512_cvtepi32_ps(q32), vscale);
965 __m512 grad = _mm512_mul_ps(w, vdy);
968 __m512 dx_cur = _mm512_loadu_ps(&dxp[chunk * 16]);
969 _mm512_storeu_ps(&dxp[chunk * 16], _mm512_add_ps(dx_cur, grad));
985 gemv_q8_0_backward_avx512(dX, W, dY, M, K);
999 for (
int n = 0; n < N; n++) {
1008float dot_q8_0(
const void *w_q8_0,
const float *x,
int K)
1015#if defined(__AVX__) || defined(__AVX2__) || defined(__AVX512F__)
1017static inline float hsum_float_8_q8_0(
const __m256 x)
1019 __m128 res = _mm256_extractf128_ps(x, 1);
1020 res = _mm_add_ps(res, _mm256_castps256_ps128(x));
1021 res = _mm_add_ps(res, _mm_movehl_ps(res, res));
1022 res = _mm_add_ss(res, _mm_movehdup_ps(res));
1023 return _mm_cvtss_f32(res);
1026#if defined(__AVX2__) || defined(__AVX512F__)
1027static inline __m256 sum_i16_pairs_float_q8_0_avx2(
const __m256i x)
1029 const __m256i ones = _mm256_set1_epi16(1);
1030 const __m256i summed_pairs = _mm256_madd_epi16(ones, x);
1031 return _mm256_cvtepi32_ps(summed_pairs);
1034static inline __m256 mul_sum_us8_pairs_float_q8_0_avx2(
const __m256i ax,
const __m256i sy)
1036#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
1037 const __m256i zero = _mm256_setzero_si256();
1038 const __m256i summed_pairs = _mm256_dpbusd_epi32(zero, ax, sy);
1039 return _mm256_cvtepi32_ps(summed_pairs);
1040#elif defined(__AVXVNNI__)
1041 const __m256i zero = _mm256_setzero_si256();
1042 const __m256i summed_pairs = _mm256_dpbusd_avx_epi32(zero, ax, sy);
1043 return _mm256_cvtepi32_ps(summed_pairs);
1045 const __m256i dot = _mm256_maddubs_epi16(ax, sy);
1046 return sum_i16_pairs_float_q8_0_avx2(dot);
1050static inline __m256 mul_sum_i8_pairs_float_q8_0_avx2(
const __m256i x,
const __m256i y)
1053 const __m256i zero = _mm256_setzero_si256();
1054 const __m256i summed_pairs = _mm256_dpbssd_epi32(zero, x, y);
1055 return _mm256_cvtepi32_ps(summed_pairs);
1057 const __m256i ax = _mm256_sign_epi8(x, x);
1058 const __m256i sy = _mm256_sign_epi8(y, x);
1059 return mul_sum_us8_pairs_float_q8_0_avx2(ax, sy);
1062#elif defined(__AVX__)
1063static inline __m128i mul_add_epi8_sse_q8_0(
const __m128i x,
const __m128i y)
1065 const __m128i ax = _mm_sign_epi8(x, x);
1066 const __m128i sy = _mm_sign_epi8(y, x);
1067 return _mm_maddubs_epi16(ax, sy);
1070static inline __m256 sum_i16_pairs_float_q8_0_avx(
const __m128i xh,
const __m128i xl)
1072 const __m128i ones = _mm_set1_epi16(1);
1073 const __m128i summed_pairsl = _mm_madd_epi16(ones, xl);
1074 const __m128i summed_pairsh = _mm_madd_epi16(ones, xh);
1075 const __m256i summed_pairs = _mm256_insertf128_si256(
1076 _mm256_castsi128_si256(summed_pairsl),
1080 return _mm256_cvtepi32_ps(summed_pairs);
1083static inline __m256 mul_sum_i8_quad_float_q8_0_avx(
const __m128i x_1_0,
1084 const __m128i x_1_1,
1085 const __m128i x_2_0,
1086 const __m128i x_2_1,
1087 const __m128i y_1_0,
1088 const __m128i y_1_1,
1089 const __m128i y_2_0,
1090 const __m128i y_2_1)
1092 const __m128i mone = _mm_set1_epi16(1);
1094 const __m128i p16_1_0 = mul_add_epi8_sse_q8_0(x_1_0, y_1_0);
1095 const __m128i p16_1_1 = mul_add_epi8_sse_q8_0(x_1_1, y_1_1);
1096 const __m128i p16_2_0 = mul_add_epi8_sse_q8_0(x_2_0, y_2_0);
1097 const __m128i p16_2_1 = mul_add_epi8_sse_q8_0(x_2_1, y_2_1);
1098 const __m128i p_1_0 = _mm_madd_epi16(p16_1_0, mone);
1099 const __m128i p_1_1 = _mm_madd_epi16(p16_1_1, mone);
1100 const __m128i p_2_0 = _mm_madd_epi16(p16_2_0, mone);
1101 const __m128i p_2_1 = _mm_madd_epi16(p16_2_1, mone);
1102 const __m128i p_1 = _mm_add_epi32(p_1_0, p_1_1);
1103 const __m128i p_2 = _mm_add_epi32(p_2_0, p_2_1);
1104 const __m256i packed = _mm256_insertf128_si256(_mm256_castsi128_si256(p_1), p_2, 1);
1105 return _mm256_cvtepi32_ps(packed);
1108static inline __m256 quad_fp16_delta_float_q8_0_avx(uint16_t x0, uint16_t y0, uint16_t x1, uint16_t y1)
1110 return _mm256_set_m128(
1142 const int qk =
QK8_0;
1143 const int nb = n / qk;
1150 for (
int ib = 0; ib < nb; ib++) {
1153 for (
int j = 0; j < qk; j++) {
1154 sumi += x[ib].
qs[j] * y[ib].
qs[j];
1163#if defined(__ARM_NEON) || defined(__aarch64__)
1164void vec_dot_q8_0_q8_0_neon(
int n,
float *s,
const void *vx,
const void *vy)
1166 const int qk =
QK8_0;
1167 const int nb = n / qk;
1174 for (
int ib = 0; ib < nb; ib++) {
1175 const int8x16_t x0 = vld1q_s8(&x[ib].qs[0]);
1176 const int8x16_t x1 = vld1q_s8(&x[ib].qs[16]);
1177 const int8x16_t y0 = vld1q_s8(&y[ib].qs[0]);
1178 const int8x16_t y1 = vld1q_s8(&y[ib].qs[16]);
1180 int32x4_t acc = vdupq_n_s32(0);
1182 const int16x8_t p0 = vmull_s8(vget_low_s8(x0), vget_low_s8(y0));
1183 const int16x8_t p1 = vmull_s8(vget_high_s8(x0), vget_high_s8(y0));
1184 const int16x8_t p2 = vmull_s8(vget_low_s8(x1), vget_low_s8(y1));
1185 const int16x8_t p3 = vmull_s8(vget_high_s8(x1), vget_high_s8(y1));
1187 acc = vaddq_s32(acc, vpaddlq_s16(p0));
1188 acc = vaddq_s32(acc, vpaddlq_s16(p1));
1189 acc = vaddq_s32(acc, vpaddlq_s16(p2));
1190 acc = vaddq_s32(acc, vpaddlq_s16(p3));
1193 vst1q_s32(lanes, acc);
1194 const int sumi = lanes[0] + lanes[1] + lanes[2] + lanes[3];
1203#if defined(__AVX2__) && !defined(__AVX512F__)
1204void vec_dot_q8_0_q8_0_avx2(
int n,
float *s,
const void *vx,
const void *vy)
1206 const int qk =
QK8_0;
1207 const int nb = n / qk;
1214 __m256 acc = _mm256_setzero_ps();
1216 for (; ib < nb; ++ib) {
1218 const __m256i qx = _mm256_loadu_si256((
const __m256i *)x[ib].qs);
1219 const __m256i qy = _mm256_loadu_si256((
const __m256i *)y[ib].qs);
1220 const __m256 q = mul_sum_i8_pairs_float_q8_0_avx2(qx, qy);
1222 acc = _mm256_fmadd_ps(d, q, acc);
1224 acc = _mm256_add_ps(_mm256_mul_ps(d, q), acc);
1228 sumf = hsum_float_8_q8_0(acc);
1237void vec_dot_q8_0_q8_0_avx512(
int n,
float *s,
const void *vx,
const void *vy)
1239 const int qk =
QK8_0;
1240 const int nb = n / qk;
1251 __m256 acc = _mm256_setzero_ps();
1252 for (
int ib = 0; ib < nb; ++ib) {
1253 const __m256 d = _mm256_set1_ps(
1255 const __m256i qx = _mm256_loadu_si256((
const __m256i *)x[ib].qs);
1256 const __m256i qy = _mm256_loadu_si256((
const __m256i *)y[ib].qs);
1257 const __m256 q = mul_sum_i8_pairs_float_q8_0_avx2(qx, qy);
1259 acc = _mm256_fmadd_ps(d, q, acc);
1261 acc = _mm256_add_ps(_mm256_mul_ps(d, q), acc);
1265 *s = hsum_float_8_q8_0(acc);
1269#if defined(__AVX__) && !defined(__AVX2__) && !defined(__AVX512F__)
1273void vec_dot_q8_0_q8_0_avx(
int n,
float *s,
const void *vx,
const void *vy)
1275 const int qk =
QK8_0;
1276 const int nb = n / qk;
1282 __m256 accum = _mm256_setzero_ps();
1284 for (; ib + 1 < nb; ib += 2) {
1285 const __m128i qx_1_0 = _mm_loadu_si128((
const __m128i *)x[ib].qs);
1286 const __m128i qx_1_1 = _mm_loadu_si128((
const __m128i *)x[ib].qs + 1);
1287 const __m128i qx_2_0 = _mm_loadu_si128((
const __m128i *)x[ib + 1].qs);
1288 const __m128i qx_2_1 = _mm_loadu_si128((
const __m128i *)x[ib + 1].qs + 1);
1289 const __m128i qy_1_0 = _mm_loadu_si128((
const __m128i *)y[ib].qs);
1290 const __m128i qy_1_1 = _mm_loadu_si128((
const __m128i *)y[ib].qs + 1);
1291 const __m128i qy_2_0 = _mm_loadu_si128((
const __m128i *)y[ib + 1].qs);
1292 const __m128i qy_2_1 = _mm_loadu_si128((
const __m128i *)y[ib + 1].qs + 1);
1294 const __m256 p = mul_sum_i8_quad_float_q8_0_avx(
1295 qx_1_0, qx_1_1, qx_2_0, qx_2_1,
1296 qy_1_0, qy_1_1, qy_2_0, qy_2_1
1298 const __m256 deltas = quad_fp16_delta_float_q8_0_avx(
1299 x[ib].d, y[ib].d, x[ib + 1].d, y[ib + 1].d
1301 accum = _mm256_add_ps(_mm256_mul_ps(deltas, p), accum);
1304 float sumf = hsum_float_8_q8_0(accum);
1305 for (; ib < nb; ++ib) {
1307 for (
int j = 0; j < qk; ++j) {
1308 sumi += x[ib].
qs[j] * y[ib].
qs[j];
1316#if defined(__SSE4_1__) && !defined(__AVX__)
1320void vec_dot_q8_0_q8_0_sse(
int n,
float *s,
const void *vx,
const void *vy)
1322 const int qk =
QK8_0;
1323 const int nb = n / qk;
1330 for (
int ib = 0; ib < nb; ib++) {
1333 __m128i acc_lo = _mm_setzero_si128();
1334 __m128i acc_hi = _mm_setzero_si128();
1337 for (
int j = 0; j < 32; j += 8) {
1339 __m128i x8 = _mm_loadl_epi64((
const __m128i *)&x[ib].qs[j]);
1340 __m128i y8 = _mm_loadl_epi64((
const __m128i *)&y[ib].qs[j]);
1343 __m128i x16 = _mm_cvtepi8_epi16(x8);
1344 __m128i y16 = _mm_cvtepi8_epi16(y8);
1347 __m128i prod = _mm_madd_epi16(x16, y16);
1350 acc_lo = _mm_add_epi32(acc_lo, prod);
1354 acc_lo = _mm_add_epi32(acc_lo, _mm_shuffle_epi32(acc_lo, _MM_SHUFFLE(1, 0, 3, 2)));
1355 acc_lo = _mm_add_epi32(acc_lo, _mm_shuffle_epi32(acc_lo, _MM_SHUFFLE(0, 1, 0, 1)));
1356 int sumi = _mm_extract_epi32(acc_lo, 0);
1358 sumf += d * (float)sumi;
1375 vec_dot_q8_0_q8_0_avx512(n, s, vx, vy);
1376#elif defined(__AVX2__)
1377 vec_dot_q8_0_q8_0_avx2(n, s, vx, vy);
1378#elif defined(__ARM_NEON) || defined(__aarch64__)
1379 vec_dot_q8_0_q8_0_neon(n, s, vx, vy);
1380#elif defined(__AVX__)
1381 vec_dot_q8_0_q8_0_avx(n, s, vx, vy);
1382#elif defined(__SSE4_1__)
1383 vec_dot_q8_0_q8_0_sse(n, s, vx, vy);
1412 const int blocks_per_row = K /
QK8_0;
1414 for (
int row = 0; row < M; row++) {
1416 &w_blocks[row * blocks_per_row],
1433#if defined(__AVX2__) || defined(__AVX512F__)
1441 const int nb = K /
QK8_0;
1443 for (; row + 3 < M; row += 4) {
1444 __m256 acc0 = _mm256_setzero_ps();
1445 __m256 acc1 = _mm256_setzero_ps();
1446 __m256 acc2 = _mm256_setzero_ps();
1447 __m256 acc3 = _mm256_setzero_ps();
1448 const block_q8_0 *w0 = w + (size_t)(row + 0) * (size_t)nb;
1449 const block_q8_0 *w1 = w + (size_t)(row + 1) * (size_t)nb;
1450 const block_q8_0 *w2 = w + (size_t)(row + 2) * (size_t)nb;
1451 const block_q8_0 *w3 = w + (size_t)(row + 3) * (size_t)nb;
1453 for (
int ib = 0; ib < nb; ++ib) {
1454 const __m256i qx = _mm256_loadu_si256((
const __m256i *)x[ib].qs);
1456 const __m256i qw0 = _mm256_loadu_si256((
const __m256i *)w0[ib].qs);
1457 const __m256i qw1 = _mm256_loadu_si256((
const __m256i *)w1[ib].qs);
1458 const __m256i qw2 = _mm256_loadu_si256((
const __m256i *)w2[ib].qs);
1459 const __m256i qw3 = _mm256_loadu_si256((
const __m256i *)w3[ib].qs);
1460 const __m256 p0 = mul_sum_i8_pairs_float_q8_0_avx2(qw0, qx);
1461 const __m256 p1 = mul_sum_i8_pairs_float_q8_0_avx2(qw1, qx);
1462 const __m256 p2 = mul_sum_i8_pairs_float_q8_0_avx2(qw2, qx);
1463 const __m256 p3 = mul_sum_i8_pairs_float_q8_0_avx2(qw3, qx);
1469 acc0 = _mm256_fmadd_ps(d0, p0, acc0);
1470 acc1 = _mm256_fmadd_ps(d1, p1, acc1);
1471 acc2 = _mm256_fmadd_ps(d2, p2, acc2);
1472 acc3 = _mm256_fmadd_ps(d3, p3, acc3);
1474 acc0 = _mm256_add_ps(_mm256_mul_ps(d0, p0), acc0);
1475 acc1 = _mm256_add_ps(_mm256_mul_ps(d1, p1), acc1);
1476 acc2 = _mm256_add_ps(_mm256_mul_ps(d2, p2), acc2);
1477 acc3 = _mm256_add_ps(_mm256_mul_ps(d3, p3), acc3);
1480 y[row + 0] = hsum_float_8_q8_0(acc0);
1481 y[row + 1] = hsum_float_8_q8_0(acc1);
1482 y[row + 2] = hsum_float_8_q8_0(acc2);
1483 y[row + 3] = hsum_float_8_q8_0(acc3);
1506 int M,
int N,
int K)
1508#if defined(__AVX2__) || defined(__AVX512F__)
1510 const int nb = K /
QK8_0;
1512 for (
int m = 0; m < M; ++m) {
1514 a + (
size_t)m * (
size_t)nb, N, K);
1521 const int nb = K /
QK8_0;
1523 for (; m + 1 < M; m += 2) {
1524 const block_q8_0 *a0 = a + (size_t)(m + 0) * (size_t)nb;
1525 const block_q8_0 *a1 = a + (size_t)(m + 1) * (size_t)nb;
1527 for (; n + 3 < N; n += 4) {
1528 const block_q8_0 *w0 = w + (size_t)(n + 0) * (size_t)nb;
1529 const block_q8_0 *w1 = w + (size_t)(n + 1) * (size_t)nb;
1530 const block_q8_0 *w2 = w + (size_t)(n + 2) * (size_t)nb;
1531 const block_q8_0 *w3 = w + (size_t)(n + 3) * (size_t)nb;
1532 __m256 acc00 = _mm256_setzero_ps();
1533 __m256 acc01 = _mm256_setzero_ps();
1534 __m256 acc02 = _mm256_setzero_ps();
1535 __m256 acc03 = _mm256_setzero_ps();
1536 __m256 acc10 = _mm256_setzero_ps();
1537 __m256 acc11 = _mm256_setzero_ps();
1538 __m256 acc12 = _mm256_setzero_ps();
1539 __m256 acc13 = _mm256_setzero_ps();
1541 for (
int ib = 0; ib < nb; ++ib) {
1542 const __m256i qa0 = _mm256_loadu_si256((
const __m256i *)a0[ib].qs);
1543 const __m256i qa1 = _mm256_loadu_si256((
const __m256i *)a1[ib].qs);
1544 const __m256i qw0 = _mm256_loadu_si256((
const __m256i *)w0[ib].qs);
1545 const __m256i qw1 = _mm256_loadu_si256((
const __m256i *)w1[ib].qs);
1546 const __m256i qw2 = _mm256_loadu_si256((
const __m256i *)w2[ib].qs);
1547 const __m256i qw3 = _mm256_loadu_si256((
const __m256i *)w3[ib].qs);
1554 const __m256 p00 = mul_sum_i8_pairs_float_q8_0_avx2(qw0, qa0);
1555 const __m256 p01 = mul_sum_i8_pairs_float_q8_0_avx2(qw1, qa0);
1556 const __m256 p02 = mul_sum_i8_pairs_float_q8_0_avx2(qw2, qa0);
1557 const __m256 p03 = mul_sum_i8_pairs_float_q8_0_avx2(qw3, qa0);
1558 const __m256 p10 = mul_sum_i8_pairs_float_q8_0_avx2(qw0, qa1);
1559 const __m256 p11 = mul_sum_i8_pairs_float_q8_0_avx2(qw1, qa1);
1560 const __m256 p12 = mul_sum_i8_pairs_float_q8_0_avx2(qw2, qa1);
1561 const __m256 p13 = mul_sum_i8_pairs_float_q8_0_avx2(qw3, qa1);
1563 acc00 = _mm256_fmadd_ps(_mm256_set1_ps(dw0 * da0), p00, acc00);
1564 acc01 = _mm256_fmadd_ps(_mm256_set1_ps(dw1 * da0), p01, acc01);
1565 acc02 = _mm256_fmadd_ps(_mm256_set1_ps(dw2 * da0), p02, acc02);
1566 acc03 = _mm256_fmadd_ps(_mm256_set1_ps(dw3 * da0), p03, acc03);
1567 acc10 = _mm256_fmadd_ps(_mm256_set1_ps(dw0 * da1), p10, acc10);
1568 acc11 = _mm256_fmadd_ps(_mm256_set1_ps(dw1 * da1), p11, acc11);
1569 acc12 = _mm256_fmadd_ps(_mm256_set1_ps(dw2 * da1), p12, acc12);
1570 acc13 = _mm256_fmadd_ps(_mm256_set1_ps(dw3 * da1), p13, acc13);
1572 acc00 = _mm256_add_ps(_mm256_mul_ps(_mm256_set1_ps(dw0 * da0), p00), acc00);
1573 acc01 = _mm256_add_ps(_mm256_mul_ps(_mm256_set1_ps(dw1 * da0), p01), acc01);
1574 acc02 = _mm256_add_ps(_mm256_mul_ps(_mm256_set1_ps(dw2 * da0), p02), acc02);
1575 acc03 = _mm256_add_ps(_mm256_mul_ps(_mm256_set1_ps(dw3 * da0), p03), acc03);
1576 acc10 = _mm256_add_ps(_mm256_mul_ps(_mm256_set1_ps(dw0 * da1), p10), acc10);
1577 acc11 = _mm256_add_ps(_mm256_mul_ps(_mm256_set1_ps(dw1 * da1), p11), acc11);
1578 acc12 = _mm256_add_ps(_mm256_mul_ps(_mm256_set1_ps(dw2 * da1), p12), acc12);
1579 acc13 = _mm256_add_ps(_mm256_mul_ps(_mm256_set1_ps(dw3 * da1), p13), acc13);
1583 C[(size_t)(m + 0) * (size_t)ldc + (n + 0)] = hsum_float_8_q8_0(acc00);
1584 C[(size_t)(m + 0) * (size_t)ldc + (n + 1)] = hsum_float_8_q8_0(acc01);
1585 C[(size_t)(m + 0) * (size_t)ldc + (n + 2)] = hsum_float_8_q8_0(acc02);
1586 C[(size_t)(m + 0) * (size_t)ldc + (n + 3)] = hsum_float_8_q8_0(acc03);
1587 C[(size_t)(m + 1) * (size_t)ldc + (n + 0)] = hsum_float_8_q8_0(acc10);
1588 C[(size_t)(m + 1) * (size_t)ldc + (n + 1)] = hsum_float_8_q8_0(acc11);
1589 C[(size_t)(m + 1) * (size_t)ldc + (n + 2)] = hsum_float_8_q8_0(acc12);
1590 C[(size_t)(m + 1) * (size_t)ldc + (n + 3)] = hsum_float_8_q8_0(acc13);
1594 w + (
size_t)n * (
size_t)nb, a0, N - n, K);
1596 w + (
size_t)n * (
size_t)nb, a1, N - n, K);
1601 a + (
size_t)m * (
size_t)nb, N, K);
1604 const int nb = K /
QK8_0;
1606 for (
int m = 0; m < M; ++m) {
1608 a + (
size_t)m * (
size_t)nb, N, K);
1616 int M,
int N,
int K)
1637 if (!y || !W || !x_q8 || M <= 0 || K <= 0)
return;
1638 if (ith < 0 || nth <= 0 || ith >= nth)
return;
1640 const int dr = (M + nth - 1) / nth;
1641 const int r0 = dr * ith;
1642 const int r1 = (r0 + dr < M) ? (r0 + dr) : M;
1644 if (r0 >= M)
return;
1648 const int blocks_per_row = K /
QK8_0;
1650 for (
int row = r0; row < r1; row++) {
1652 &w_blocks[row * blocks_per_row],
1669 if (!y || !W || !x_q8 || M <= 0 || K <= 0)
return;
1670 if (ith < 0 || nth <= 0 || ith >= nth)
return;
1672 const int dr = (M + nth - 1) / nth;
1673 const int r0 = dr * ith;
1674 const int r1 = (r0 + dr < M) ? (r0 + dr) : M;
1676 if (r0 >= M)
return;
1680 const int blocks_per_row = K /
QK8_0;
1682#if defined(__AVX__) || defined(__SSE4_1__)
1684 const int PREFETCH_ROWS = 4;
1685 for (
int p = 0; p < PREFETCH_ROWS && r0 + p < r1; ++p) {
1686 const char *row_ptr = (
const char *)(w_blocks + (r0 + p) * blocks_per_row);
1687 _mm_prefetch(row_ptr, _MM_HINT_T0);
1688 _mm_prefetch(row_ptr + 64, _MM_HINT_T0);
1691 for (
int row = r0; row < r1; ++row) {
1693 if (row + PREFETCH_ROWS < r1) {
1694 const char *pf = (
const char *)(w_blocks + (row + PREFETCH_ROWS) * blocks_per_row);
1695 _mm_prefetch(pf, _MM_HINT_T0);
1696 _mm_prefetch(pf + 64, _MM_HINT_T0);
1700 &w_blocks[row * blocks_per_row],
1705 for (
int row = r0; row < r1; row++) {
1707 &w_blocks[row * blocks_per_row],
1722 if (!y || !W || !x || M <= 0 || K <= 0)
return;
1723 if (ith < 0 || nth <= 0 || ith >= nth)
return;
1725 const int dr = (M + nth - 1) / nth;
1726 const int r0 = dr * ith;
1727 const int r1 = (r0 + dr < M) ? (r0 + dr) : M;
1729 if (r0 >= M)
return;
1732 const int blocks_per_row = K /
QK8_0;
1734#if defined(__AVX__) || defined(__SSE4_1__)
1735 const int PREFETCH_ROWS = 4;
1736 for (
int p = 0; p < PREFETCH_ROWS && r0 + p < r1; ++p) {
1737 const char *row_ptr = (
const char *)(blocks + (r0 + p) * blocks_per_row);
1738 _mm_prefetch(row_ptr, _MM_HINT_T0);
1739 _mm_prefetch(row_ptr + 64, _MM_HINT_T0);
1742 for (
int row = r0; row < r1; ++row) {
1743 if (row + PREFETCH_ROWS < r1) {
1744 const char *pf = (
const char *)(blocks + (row + PREFETCH_ROWS) * blocks_per_row);
1745 _mm_prefetch(pf, _MM_HINT_T0);
1746 _mm_prefetch(pf + 64, _MM_HINT_T0);
1750#if defined(__AVX512F__)
1751 gemv_q8_0_avx512(&y[row],
1752 (
const char *)blocks + row * blocks_per_row *
sizeof(
block_q8_0),
1754#elif defined(__AVX2__)
1755 gemv_q8_0_avx2(&y[row],
1756 (
const char *)blocks + row * blocks_per_row *
sizeof(
block_q8_0),
1758#elif defined(__AVX__)
1759 gemv_q8_0_avx(&y[row],
1760 (
const char *)blocks + row * blocks_per_row *
sizeof(
block_q8_0),
1762#elif defined(__SSE4_1__)
1763 gemv_q8_0_sse(&y[row],
1764 (
const char *)blocks + row * blocks_per_row *
sizeof(
block_q8_0),
1768 (
const char *)blocks + row * blocks_per_row *
sizeof(
block_q8_0),
1773 for (
int row = r0; row < r1; row++) {
1775 (
const char *)blocks + row * blocks_per_row *
sizeof(
block_q8_0),
CPU feature detection and dispatch macros.
static int ck_env_truthy_or_qwen3vl_ocr_profile(const char *name)
Quantization block structures for weight-only quantization.
#define CK_FP16_TO_FP32(x)
#define CK_FP32_TO_FP16(x)
void gemv_q8_0_parallel_simd(float *y, const void *W, const float *x, int M, int K, int ith, int nth)
Parallel SIMD GEMV for Q8_0 weights x FP32 input with prefetching.
static void gemm_nt_q8_0_rowloop(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
Matrix-matrix multiply: C[M,N] = A[M,K] @ B[N,K]^T + bias.
static int ck_q8_0_fp32_m4n4_enabled(void)
void gemv_q8_0(float *y, const void *W, const float *x, int M, int K)
Auto-dispatch GEMV for Q8_0 weights based on CPU features.
void quantize_batch_q8_0(const float *x, void *vy, int num_rows, int k)
Batch quantize FP32 to Q8_0 format (row-major output)
static int ck_q8_0_q8_0_debug_ref(void)
static int ck_nearest_int_q8_0(float fval)
void gemv_q8_0_q8_0_parallel_simd(float *y, const void *W, const void *x_q8, int M, int K, int ith, int nth)
Parallel SIMD GEMV for Q8_0 x Q8_0 with prefetching.
void quantize_batch_q8_k(const float *x, void *vy, int num_rows, int k)
Batch quantize FP32 to Q8_K format (row-major output)
void quantize_row_q8_k(const float *x, void *vy, int k)
void gemm_q8_0_backward(float *dX, const void *W, const float *dY, int M, int N, int K)
Batched backward pass.
void vec_dot_q8_0_q8_0_ref(int n, float *s, const void *vx, const void *vy)
Quantized dot product: Q8_0 weights x Q8_0 input (scalar reference)
void gemv_q8_0_q8_0_parallel(float *y, const void *W, const void *x_q8, int M, int K, int ith, int nth)
Parallel reference GEMV for Q8_0 x Q8_0.
void gemv_q8_0_q8_0_x4(float *y, const void *W, const void *x_q8, int M, int K)
void gemm_q8_0(float *Y, const void *W, const float *X, int M, int N, int K)
Matrix-matrix multiply with Q8_0 weights.
void gemm_nt_q8_0(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
void gemm_q8_0_q8_0_m2n4(float *C, const void *W, const void *A_q8, int M, int N, int K)
static int ck_q8_0_debug_ref(void)
void gemv_q8_0_q8_0(float *y, const void *W, const void *x_q8, int M, int K)
Matrix-vector multiply with Q8_0 weights and Q8_0 input.
void quantize_row_q8_0(const float *x, void *vy, int k)
Quantize FP32 to Q8_0 format (scalar reference)
void vec_dot_q8_0_q8_0(int n, float *s, const void *vx, const void *vy)
Auto-dispatch quantized dot product Q8_0 x Q8_0.
void gemm_q8_0_q8_0_m2n4_strided(float *C, int ldc, const void *W, const void *A_q8, int M, int N, int K)
void gemv_q8_0_backward_ref(float *dX, const void *W, const float *dY, int M, int K)
Backward pass: compute input gradient (scalar reference)
float dot_q8_0(const void *w_q8_0, const float *x, int K)
void gemv_q8_0_backward(float *dX, const void *W, const float *dY, int M, int K)
Auto-dispatch backward.
void gemv_q8_0_ref(float *y, const void *W, const float *x, int M, int K)
Matrix-vector multiply with Q8_0 weights (scalar reference)