← Back to C-Kernel-Engine Docs Doxygen Source Documentation
 
Loading...
Searching...
No Matches
gemm_batch_int8.c File Reference

Batch GEMM kernels for quantized weights with INT8 activations. More...

#include <stdint.h>
#include <stddef.h>
#include <string.h>
#include <math.h>
#include "ckernel_quant.h"

Go to the source code of this file.

Macros

#define AMX_TILE_K   64
 
#define AMX_TILE_M   16
 
#define AMX_TILE_N   16
 
#define HAS_AMX   0
 
#define QK5_0   32 /* Q5_0: 32 weights per block */
 
#define QK8_0   32 /* Q8_0: 32 weights per block */
 

Functions

const char * gemm_batch_int8_impl_name (void)
 Get the best implementation name for logging/debugging.
 
void gemm_nt_q5_0_q8_0_ref (const void *A, const void *B, float *C, int M, int N, int K)
 Dispatcher for gemm_nt_q8_0_q8_0.
 
void gemm_nt_q8_0_q8_0 (const void *A, const void *B, const float *bias, float *C, int M, int N, int K)
 gemm_nt_q8_0_q8_0 with optional bias (matches header signature)
 
void gemm_nt_q8_0_q8_0_m2n4 (const void *A, const void *B, const float *bias, float *C, int M, int N, int K)
 
void gemm_nt_q8_0_q8_0_m2n4_tile (const void *A, const void *B, const float *bias, float *C, int M, int N, int K, int ldc)
 
void gemm_nt_q8_0_q8_0_ref (const void *A, const void *B, float *C, int M, int N, int K)
 Scalar reference: gemm_nt_q8_0_q8_0.
 
void gemm_q8_0_q8_0_m2n4 (float *C, const void *W, const void *A_q8, int M, int N, int K)
 
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_q8_0_x4 (float *y, const void *W, const void *x_q8, int M, int K)
 

Detailed Description

Batch GEMM kernels for quantized weights with INT8 activations.

CK-ENGINE KERNEL RULES:

  1. NO malloc/free - memory via bump allocator, pointers passed in
  2. NO OpenMP - parallelization at orchestrator/codegen layer
  3. API must define: inputs, outputs, workspace, and memory layouts
  4. Pure computation - deterministic, no side effects

After changes: make test && make llamacpp-parity-full

Implements batch matrix multiplication where:

  • Activations (A): Q8_0 quantized (INT8 + scale)
  • Weights (B): Q5_0 or Q8_0 quantized
  • Output (C): FP32

Operation: C[M,N] = A[M,K] @ B[N,K]^T (B is transposed/row-major weights)

Instruction Set Implementations:

  • Scalar: Reference implementation for correctness verification
  • AVX: 256-bit SIMD (8 floats, or 32 int8s)
  • AVX-512: 512-bit SIMD (16 floats, or 64 int8s)
  • AMX: Intel Advanced Matrix Extensions (tile-based, requires Sapphire Rapids+)

Design Philosophy:

  • Every kernel MUST produce bit-identical results to scalar reference
  • Comprehensive testing against llama.cpp ensures correctness
  • Performance optimizations never compromise accuracy
Author
C-Kernel-Engine Team
Date
2024

Definition in file gemm_batch_int8.c.

Macro Definition Documentation

◆ AMX_TILE_K

#define AMX_TILE_K   64

Definition at line 65 of file gemm_batch_int8.c.

◆ AMX_TILE_M

#define AMX_TILE_M   16

Definition at line 63 of file gemm_batch_int8.c.

◆ AMX_TILE_N

#define AMX_TILE_N   16

Definition at line 64 of file gemm_batch_int8.c.

◆ HAS_AMX

#define HAS_AMX   0

Definition at line 52 of file gemm_batch_int8.c.

◆ QK5_0

#define QK5_0   32 /* Q5_0: 32 weights per block */

Definition at line 60 of file gemm_batch_int8.c.

◆ QK8_0

#define QK8_0   32 /* Q8_0: 32 weights per block */

Definition at line 59 of file gemm_batch_int8.c.

Function Documentation

◆ gemm_batch_int8_impl_name()

const char * gemm_batch_int8_impl_name ( void  )

Get the best implementation name for logging/debugging.

Definition at line 522 of file gemm_batch_int8.c.

523{
524#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
525 return "AVX-512 VNNI";
526#elif HAS_AMX
527 /* The AMX entry points in this file currently fall back to AVX-512/ref. */
528 return "AMX fallback";
529#elif defined(__AVX512F__)
530 return "AVX-512";
531#elif defined(__AVX2__)
532 return "AVX2";
533#elif defined(__AVX__)
534 return "AVX";
535#else
536 return "Scalar";
537#endif
538}

◆ gemm_nt_q5_0_q8_0_ref()

void gemm_nt_q5_0_q8_0_ref ( const void *  A,
const void *  B,
float *  C,
int  M,
int  N,
int  K 
)

Dispatcher for gemm_nt_q8_0_q8_0.

Selects the best available implementation at runtime.

Scalar reference: gemm_nt_q5_0_q8_0

Q5_0 weight reconstruction: weight[j] = d * ((qs_nibble | (qh_bit << 4)) - 16)

For j in 0..15: use low nibble + qh bit j For j in 16..31: use high nibble + qh bit (j+16) -> actually bit (j) for j=16..31

Parameters
AInput activations [M, K] in Q8_0 format
BWeight matrix [N, K] in Q5_0 format
COutput matrix [M, N] in FP32
MNumber of tokens (batch size)
NNumber of output features
KNumber of input features (must be multiple of 32)

Definition at line 360 of file gemm_batch_int8.c.

365{
366 const int nb = K / QK5_0;
367 const block_q8_0 *a_blocks = (const block_q8_0 *)A;
368 const block_q5_0 *b_blocks = (const block_q5_0 *)B;
369
370 for (int m = 0; m < M; m++) {
371 const block_q8_0 *a_row = a_blocks + (size_t)m * nb;
372
373 for (int n = 0; n < N; n++) {
374 const block_q5_0 *b_row = b_blocks + (size_t)n * nb;
375 float sum = 0.0f;
376
377 for (int ib = 0; ib < nb; ib++) {
378 const float d_a = CK_FP16_TO_FP32(a_row[ib].d);
379 const float d_b = CK_FP16_TO_FP32(b_row[ib].d);
380 const float d = d_a * d_b;
381
382 /* Load high bits as 32-bit value */
383 uint32_t qh;
384 memcpy(&qh, b_row[ib].qh, sizeof(qh));
385
386 int32_t sumi = 0;
387
388 /* Process 32 weights: j=0..15 uses low nibble, j=16..31 uses high nibble */
389 for (int j = 0; j < 16; j++) {
390 /* First 16 weights: low nibble + qh bit j */
391 const uint8_t xh_0 = ((qh >> j) & 1) << 4;
392 const int8_t w0 = (int8_t)(((b_row[ib].qs[j] & 0x0F) | xh_0) - 16);
393
394 /* Second 16 weights: high nibble + qh bit (j+16) */
395 const uint8_t xh_1 = ((qh >> (j + 16)) & 1) << 4;
396 const int8_t w1 = (int8_t)(((b_row[ib].qs[j] >> 4) | xh_1) - 16);
397
398 /* Accumulate with activation values */
399 sumi += (int32_t)w0 * (int32_t)a_row[ib].qs[j];
400 sumi += (int32_t)w1 * (int32_t)a_row[ib].qs[j + 16];
401 }
402
403 sum += d * (float)sumi;
404 }
405
406 C[(size_t)m * N + n] = sum;
407 }
408 }
409}
#define CK_FP16_TO_FP32(x)
#define QK5_0
#define C(color)
Definition show_config.c:39
int8_t qs[32]

References C, CK_FP16_TO_FP32, QK5_0, and block_q8_0::qs.

◆ gemm_nt_q8_0_q8_0()

void gemm_nt_q8_0_q8_0 ( const void *  A,
const void *  B,
const float *  bias,
float *  C,
int  M,
int  N,
int  K 
)

gemm_nt_q8_0_q8_0 with optional bias (matches header signature)

C[m,n] = A[m,K] @ B[n,K]^T + bias[n]

Definition at line 552 of file gemm_batch_int8.c.

558{
559 /* First compute GEMM */
560#if defined(__AVX2__)
561 /*
562 * The production contract is bit-exact with llama.cpp's eight-lane
563 * FP32 accumulation tree. AVX-512/VNNI changes the integer dot, but it
564 * must not silently replace that FP32 reduction with the scalar-per-block
565 * candidate above. The AVX2-named entry point delegates each activation
566 * row to the certified x4 provider, which also uses VNNI instructions when
567 * they are available while preserving the declared reduction order.
568 */
569 gemm_nt_q8_0_q8_0_avx2(A, B, C, M, N, K);
570#elif defined(__AVX__)
571 gemm_nt_q8_0_q8_0_avx(A, B, C, M, N, K);
572#else
573 gemm_nt_q8_0_q8_0_ref(A, B, C, M, N, K);
574#endif
575
576 /* Add bias if provided */
577 if (bias != NULL) {
578 for (int m = 0; m < M; m++) {
579 for (int n = 0; n < N; n++) {
580 C[(size_t)m * N + n] += bias[n];
581 }
582 }
583 }
584}
void gemm_nt_q8_0_q8_0_ref(const void *A, const void *B, float *C, int M, int N, int K)
Scalar reference: gemm_nt_q8_0_q8_0.

References C, and gemm_nt_q8_0_q8_0_ref().

Referenced by ck_test_gemm_q8_0(), gemm_nt_q8_0_dispatch(), and gemm_nt_q8_0_mlp_dispatch().

◆ gemm_nt_q8_0_q8_0_m2n4()

void gemm_nt_q8_0_q8_0_m2n4 ( const void *  A,
const void *  B,
const float *  bias,
float *  C,
int  M,
int  N,
int  K 
)

Definition at line 586 of file gemm_batch_int8.c.

592{
593 gemm_q8_0_q8_0_m2n4(C, B, A, M, N, K);
594 if (bias != NULL) {
595 for (int m = 0; m < M; ++m) {
596 for (int n = 0; n < N; ++n) {
597 C[(size_t)m * (size_t)N + n] += bias[n];
598 }
599 }
600 }
601}
void gemm_q8_0_q8_0_m2n4(float *C, const void *W, const void *A_q8, int M, int N, int K)

References C, and gemm_q8_0_q8_0_m2n4().

◆ gemm_nt_q8_0_q8_0_m2n4_tile()

void gemm_nt_q8_0_q8_0_m2n4_tile ( const void *  A,
const void *  B,
const float *  bias,
float *  C,
int  M,
int  N,
int  K,
int  ldc 
)

Definition at line 603 of file gemm_batch_int8.c.

609{
610 gemm_q8_0_q8_0_m2n4_strided(C, ldc, B, A, M, N, K);
611 if (bias != NULL) {
612 for (int m = 0; m < M; ++m) {
613 for (int n = 0; n < N; ++n) {
614 C[(size_t)m * (size_t)ldc + n] += bias[n];
615 }
616 }
617 }
618}
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)

References C, and gemm_q8_0_q8_0_m2n4_strided().

◆ gemm_nt_q8_0_q8_0_ref()

void gemm_nt_q8_0_q8_0_ref ( const void *  A,
const void *  B,
float *  C,
int  M,
int  N,
int  K 
)

Scalar reference: gemm_nt_q8_0_q8_0.

C[m,n] = sum_k( dequant(A[m,k]) * dequant(B[n,k]) ) = sum_blocks( d_a * d_b * sum_j(a_qs[j] * b_qs[j]) )

Parameters
AInput activations [M, K] in Q8_0 format
BWeight matrix [N, K] in Q8_0 format (row-major, each row is one output)
COutput matrix [M, N] in FP32
MNumber of tokens (batch size)
NNumber of output features (rows in B)
KNumber of input features (must be multiple of 32)

Definition at line 118 of file gemm_batch_int8.c.

123{
124 const int nb = K / QK8_0; /* Number of blocks per row */
125 const block_q8_0 *a_blocks = (const block_q8_0 *)A;
126 const block_q8_0 *b_blocks = (const block_q8_0 *)B;
127
128 for (int m = 0; m < M; m++) {
129 const block_q8_0 *a_row = a_blocks + (size_t)m * nb;
130
131 for (int n = 0; n < N; n++) {
132 const block_q8_0 *b_row = b_blocks + (size_t)n * nb;
133 float sum = 0.0f;
134
135 for (int ib = 0; ib < nb; ib++) {
136 const float d_a = CK_FP16_TO_FP32(a_row[ib].d);
137 const float d_b = CK_FP16_TO_FP32(b_row[ib].d);
138 const float d = d_a * d_b;
139
140 int32_t sumi = 0;
141 for (int j = 0; j < QK8_0; j++) {
142 sumi += (int32_t)a_row[ib].qs[j] * (int32_t)b_row[ib].qs[j];
143 }
144
145 sum += d * (float)sumi;
146 }
147
148 C[(size_t)m * N + n] = sum;
149 }
150 }
151}
#define QK8_0

References C, CK_FP16_TO_FP32, QK8_0, and block_q8_0::qs.

Referenced by gemm_nt_q8_0_q8_0().

◆ gemm_q8_0_q8_0_m2n4()

void gemm_q8_0_q8_0_m2n4 ( float *  C,
const void *  W,
const void *  A_q8,
int  M,
int  N,
int  K 
)

Definition at line 1613 of file gemm_kernels_q8_0.c.

1617{
1618 gemm_q8_0_q8_0_m2n4_strided(C, N, W, A_q8, M, N, K);
1619}
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)

References C, and gemm_q8_0_q8_0_m2n4_strided().

Referenced by gemm_nt_q8_0_q8_0_m2n4().

◆ gemm_q8_0_q8_0_m2n4_strided()

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 
)

Definition at line 1502 of file gemm_kernels_q8_0.c.

1507{
1508#if defined(__AVX2__) || defined(__AVX512F__)
1509 if (ck_q8_0_q8_0_debug_ref() || (K % QK8_0) != 0) {
1510 const int nb = K / QK8_0;
1511 const block_q8_0 *a = (const block_q8_0 *)A_q8;
1512 for (int m = 0; m < M; ++m) {
1513 gemv_q8_0_q8_0(C + (size_t)m * (size_t)ldc, W,
1514 a + (size_t)m * (size_t)nb, N, K);
1515 }
1516 return;
1517 }
1518
1519 const block_q8_0 *w = (const block_q8_0 *)W;
1520 const block_q8_0 *a = (const block_q8_0 *)A_q8;
1521 const int nb = K / QK8_0;
1522 int m = 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;
1526 int n = 0;
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();
1540
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);
1548 const float da0 = CK_FP16_TO_FP32(a0[ib].d);
1549 const float da1 = CK_FP16_TO_FP32(a1[ib].d);
1550 const float dw0 = CK_FP16_TO_FP32(w0[ib].d);
1551 const float dw1 = CK_FP16_TO_FP32(w1[ib].d);
1552 const float dw2 = CK_FP16_TO_FP32(w2[ib].d);
1553 const float dw3 = CK_FP16_TO_FP32(w3[ib].d);
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);
1562#if defined(__FMA__)
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);
1571#else
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);
1580#endif
1581 }
1582
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);
1591 }
1592 if (n < N) {
1593 gemv_q8_0_q8_0(C + (size_t)(m + 0) * (size_t)ldc + n,
1594 w + (size_t)n * (size_t)nb, a0, N - n, K);
1595 gemv_q8_0_q8_0(C + (size_t)(m + 1) * (size_t)ldc + n,
1596 w + (size_t)n * (size_t)nb, a1, N - n, K);
1597 }
1598 }
1599 if (m < M) {
1600 gemv_q8_0_q8_0_x4(C + (size_t)m * (size_t)ldc, W,
1601 a + (size_t)m * (size_t)nb, N, K);
1602 }
1603#else
1604 const int nb = K / QK8_0;
1605 const block_q8_0 *a = (const block_q8_0 *)A_q8;
1606 for (int m = 0; m < M; ++m) {
1607 gemv_q8_0_q8_0(C + (size_t)m * (size_t)ldc, W,
1608 a + (size_t)m * (size_t)nb, N, K);
1609 }
1610#endif
1611}
#define QK8_0
static int ck_q8_0_q8_0_debug_ref(void)
void gemv_q8_0_q8_0_x4(float *y, const void *W, const void *x_q8, int M, int K)
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.

References C, CK_FP16_TO_FP32, ck_q8_0_q8_0_debug_ref(), gemv_q8_0_q8_0(), gemv_q8_0_q8_0_x4(), and QK8_0.

Referenced by gemm_nt_q8_0_q8_0_m2n4_tile(), and gemm_q8_0_q8_0_m2n4().

◆ gemv_q8_0_q8_0_x4()

void gemv_q8_0_q8_0_x4 ( float *  y,
const void *  W,
const void *  x_q8,
int  M,
int  K 
)

Definition at line 1428 of file gemm_kernels_q8_0.c.

1432{
1433#if defined(__AVX2__) || defined(__AVX512F__)
1434 if (ck_q8_0_q8_0_debug_ref() || (K % QK8_0) != 0) {
1435 gemv_q8_0_q8_0(y, W, x_q8, M, K);
1436 return;
1437 }
1438
1439 const block_q8_0 *w = (const block_q8_0 *)W;
1440 const block_q8_0 *x = (const block_q8_0 *)x_q8;
1441 const int nb = K / QK8_0;
1442 int row = 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;
1452
1453 for (int ib = 0; ib < nb; ++ib) {
1454 const __m256i qx = _mm256_loadu_si256((const __m256i *)x[ib].qs);
1455 const float dx = CK_FP16_TO_FP32(x[ib].d);
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);
1464 const __m256 d0 = _mm256_set1_ps(CK_FP16_TO_FP32(w0[ib].d) * dx);
1465 const __m256 d1 = _mm256_set1_ps(CK_FP16_TO_FP32(w1[ib].d) * dx);
1466 const __m256 d2 = _mm256_set1_ps(CK_FP16_TO_FP32(w2[ib].d) * dx);
1467 const __m256 d3 = _mm256_set1_ps(CK_FP16_TO_FP32(w3[ib].d) * dx);
1468#if defined(__FMA__)
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);
1473#else
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);
1478#endif
1479 }
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);
1484 }
1485 if (row < M) {
1486 gemv_q8_0_q8_0(y + row, w + (size_t)row * (size_t)nb,
1487 x, M - row, K);
1488 }
1489#else
1490 gemv_q8_0_q8_0(y, W, x_q8, M, K);
1491#endif
1492}

Referenced by gemm_q8_0_q8_0_m2n4_strided().