Skip to content

Fused int4 and int8 dequantize-matmul in C

ModuleL9.5 · side · C · Pass 6 · 4 to 5 h
You buildc/src/kernels/qmatmul.c: tl_matmul_q4_f32 (W4A32: signed 4-bit weights, one f16 scale per group) and tl_matmul_q8_f32 (W8A32: int8 weights, one f32 scale per output row), each computing y=xW⊤y = x W^\top without ever writing WW out in float32
Contractcourse/contracts/c/include/tinyllm/qmatmul.h · byte layout: formats/safetensors.md (Int4 weights) · rules: c/ABI.md
Testscourse/tests/L9.5/test_qmatmul.c (standalone C tests under ASan and UBSan, against dequantize-then-tl_matmul_f32) and course/tests/L9.5/bench/q4_gemv_2048.c
Needsrt.02 the loader · rt.03 the pool · M09.7 tl_f16_to_f32 · L9.1 your float32 matmul, the baseline · L8.5 your Python quantizer (or --ref-deps). Reading: M09.3 (error bounds)
Used byNone: this optional C exercise is tested as a standalone binary; the Rust engine implements quantized inference independently
MilestoneMS-L9
Optional depthFrantar et al., “GPTQ” (2022); Lin et al., “AWQ” (2023); Dettmers et al., “LLM.int8()” (2022); the llama.cpp Q4_0 block format (ggml-quants.c)
  • Decode is memory-bound, so reading 4.5 bits per weight instead of 32 makes it faster even though every weight must be unpacked first: the reference int4 GEMV is about 3.9 times faster than your float32 kernel at N=K=2048N = K = 2048 (ol bench L9.5).
  • The layout is a contract: byte bb of row nn holds column 2b2b in the low nibble and 2b+12b + 1 in the high nibble, each a signed 4-bit two’s complement value in [−8,7][-8, 7] (every_nibble_code, hand_example).
  • One scale per group of columns, an IEEE half-precision bit pattern, multiplies the group’s partial sum once: y=∑gsg∑k∈gxkqky = \sum_g s_g \sum_{k \in g} x_k q_k (matches_dequantize_then_matmul, f16_scales_decode_exactly).
  • Fused means the float32 weight never exists: each chunk of 64 weights is unpacked into a small buffer and consumed immediately.
  • Speed comes from independent accumulators the compiler can keep in registers. Eight interleaved lanes per group, combined in a fixed tree, keep the kernel fast and every row’s bits independent of the batch (batch_invariant_rows).
Terminal window
ol start L9.5 # stubs c/src/kernels/qmatmul.c into your repo
ol tests L9.5 # read the test catalog first
ol check L9.5 # exit code is the verdict
ol check L9.5 --ref-deps # only if rt.02, rt.03, M09.7, L9.1, or L8.5 is not passing yet
ol bench L9.5 --assert # at least 2x your float32 kernel on one decode step (local only)
ol parity quant.int4 # your L8.5 encoder and this decoder against one golden
ol diff L9.5 # after passing: your code against the reference

L8.5 taught your Python stack to quantize a Llama: quantize_int4_group packs every projection into 4-bit codes with float16 scales, and export_q4 writes them as *.qweight and *.scales tensors. Python’s QuantLinear dequantizes to float32 and calls numpy, so it saves disk, not time. This optional standalone C exercise reads packed bytes directly and compares its output with dequantize-then-matmul. The Rust serving path is a separate implementation that consumes the same safetensors format without calling this C kernel.

SymbolMeaningType / shape
x∈RM×Kx \in \mathbb{R}^{M \times K}activations, MM rows (1 for decode)float[M][K]
W∈RN×KW \in \mathbb{R}^{N \times K}the Linear weight, [out, in] as storednever materialized
y=xW⊤y = x W^\topthe outputfloat[M][N]
GGgroup: columns that share one scaleint64_t
qnk∈{−8,…,7}q_{nk} \in \{-8, \dots, 7\}the 4-bit code of WnkW_{nk}4 bits
sn,gs_{n,g}the scale of row nn, group g=⌊k/G⌋g = \lfloor k / G \rfloorIEEE f16
Wnk=qnk sn,⌊k/G⌋W_{nk} = q_{nk}\, s_{n, \lfloor k/G \rfloor}the value the code stands for
ℓ\ella lane index, 0≤ℓ<80 \le \ell < 8

2.1 Why quantized decode is faster: bytes per weight

Section titled “2.1 Why quantized decode is faster: bytes per weight”

A decode step multiplies one activation row by every weight matrix: 2NK2NK FLOPs for NKNK weights, one multiply-add per weight loaded. The arithmetic intensity is about 0.5 FLOP per byte in float32, far below what a core can compute per byte of bandwidth (L9.1 2.1), so the step takes as long as it takes to read the weights. int4 with groups of 32 stores 4+16/32=4.54 + 16/32 = 4.5 bits per weight, a factor of 7.1 fewer bytes; int8 per channel stores about 8. The kernel can spend several instructions per weight on unpacking and still win, as long as it keeps up with memory.

formats/safetensors.md (shared with L8.5, which writes it, and the Rust engine, which loads the same file format independently) fixes:

  • qweight: uint8 [N, K/2]. Byte bb of row nn holds column 2b2b in its low nibble (bits 0 to 3) and column 2b+12b + 1 in its high nibble (bits 4 to 7).
  • Each nibble is signed 4-bit two’s complement: the unsigned value u∈[0,15]u \in [0, 15] means uu when u<8u < 8 and u−16u - 16 otherwise, so 0x8 is −8-8 and 0xF is −1-1. In C: u - ((u & 8) << 1).
  • scales: f16 [N, K/G], the raw bit patterns of IEEE half precision (1 sign, 5 exponent, 10 mantissa bits). Decoded by your tl_f16_to_f32 from M09.7, not by shifting into the top of a float32, which is what bfloat16 would be.

The kernel validates the shape before touching anything: KK even (two codes per byte), GG even and positive (a group is whole bytes), KK a multiple of GG (whole groups); otherwise TL_ESHAPE.

Dequantize-then-multiply would write WW in float32 (as many bytes as the model you were trying not to read) and read it back. The fused kernel instead factors each group’s scale out of its sum:

ymn=∑kxmk qnk sn,⌊k/G⌋=∑gsn,g∑k∈gxmk qnk⏟partialg.y_{mn} = \sum_{k} x_{mk}\, q_{nk}\, s_{n, \lfloor k/G \rfloor} = \sum_{g} s_{n,g} \underbrace{\sum_{k \in g} x_{mk}\, q_{nk}}_{\text{partial}_g} .

One multiply by the scale per group instead of one per weight, and the codes are turned into floats only in a small buffer: 64 at a time (32 bytes), consumed by the next loop while they are still in L1. int8 is the same with one group per row: ymn=sn∑kxmkqnky_{mn} = s_n \sum_k x_{mk} q_{nk}.

A single running sum is a chain: each addition waits for the previous one, about 4 cycles on a modern core, so one chain caps the kernel at a quarter of an addition per cycle whatever the vector width. The kernel therefore keeps 8 independent lanes per group: lane ℓ\ell accumulates the products of positions ℓ,ℓ+8,ℓ+16,…\ell, \ell + 8, \ell + 16, \dots of the group, in increasing order, and the compiler maps the 8 lanes onto vector registers (two 4-wide NEON or one 8-wide AVX register). At the end of the group the lanes combine in a fixed tree,

partialg=((ℓ0+ℓ1)+(ℓ2+ℓ3))+((ℓ4+ℓ5)+(ℓ6+ℓ7)),\text{partial}_g = \big((\ell_0 + \ell_1) + (\ell_2 + \ell_3)\big) + \big((\ell_4 + \ell_5) + (\ell_6 + \ell_7)\big),

and the output accumulates sn,g partialgs_{n,g}\,\text{partial}_g in increasing gg. Every one of those steps depends only on (m,n)(m, n): not on MM, not on the thread, not on the row’s position. That keeps the batch invariance of c/ABI.md rule 10 (the order is fixed, though it is not a single left-to-right sum).

Two C details decide whether the lanes stay in registers. If the lane array is a parameter, the compiler must assume it could alias x or the code buffer and reloads it on every step; copying it into a local array and marking the pointers restrict lets the lanes live in registers. In the reference that one change took the kernel from 1.75 ms to 0.99 ms per step. And #pragma STDC FP_CONTRACT OFF keeps the compiler from fusing a multiply and an add on some paths but not others (L9.1 2.4).

Output rows nn split across rt.03 workers in ranges of 8; each ymny_{mn} is computed entirely by one worker with the arithmetic above, so threads never change the bits.

WW is 2×42 \times 4, G=2G = 2, x=[1,2,3,4]x = [1, 2, 3, 4], M=1M = 1:

Rowcodes qqscales ss (f16 bits)bytes
0[1,−2∣3,−8][1, -2 \mid 3, -8][0.5∣2][0.5 \mid 2] = 0x3800, 0x40000xE1, 0x83
1[7,0∣−1,4][7, 0 \mid -1, 4][1∣0.25][1 \mid 0.25] = 0x3C00, 0x34000x07, 0x4F

Packing row 0: $1 = $ 0x1 (low), $-2 = 16 - 2 = 14 = $ 0xE (high), so byte 0 is 0xE1; $3 = $ 0x3, $-8 = $ 0x8, byte 1 is 0x83. Row 1: 0x07 and ($-1 = $ 0xF low, 44 high) 0x4F.

Decoding byte 0xE1: low nibble 1<81 < 8 gives 11; high nibble 14≥814 \ge 8 gives 14−16=−214 - 16 = -2.

y0=0.5 (1⋅1+2⋅(−2))+2 (3⋅3+4⋅(−8))=0.5⋅(−3)+2⋅(−23)=−47.5y_0 = 0.5\,(1 \cdot 1 + 2 \cdot (-2)) + 2\,(3 \cdot 3 + 4 \cdot (-8)) = 0.5 \cdot (-3) + 2 \cdot (-23) = -47.5

y1=1 (1⋅7+2⋅0)+0.25 (3⋅(−1)+4⋅4)=7+0.25⋅13=10.25y_1 = 1\,(1 \cdot 7 + 2 \cdot 0) + 0.25\,(3 \cdot (-1) + 4 \cdot 4) = 7 + 0.25 \cdot 13 = 10.25

Every value is exact in float32. With the same codes as int8 and one scale per row (0.50.5 and 0.250.25): y=[0.5⋅(1−4+9−32),0.25⋅(7+0−3+16)]=[−13,5]y = [0.5 \cdot (1 - 4 + 9 - 32), 0.25 \cdot (7 + 0 - 3 + 16)] = [-13, 5]. These are the values checked by hand_example and q8_hand_example.

/* tinyllm/qmatmul.h: y [M, N] = x [M, K] @ W^T, W [N, K] quantized */
tl_status tl_matmul_q4_f32(const float *x, const uint8_t *wq, const uint16_t *scales_f16,
float *y, int64_t M, int64_t N, int64_t K, int64_t group, tl_pool *tp);
tl_status tl_matmul_q8_f32(const float *x, const int8_t *wq, const float *scales,
float *y, int64_t M, int64_t N, int64_t K, tl_pool *tp);
/* TL_ESHAPE: K odd, group odd or <= 0, K % group != 0. TL_EINVAL: negative dims,
NULL pointers. M or N == 0: no-op. K == 0: y = 0. */
TestKINDChecksWhy it matters downstream
hand_exampleunit, smokesection 3 exactlythe nibble order, the sign, the scale per group
q8_hand_exampleunitthe int8 version of section 3one scale per row
every_nibble_codeboundaryall 16 codes in both nibbles decode to 0..7,−8..−10..7, -8..-1an unsigned or swapped decode is caught exactly
f16_scales_decode_exactlyboundarythe f16 bits 0x0001 (2−242^{-24}) and 0x3555half precision, not bfloat16
matches_dequantize_then_matmuldifferentialgroups 2, 32, 64, and all of K=384K = 384 against dequantize + your tl_matmul_f32the fused kernel equals the obvious recipe
q8_matches_dequantize_then_matmuldifferentialK=200K = 200 (not a multiple of the chunk) against the same recipeint8 per channel
batch_invariant_rowsproperty12 rows alone vs in batches, bitwisedecode and batched decode agree
shape_and_argument_errorsboundaryTL_ESHAPE and TL_EINVAL cases leave yy untouched; empty dims; K=0K = 0errors before any read
pool_result_equals_serial_bitwiseproperty4 threads vs serial, N=37N = 37threads never change bits
hand_exampleunit, smokenumpy packs section 3’s bytes; the kernel returns [−47.5,10.25][-47.5, 10.25]the layout from Python

The bench (ol bench L9.5, local only) runs one decode step, M=1M = 1, N=K=2048N = K = 2048, group 32, and reports q4_speedup_vs_f32 (budget ≥2\ge 2; the reference reaches about 3.9) and the int8 ratio.

PitfallSymptomCaught by
reading nibbles as unsigned 0..150..15negative weights become large positive onesevery_nibble_code (mutant s01)
the high nibble as the even columncolumns swapped in pairs: plausible, wrongevery_nibble_code (mutant s02)
indexing scales by group only, not by rowevery row uses row 0’s scaleshand_example (mutant s03)
not resetting the partial sum per groupeach scale multiplies the sum of all earlier groups toohand_example (mutant s04)
decoding the f16 scale as bfloat16scales off by orders of magnitudef16_scales_decode_exactly (mutant s05)
unpacking at the chunk offset instead of the group’severy group reads the first group’s codesmatches_dequantize_then_matmul (mutant s06)
the int8 scale indexed by the activation rowwrong for M>1M > 1 onlyq8_matches_dequantize_then_matmul (mutant s07)
forgetting the int8 scaleoutputs 50 to 1000 times too largeq8_hand_example (mutant s08)
a special summation order for M=1M = 1decode and batched decode differ in the last bitbatch_invariant_rows (mutant s09)
accepting a group that does not divide KKreads past x and the weightsshape_and_argument_errors (mutant s10)
skipping the NULL check on the weightsa crash instead of TL_EINVALshape_and_argument_errors (mutant s11)
a pooled partition that drops a rowwrong only with threadspool_result_equals_serial_bitwise (mutant s12)
accumulating through a pointer parametercorrect but about 1.8 times slower: the lanes are reloaded every stepthe bench (ol bench L9.5, local only; no course test checks speed)
DirectionModuleHow it uses this
Backrt.02the error slot and allocator support
Backrt.03tl_parallel_for over output rows
BackM09.7tl_f16_to_f32 decodes every scale (and tl_f32_to_f16 encodes them in the tests)
BackL9.1the float32 baseline: dequantize, then tl_matmul_f32; the bench’s comparison
BackL8.5the quantizer and the byte layout this kernel decodes
Your pieceProduction equivalentWhat it addsWhere to look
int4 with f16 group scalesllama.cpp Q4_0, Q4_Kblocks of 32 with the scale stored inline next to the codes (one cache line), super-blocks with 6-bit sub-scalesggml/src/ggml-quants.c
W4A32W4A8 / W4A16 kernelsactivations quantized too, integer dot products (sdot, VNNI)llama.cpp ggml_vec_dot_q4_0_q8_0
absmax quantizationGPTQ, AWQchoose codes to minimize the layer’s output error, not each weight’sthe GPTQ and AWQ papers
CPU GEMVMarlin, ExLlamaV2GPU kernels that reach memory bandwidth at batch sizes up to 16IST-DASLab marlin