← Back to C-Kernel-Engine Docs Doxygen Source Documentation
 
Loading...
Searching...
No Matches
gemm_kernels_q4k_q8k_avx2.c
Go to the documentation of this file.
1/**
2 * @file gemm_kernels_q4k_q8k_avx2.c
3 * @brief AVX2 Q4_K x Q8_K matvec kernel (inference only)
4 *
5 * CK-ENGINE KERNEL RULES:
6 * =======================
7 * 1. NO malloc/free - memory via bump allocator, pointers passed in
8 * 2. NO OpenMP - parallelization at orchestrator/codegen layer
9 * 3. API must define: inputs, outputs, workspace, and memory layouts
10 * 4. Pure computation - deterministic, no side effects
11 *
12 * After changes: make test && make llamacpp-parity-full
13 *
14 * Requires AVX2 for 256-bit integer operations.
15 */
16
17#include <stddef.h>
18#include <stdint.h>
19#include <string.h>
20
21#include "ckernel_quant.h"
22
23#if defined(__AVX2__)
24#include <immintrin.h>
25#endif
26
27void gemv_q4_k_q8_k_ref(float *y,
28 const void *W,
29 const void *x_q8,
30 int M, int K);
31
32#if defined(__AVX2__)
33static inline float hsum256_ps_q4k(__m256 v) {
34 __m128 sum = _mm_add_ps(_mm256_castps256_ps128(v), _mm256_extractf128_ps(v, 1));
35 sum = _mm_add_ps(sum, _mm_movehl_ps(sum, sum));
36 sum = _mm_add_ss(sum, _mm_movehdup_ps(sum));
37 return _mm_cvtss_f32(sum);
38}
39
40static inline __m256i q4k_scale_pair_shuffle(int pair)
41{
42 const int lo = 2 * pair;
43 return _mm256_set1_epi16((short)(((lo + 1) << 8) | lo));
44}
45
46static float dot_q4_k_q8_k_avx2(const block_q4_K *w,
47 const block_q8_K *x,
48 int k)
49{
50 const int nb = k / QK_K;
51 __m256 acc = _mm256_setzero_ps();
52 __m128 acc_min = _mm_setzero_ps();
53 const __m256i nibble_mask = _mm256_set1_epi8(0x0F);
54 const uint32_t mask_6bit = UINT32_C(0x3f3f3f3f);
55 const uint32_t mask_4bit = UINT32_C(0x0f0f0f0f);
56 const uint32_t mask_2bit = UINT32_C(0x03030303);
57
58 for (int i = 0; i < nb; ++i) {
59 uint32_t packed[4];
60 memcpy(packed, w[i].scales, 12);
61 packed[3] = ((packed[2] >> 4) & mask_4bit) |
62 (((packed[1] >> 6) & mask_2bit) << 4);
63 const uint32_t upper_scales = packed[1] & mask_6bit;
64 packed[1] = (packed[2] & mask_4bit) |
65 (((packed[0] >> 6) & mask_2bit) << 4);
66 packed[2] = upper_scales;
67 packed[0] &= mask_6bit;
68
69 const __m256i mins_and_scales = _mm256_cvtepu8_epi16(
70 _mm_set_epi32((int)packed[3], (int)packed[2],
71 (int)packed[1], (int)packed[0]));
72 const __m128i scale_bytes =
73 _mm256_extracti128_si256(mins_and_scales, 0);
74 const __m256i scales = _mm256_set_m128i(scale_bytes, scale_bytes);
75
76 const float d = CK_FP16_TO_FP32(w[i].d) * x[i].d;
77 const float dmin = -CK_FP16_TO_FP32(w[i].dmin) * x[i].d;
78
79 const __m128i q8_sums = _mm_hadd_epi16(
80 _mm_loadu_si128((const __m128i *)&x[i].bsums[0]),
81 _mm_loadu_si128((const __m128i *)&x[i].bsums[8]));
82 const __m128i min_products = _mm_madd_epi16(
83 _mm256_extracti128_si256(mins_and_scales, 1), q8_sums);
84 acc_min = _mm_fmadd_ps(
85 _mm_set1_ps(dmin), _mm_cvtepi32_ps(min_products), acc_min);
86
87 __m256i block_sum = _mm256_setzero_si256();
88 for (int group = 0; group < QK_K / 64; ++group) {
89 const __m256i packed = _mm256_loadu_si256(
90 (const __m256i *)&w[i].qs[group * 32]);
91 const __m256i q4_lo = _mm256_and_si256(packed, nibble_mask);
92 const __m256i q4_hi = _mm256_and_si256(
93 _mm256_srli_epi16(packed, 4), nibble_mask);
94 const __m256i q8_lo = _mm256_loadu_si256(
95 (const __m256i *)&x[i].qs[group * 64]);
96 const __m256i q8_hi = _mm256_loadu_si256(
97 (const __m256i *)&x[i].qs[group * 64 + 32]);
98 __m256i lo = _mm256_maddubs_epi16(q4_lo, q8_lo);
99 __m256i hi = _mm256_maddubs_epi16(q4_hi, q8_hi);
100 const __m256i scale_lo = _mm256_shuffle_epi8(
101 scales, q4k_scale_pair_shuffle(2 * group));
102 const __m256i scale_hi = _mm256_shuffle_epi8(
103 scales, q4k_scale_pair_shuffle(2 * group + 1));
104 lo = _mm256_madd_epi16(scale_lo, lo);
105 hi = _mm256_madd_epi16(scale_hi, hi);
106 block_sum = _mm256_add_epi32(block_sum, _mm256_add_epi32(lo, hi));
107 }
108 acc = _mm256_fmadd_ps(
109 _mm256_set1_ps(d), _mm256_cvtepi32_ps(block_sum), acc);
110 }
111
112 acc_min = _mm_add_ps(acc_min, _mm_movehl_ps(acc_min, acc_min));
113 acc_min = _mm_add_ss(acc_min, _mm_movehdup_ps(acc_min));
114 return hsum256_ps_q4k(acc) + _mm_cvtss_f32(acc_min);
115}
116#endif
117
119 const void *W,
120 const void *x_q8,
121 int M, int K)
122{
123#if defined(__AVX2__)
124 if (!y || !W || !x_q8 || M <= 0 || K <= 0) {
125 return;
126 }
127
128 const block_q4_K *blocks = (const block_q4_K *)W;
129 const block_q8_K *x = (const block_q8_K *)x_q8;
130 const int blocks_per_row = K / QK_K;
131
132 for (int row = 0; row < M; ++row) {
133 const block_q4_K *w_row = blocks + (size_t)row * (size_t)blocks_per_row;
134 y[row] = dot_q4_k_q8_k_avx2(w_row, x, K);
135 }
136#else
137 gemv_q4_k_q8_k_ref(y, W, x_q8, M, K);
138#endif
139}
Quantization block structures for weight-only quantization.
#define CK_FP16_TO_FP32(x)
#define QK_K
void gemv_q4_k_q8_k_avx2(float *y, const void *W, const void *x_q8, int M, int K)
void gemv_q4_k_q8_k_ref(float *y, const void *W, const void *x_q8, int M, int K)