43#if defined(__AVX512F__) || defined(__AVX2__) || defined(__AVX__) || defined(__SSE4_1__)
48#if defined(__AMX_INT8__) && defined(__AVX512VNNI__)
73 const void *A_q8,
int M,
int N,
int K);
75#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
76static inline int32_t hsum256_epi32_q8_batch(__m256i v)
78 __m128i lo = _mm256_castsi256_si128(v);
79 __m128i hi = _mm256_extracti128_si256(v, 1);
80 __m128i sum = _mm_add_epi32(lo, hi);
81 sum = _mm_add_epi32(sum, _mm_srli_si128(sum, 8));
82 sum = _mm_add_epi32(sum, _mm_srli_si128(sum, 4));
83 return _mm_cvtsi128_si32(sum);
86static inline int32_t dot_q8_0_q8_0_32_vnni_i32(
const int8_t *a,
const int8_t *b)
88 const __m256i va = _mm256_loadu_si256((
const __m256i *)a);
89 const __m256i vb = _mm256_loadu_si256((
const __m256i *)b);
90 const __m256i va_u = _mm256_xor_si256(va, _mm256_set1_epi8((
char)0x80));
91 const __m256i dot_u_s = _mm256_dpbusd_epi32(_mm256_setzero_si256(), va_u, vb);
92 const __m256i sum_b = _mm256_dpbusd_epi32(_mm256_setzero_si256(), _mm256_set1_epi8(1), vb);
93 const __m256i correction = _mm256_slli_epi32(sum_b, 7);
94 return hsum256_epi32_q8_batch(_mm256_sub_epi32(dot_u_s, correction));
124 const int nb = K /
QK8_0;
128 for (
int m = 0; m < M; m++) {
129 const block_q8_0 *a_row = a_blocks + (size_t)m * nb;
131 for (
int n = 0; n < N; n++) {
132 const block_q8_0 *b_row = b_blocks + (size_t)n * nb;
135 for (
int ib = 0; ib < nb; ib++) {
138 const float d = d_a * d_b;
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];
145 sum += d * (float)sumi;
148 C[(size_t)m * N + n] = sum;
161void gemm_nt_q8_0_q8_0_avx2(
167 const int nb = K /
QK8_0;
170 for (
int m = 0; m < M; m++) {
171 const block_q8_0 *a_row = a_blocks + (size_t)m * nb;
177#if defined(__AVX__) && !defined(__AVX2__)
185void gemm_nt_q8_0_q8_0_avx(
191 const int nb = K /
QK8_0;
195 for (
int m = 0; m < M; m++) {
196 const block_q8_0 *a_row = a_blocks + (size_t)m * nb;
197 for (
int n = 0; n < N; n++) {
198 const block_q8_0 *b_row = b_blocks + (size_t)n * nb;
201 for (
int ib = 0; ib < nb; ib++) {
204 const int8_t *a_qs = a_row[ib].
qs;
205 const int8_t *b_qs = b_row[ib].
qs;
208 __m128i d0 = _mm_madd_epi16(
209 _mm_cvtepi8_epi16(_mm_loadl_epi64((
const __m128i *)(a_qs + 0))),
210 _mm_cvtepi8_epi16(_mm_loadl_epi64((
const __m128i *)(b_qs + 0))));
211 __m128i d1 = _mm_madd_epi16(
212 _mm_cvtepi8_epi16(_mm_loadl_epi64((
const __m128i *)(a_qs + 8))),
213 _mm_cvtepi8_epi16(_mm_loadl_epi64((
const __m128i *)(b_qs + 8))));
214 __m128i d2 = _mm_madd_epi16(
215 _mm_cvtepi8_epi16(_mm_loadl_epi64((
const __m128i *)(a_qs + 16))),
216 _mm_cvtepi8_epi16(_mm_loadl_epi64((
const __m128i *)(b_qs + 16))));
217 __m128i d3 = _mm_madd_epi16(
218 _mm_cvtepi8_epi16(_mm_loadl_epi64((
const __m128i *)(a_qs + 24))),
219 _mm_cvtepi8_epi16(_mm_loadl_epi64((
const __m128i *)(b_qs + 24))));
222 __m128i s4 = _mm_add_epi32(_mm_add_epi32(d0, d1),
223 _mm_add_epi32(d2, d3));
224 s4 = _mm_add_epi32(s4, _mm_srli_si128(s4, 8));
225 s4 = _mm_add_epi32(s4, _mm_srli_si128(s4, 4));
226 sum += d * (float)_mm_cvtsi128_si32(s4);
228 C[(size_t)m * N + n] = sum;
234#if defined(__AVX512F__)
241void gemm_nt_q8_0_q8_0_avx512(
247 const int nb = K /
QK8_0;
251 for (
int m = 0; m < M; m++) {
252 const block_q8_0 *a_row = a_blocks + (size_t)m * nb;
254 for (
int n = 0; n < N; n++) {
255 const block_q8_0 *b_row = b_blocks + (size_t)n * nb;
258 for (
int ib = 0; ib < nb; ib++) {
261 const float d = d_a * d_b;
264 __m256i va_256 = _mm256_loadu_si256((
const __m256i *)a_row[ib].qs);
265 __m256i vb_256 = _mm256_loadu_si256((
const __m256i *)b_row[ib].qs);
268 __m512i va_16 = _mm512_cvtepi8_epi16(va_256);
269 __m512i vb_16 = _mm512_cvtepi8_epi16(vb_256);
272 __m512i prod = _mm512_mullo_epi16(va_16, vb_16);
275 __m512i sum_32 = _mm512_madd_epi16(prod, _mm512_set1_epi16(1));
278 int32_t sumi = _mm512_reduce_add_epi32(sum_32);
280 sum += d * (float)sumi;
283 C[(size_t)m * N + n] = sum;
288#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
298void gemm_nt_q8_0_q8_0_vnni(
304 const int nb = K /
QK8_0;
308 for (
int m = 0; m < M; m++) {
309 const block_q8_0 *a_row = a_blocks + (size_t)m * nb;
311 for (
int n = 0; n < N; n++) {
312 const block_q8_0 *b_row = b_blocks + (size_t)n * nb;
315 for (
int ib = 0; ib < nb; ib++) {
318 const int32_t sumi = dot_q8_0_q8_0_32_vnni_i32(a_row[ib].qs, b_row[ib].qs);
319 sum += (d_a * d_b) * (
float)sumi;
322 C[(size_t)m * N + n] = sum;
366 const int nb = K /
QK5_0;
370 for (
int m = 0; m < M; m++) {
371 const block_q8_0 *a_row = a_blocks + (size_t)m * nb;
373 for (
int n = 0; n < N; n++) {
374 const block_q5_0 *b_row = b_blocks + (size_t)n * nb;
377 for (
int ib = 0; ib < nb; ib++) {
380 const float d = d_a * d_b;
384 memcpy(&qh, b_row[ib].qh,
sizeof(qh));
389 for (
int j = 0; j < 16; 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);
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);
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];
403 sum += d * (float)sumi;
406 C[(size_t)m * N + n] = sum;
434 uint8_t reserved[14];
439static void amx_tile_config_init(
void)
441 static __thread
int initialized = 0;
442 if (initialized)
return;
444 tile_config_t tc = {0};
452 tc.rows[0] = 16; tc.colsb[0] = 64;
453 tc.rows[1] = 16; tc.colsb[1] = 64;
454 tc.rows[2] = 16; tc.colsb[2] = 64;
455 tc.rows[3] = 16; tc.colsb[3] = 64;
456 tc.rows[4] = 16; tc.colsb[4] = 64;
457 tc.rows[5] = 16; tc.colsb[5] = 64;
458 tc.rows[6] = 16; tc.colsb[6] = 64;
459 tc.rows[7] = 16; tc.colsb[7] = 64;
461 _tile_loadconfig(&tc);
474void gemm_nt_q8_0_q8_0_amx(
480 amx_tile_config_init();
490 gemm_nt_q8_0_q8_0_avx512(A, B,
C, M, N, K);
493void gemm_nt_q5_0_q8_0_amx(
499 amx_tile_config_init();
524#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
525 return "AVX-512 VNNI";
528 return "AMX fallback";
529#elif defined(__AVX512F__)
531#elif defined(__AVX2__)
533#elif defined(__AVX__)
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);
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];
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];
608 int M,
int N,
int K,
int ldc)
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];
Quantization block structures for weight-only quantization.
#define CK_FP16_TO_FP32(x)
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_ref(const void *A, const void *B, float *C, int M, int N, int K)
Scalar reference: gemm_nt_q8_0_q8_0.
void gemv_q8_0_q8_0_x4(float *y, const void *W, const void *x_q8, int M, int K)
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_q8_0_q8_0_m2n4(float *C, const void *W, const void *A_q8, int M, int N, int K)
const char * gemm_batch_int8_impl_name(void)
Get the best implementation name for logging/debugging.
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_q8_0_q8_0_m2n4_strided(float *C, int ldc, const void *W, const void *A_q8, int M, int N, int K)
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)