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

GEMM/GEMV kernels with Q4_1 quantized weights. More...

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

Go to the source code of this file.

Functions

float dot_q4_1 (const void *w_q4_1, const float *x, int K)
 
void gemm_nt_q4_1 (const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
 GEMM with transposed Q4_1 weights: C = A @ B^T.
 
void gemm_q4_1 (float *Y, const void *W, const float *X, int M, int N, int K)
 Matrix-matrix multiply with Q4_1 weights.
 
void gemm_q4_1_backward (float *dX, const void *W, const float *dY, int M, int N, int K)
 Batched backward pass.
 
void gemv_q4_1 (float *y, const void *W, const float *x, int M, int K)
 Auto-dispatch GEMV.
 
void gemv_q4_1_backward (float *dX, const void *W, const float *dY, int M, int K)
 Auto-dispatch backward.
 
void gemv_q4_1_backward_ref (float *dX, const void *W, const float *dY, int M, int K)
 Backward pass: compute input gradient.
 
void gemv_q4_1_ref (float *y, const void *W, const float *x, int M, int K)
 Matrix-vector multiply with Q4_1 weights (scalar reference)
 

Detailed Description

GEMM/GEMV kernels with Q4_1 quantized weights.

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

Q4_1 Format:

  • 32 weights per block
  • 1 FP16 scale (d) per block
  • 1 FP16 minimum (m) per block
  • 20 bytes per 32 weights = 5.0 bits/weight

Dequantization: w = d * q + m where q is the 4-bit unsigned value (0-15)

Operations: Forward: Y = W @ X (W is Q4_1, X and Y are FP32) Backward: dX = W^T @ dY (gradient w.r.t. input)

Definition in file gemm_kernels_q4_1.c.

Function Documentation

◆ dot_q4_1()

float dot_q4_1 ( const void *  w_q4_1,
const float *  x,
int  K 
)

Definition at line 299 of file gemm_kernels_q4_1.c.

300{
301 float result;
302 gemv_q4_1(&result, w_q4_1, x, 1, K);
303 return result;
304}
void gemv_q4_1(float *y, const void *W, const float *x, int M, int K)
Auto-dispatch GEMV.

References gemv_q4_1().

◆ gemm_nt_q4_1()

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

GEMM with transposed Q4_1 weights: C = A @ B^T.

Parameters
AInput activations [M x K], row-major FP32
BWeight matrix in Q4_1 format [N x K], row-major quantized
biasOptional bias [N], NULL if not used
COutput [M x N], row-major FP32
MBatch size (number of tokens)
NOutput dimension
KInput dimension

Definition at line 256 of file gemm_kernels_q4_1.c.

261{
262 const block_q4_1 *blocks = (const block_q4_1 *)B;
263 const int blocks_per_row = K / QK4_1;
264
265 for (int m = 0; m < M; m++) {
266 const float *a_row = &A[m * K];
267
268 for (int n = 0; n < N; n++) {
269 float sum = 0.0f;
270
271 for (int b = 0; b < blocks_per_row; b++) {
272 const block_q4_1 *block = &blocks[n * blocks_per_row + b];
273 const float d = CK_FP16_TO_FP32(block->d);
274 const float min = CK_FP16_TO_FP32(block->m);
275 const float *ap = &a_row[b * QK4_1];
276
277 for (int i = 0; i < QK4_1 / 2; i++) {
278 const uint8_t packed = block->qs[i];
279 const int q0 = (packed & 0x0F);
280 const int q1 = (packed >> 4);
281
282 const float w0 = d * (float)q0 + min;
283 const float w1 = d * (float)q1 + min;
284
285 sum += w0 * ap[2 * i + 0];
286 sum += w1 * ap[2 * i + 1];
287 }
288 }
289
290 C[m * N + n] = sum + (bias ? bias[n] : 0.0f);
291 }
292 }
293}
#define CK_FP16_TO_FP32(x)
#define QK4_1
#define C(color)
Definition show_config.c:39
uint8_t qs[32/2]

References C, CK_FP16_TO_FP32, block_q4_1::d, block_q4_1::m, QK4_1, and block_q4_1::qs.

Referenced by ck_gemm_nt_quant().

◆ gemm_q4_1()

void gemm_q4_1 ( float *  Y,
const void *  W,
const float *  X,
int  M,
int  N,
int  K 
)

Matrix-matrix multiply with Q4_1 weights.

Definition at line 158 of file gemm_kernels_q4_1.c.

162{
163 for (int n = 0; n < N; n++) {
164 gemv_q4_1(&Y[n * M], W, &X[n * K], M, K);
165 }
166}

References gemv_q4_1().

◆ gemm_q4_1_backward()

void gemm_q4_1_backward ( float *  dX,
const void *  W,
const float *  dY,
int  M,
int  N,
int  K 
)

Batched backward pass.

Definition at line 231 of file gemm_kernels_q4_1.c.

235{
236 for (int n = 0; n < N; n++) {
237 gemv_q4_1_backward(&dX[n * K], W, &dY[n * M], M, K);
238 }
239}
void gemv_q4_1_backward(float *dX, const void *W, const float *dY, int M, int K)
Auto-dispatch backward.

References gemv_q4_1_backward().

◆ gemv_q4_1()

void gemv_q4_1 ( float *  y,
const void *  W,
const float *  x,
int  M,
int  K 
)

Auto-dispatch GEMV.

Definition at line 139 of file gemm_kernels_q4_1.c.

143{
144#ifdef __AVX512F__
145 gemv_q4_1_avx512(y, W, x, M, K);
146#else
147 gemv_q4_1_ref(y, W, x, M, K);
148#endif
149}
void gemv_q4_1_ref(float *y, const void *W, const float *x, int M, int K)
Matrix-vector multiply with Q4_1 weights (scalar reference)

References gemv_q4_1_ref().

Referenced by dot_q4_1(), and gemm_q4_1().

◆ gemv_q4_1_backward()

void gemv_q4_1_backward ( float *  dX,
const void *  W,
const float *  dY,
int  M,
int  K 
)

Auto-dispatch backward.

Definition at line 220 of file gemm_kernels_q4_1.c.

224{
225 gemv_q4_1_backward_ref(dX, W, dY, M, K);
226}
void gemv_q4_1_backward_ref(float *dX, const void *W, const float *dY, int M, int K)
Backward pass: compute input gradient.

References gemv_q4_1_backward_ref().

Referenced by gemm_q4_1_backward().

◆ gemv_q4_1_backward_ref()

void gemv_q4_1_backward_ref ( float *  dX,
const void *  W,
const float *  dY,
int  M,
int  K 
)

Backward pass: compute input gradient.

Parameters
dXOutput gradient w.r.t. input [K]
WWeight matrix in Q4_1 format [M x K]
dYGradient w.r.t. output [M]
MNumber of output rows
KNumber of columns (input dimension)

Definition at line 181 of file gemm_kernels_q4_1.c.

185{
186 const block_q4_1 *blocks = (const block_q4_1 *)W;
187 const int blocks_per_row = K / QK4_1;
188
189 /* Zero output gradient */
190 memset(dX, 0, K * sizeof(float));
191
192 /* Accumulate: dX += W^T @ dY */
193 for (int row = 0; row < M; row++) {
194 const float dy = dY[row];
195
196 for (int b = 0; b < blocks_per_row; b++) {
197 const block_q4_1 *block = &blocks[row * blocks_per_row + b];
198 const float d = CK_FP16_TO_FP32(block->d);
199 const float m = CK_FP16_TO_FP32(block->m);
200 float *dxp = &dX[b * QK4_1];
201
202 for (int i = 0; i < QK4_1 / 2; i++) {
203 const uint8_t packed = block->qs[i];
204 const int q0 = (packed & 0x0F);
205 const int q1 = (packed >> 4);
206
207 const float w0 = d * (float)q0 + m;
208 const float w1 = d * (float)q1 + m;
209
210 dxp[2*i + 0] += w0 * dy;
211 dxp[2*i + 1] += w1 * dy;
212 }
213 }
214 }
215}

References CK_FP16_TO_FP32, block_q4_1::d, block_q4_1::m, QK4_1, and block_q4_1::qs.

Referenced by gemv_q4_1_backward().

◆ gemv_q4_1_ref()

void gemv_q4_1_ref ( float *  y,
const void *  W,
const float *  x,
int  M,
int  K 
)

Matrix-vector multiply with Q4_1 weights (scalar reference)

Parameters
yOutput vector [M]
WWeight matrix in Q4_1 format [M x K]
xInput vector [K]
MNumber of output rows
KNumber of columns (must be multiple of 32)

Definition at line 50 of file gemm_kernels_q4_1.c.

54{
55 const block_q4_1 *blocks = (const block_q4_1 *)W;
56 const int blocks_per_row = K / QK4_1;
57
58 for (int row = 0; row < M; row++) {
59 float sum = 0.0f;
60
61 for (int b = 0; b < blocks_per_row; b++) {
62 const block_q4_1 *block = &blocks[row * blocks_per_row + b];
63 const float d = CK_FP16_TO_FP32(block->d);
64 const float m = CK_FP16_TO_FP32(block->m);
65 const float *xp = &x[b * QK4_1];
66
67 for (int i = 0; i < QK4_1 / 2; i++) {
68 const uint8_t packed = block->qs[i];
69 const int q0 = (packed & 0x0F);
70 const int q1 = (packed >> 4);
71
72 /* Dequantize: w = d * q + m */
73 const float w0 = d * (float)q0 + m;
74 const float w1 = d * (float)q1 + m;
75
76 sum += w0 * xp[2*i + 0];
77 sum += w1 * xp[2*i + 1];
78 }
79 }
80
81 y[row] = sum;
82 }
83}

References CK_FP16_TO_FP32, block_q4_1::d, block_q4_1::m, QK4_1, and block_q4_1::qs.

Referenced by gemv_q4_1().