🚧 Week 2 Day 2: Keep W4 Packed
Complete Day 1: Cache and Measure first. Its
capacity-cache checkpoint gives you a matched cached-model baseline, bounded
dense K/V storage, and a decode roofline. Keep that baseline fixed while changing the
projection representation. The Day 2 starter
already supplies the packed-weight container and its
QuantizedWeights.from_mlx_layer loader, the extension declaration and
binding, fail-closed C++/Metal stubs, and cumulative model switches. Your work
is to:
- dequantize selected embedding rows without expanding the full table;
- define the lazy quantized-matmul primitive and its validation boundary;
- implement the readable Metal matrix control and the decode-shaped SIMD matvec; and
- wire packed projections and the tied output head into the live cached model.
Build the learner and reference extensions before running Day 2’s native tests. Begin with the Python wrapper gate, then exercise the GPU operator:
pdm run build-ext
pdm run build-ext-ref
pdm run test --week 2 --day 2 -- -k task_1
pdm run test --week 2 --day 2 -- -k gpu
Finish with the complete Day 2 gate and the quantized-matvec model
checkpoint. The result is complete only when the cached model dispatches
through your quantized path; packed storage and an isolated fast kernel are
intermediate steps.
📚 Readings
Debug Metal Without a CPU Twin
You do not need a C++ CPU twin. Bring up the operator with this three-level validation ladder:
- Write the equation in Python with
mlx.core. This is the semantic oracle. - Translate it into a deliberately simple Metal kernel, usually with one thread responsible for one output element.
- Optimize the validated Metal kernel with SIMD groups, vectorized loads, or SIMD-group matrix operations.
Compare each level with the one immediately above it. Full-model text output is too indirect to diagnose an optimized kernel.
Make Failures Small and Synchronous
Use deterministic fixtures whose expected values are easy to inspect: zeros, ones, ramps, identity-like weights, and a fixed random seed. Exercise a small aligned shape and then a tail shape. For example, test 8 and 10 rows for an 8-row tile, or sequence lengths 32 and 35 for a 32-token block.
MLX execution is lazy, so force evaluation directly after the operator under test. This turns a delayed compile or GPU execution failure into a failure at the responsible call site:
expected = python_reference(*inputs)
actual = metal_operator(*inputs)
mx.eval(expected, actual)
assert actual.shape == expected.shape
assert actual.dtype == mx.bfloat16
assert mx.allclose(actual, expected, rtol=2e-2, atol=2e-2).item()
When a check fails, inspect the wrapper boundary before the arithmetic. Assert the tensor rank, shape, dtype, and contiguity assumptions in Python or C++, and verify that the encoded buffer indices match the Metal function signature. Then classify the failure:
- a pipeline creation error usually means the kernel name, specialization, or Metal compilation is wrong;
- an execution or address error usually means a grid, bounds check, stride, or buffer binding is wrong;
- a finite but inaccurate result usually means the indexing, reduction, mask, dequantization, or accumulator update is wrong.
For a numerical mismatch, simplify the schedule temporarily. Assign one output to one thread, remove cooperative loads, and compare an intermediate such as a dequantized weight group, a partial dot product, or an online-softmax row. A small debug-only output buffer is usually more useful than printing from every GPU thread. Restore one optimization at a time, rerunning the aligned and tail-shape tests after each change.
Represent Weights With Fewer Bits
Quantization stores each floating-point weight as a value from a small integer codebook plus the parameters needed to reconstruct an approximation. This course uses weight-only 4-bit quantization:
- W4 means that each logical weight is represented by a 4-bit code.
- A16 means that activations and outputs remain 16-bit floating point.
- The resulting path is called W4A16. This course uses BF16 for its activations, scales, biases, and outputs.
The 16 possible codes approximate the original values. In return for some numerical precision, the smaller representation reduces memory traffic.
The kernel does not materialize a dense BF16 weight matrix. It unpacks each 4-bit code, reconstructs the weight in registers, and immediately multiplies it by the corresponding BF16 activation.
Group-Wise Affine Quantization
Rather than applying one scale to an entire weight matrix, divide each row into groups and quantize each group independently. A local scale and bias retain more information about that group’s weight distribution.
For a weight matrix W of shape (K, N), divide each row into groups of size
G. The Qwen3-4B MLX 4-bit checkpoint used in this course has a fixed group
size of 128:
Logical weight matrix W: K × N
Group size: G = 128
Number of groups per row = N / G
For each stored group of G consecutive values in a row:
1. Unpack each unsigned 4-bit code q in [0, 15]
2. Load the group's stored scale s and bias b
3. Reconstruct each value as q * s + b
Reconstruct a Stored Group
The checkpoint already contains packed codes and affine parameters. For an
unpacked unsigned code q, use the stored scale s and bias b directly:
Read the equation as: reconstructed weight equals the unsigned code q
multiplied by its stored signed scale s, plus its stored bias b.
The codes are unsigned, but the stored scale is signed. A positive scale maps
code 0 to the lower endpoint and code 15 toward the upper endpoint. A negative
scale reverses the mapping: code 0 is the upper endpoint and code 15 moves
toward the lower endpoint. Both orientations occur in the shipped Qwen3-4B MLX
checkpoint. Use scale and bias as stored instead of reconstructing them from
an assumed min/max orientation.
For example, these two stored parameter pairs reconstruct the same endpoint range in opposite code order:
positive orientation: scale = 0.0867, bias = -0.5 => q=0 is -0.5, q=15 is about 0.8
negative orientation: scale = -0.0867, bias = 0.8 => q=0 is 0.8, q=15 is about -0.5
All required quantized-matmul tests use group_size = 128 and BF16 scales,
biases, activations, and outputs. Normalize those tensors to BF16 in your
solution’s model loader so every later kernel receives one model dtype.
Packed Storage Layout
The 4-bit codes are packed for compact storage and efficient access:
Logical weight matrix: K × N
Dense BF16 storage: K × N bfloat16 (2 bytes each) = 2KN bytes
W4 code storage: K × N int4 (0.5 bytes each) = 0.5KN bytes
Packing: 8 × 4-bit values fit in one uint32 (32 bits)
Packed codes shape: K × (N / 8) uint32
Scales shape: K × (N / G) bfloat16
Biases shape: K × (N / G) bfloat16
Example packing for 8 consecutive 4-bit values [a, b, c, d, e, f, g, h]:
uint32_value = (h << 28) | (g << 24) | (f << 20) | (e << 16) |
(d << 12) | (c << 8) | (b << 4) | a
Unpacking:
a = (uint32_value >> 0) & 0xF
b = (uint32_value >> 4) & 0xF
c = (uint32_value >> 8) & 0xF
...
h = (uint32_value >> 28) & 0xF
Revisit the Decode Roofline
The packed codes are not the entire W4 representation. Each group of 128 weights also stores one BF16 scale and one BF16 bias:
bytes per W4 weight = 0.5 + (2 + 2) / 128 = 0.53125 bytes
streamed W4 bytes = 4,022,272,000 × 0.53125 = 2.137 GB per token
arithmetic intensity = 8.045 GFLOPs / 2.137 GB = 3.765 FLOPs/byte
Now W4 can be added to the dense comparison:
| Weight format | Value bits | Metadata per 128 weights | Effective bytes per weight | Streamed weight bytes per token | Weight arithmetic intensity |
|---|---|---|---|---|---|
| FP16 | 16 | None | 2 | 8.045 GB | 1.0 FLOP/byte |
| BF16 | 16 | None | 2 | 8.045 GB | 1.0 FLOP/byte |
| W4 | 4 | One BF16 scale and one BF16 bias | 0.53125 | 2.137 GB | 3.765 FLOPs/byte |
This representation reduces projection weight traffic by 3.765×. Treat that ratio as a weight-traffic advantage for one-token decode, not as an end-to-end speedup promise.
Theoretical Decode Roofline Across Apple Silicon
Apple publishes unified-memory bandwidth but not a directly comparable BF16 GPU TFLOPS figure. A bandwidth roofline can therefore be calculated without assuming a compute ceiling:
ideal tokens/s = advertised memory bandwidth / streamed weight bytes per token
The table uses the highest-bandwidth configuration of each named chip listed in the cited Apple product specifications. GB is decimal, matching Apple’s specifications. The results are theoretical ceilings, not benchmark measurements.
| Chip | Bandwidth | FP16/BF16 roofline | W4 roofline |
|---|---|---|---|
| M1 Pro | 200 GB/s | 24.9 tok/s | 93.6 tok/s |
| M1 Max | 400 GB/s | 49.7 tok/s | 187.2 tok/s |
| M1 Ultra | 800 GB/s | 99.4 tok/s | 374.4 tok/s |
| M2 Pro | 200 GB/s | 24.9 tok/s | 93.6 tok/s |
| M2 Max | 400 GB/s | 49.7 tok/s | 187.2 tok/s |
| M2 Ultra | 800 GB/s | 99.4 tok/s | 374.4 tok/s |
| M3 Pro | 150 GB/s | 18.6 tok/s | 70.2 tok/s |
| M3 Max | 400 GB/s | 49.7 tok/s | 187.2 tok/s |
| M3 Ultra | 819 GB/s | 101.8 tok/s | 383.3 tok/s |
| M4 Pro | 273 GB/s | 33.9 tok/s | 127.8 tok/s |
| M4 Max | 546 GB/s | 67.9 tok/s | 255.5 tok/s |
The advertised bandwidths come from Apple’s specifications for M1 Pro and Max, M1 Ultra, M2 Pro and Max, M2 Ultra, M3 Pro and Max, M3 Ultra in the 2025 Mac Studio, and M4 Pro and Max. The 2025 Mac Studio specifications list M3 Ultra alongside M4 Max, not an M4 Ultra; this dated comparison stops at M4 Pro/Max and makes no claim about subsequent products.
These values assume peak advertised bandwidth, one read of every projection weight, and no other traffic or work. A complete model also reads activations and KV, launches other operators, and cannot sustain peak bandwidth continuously, so actual throughput is lower. The performance appendix records measured results separately from this theoretical exercise.
This roofline describes one-token decode, where M = 1 and each streamed
weight serves one activation row. Prefill reuses each weight tile across many
rows and raises arithmetic intensity, so it needs a matrix schedule. The decode
bandwidth ratio does not predict prefill performance.
Quantized Matrix Multiplication
Mathematical Formulation
For the matrix product C = A B^T, multiply A by the transpose of B.
The arrays are:
A: shape(M, N), bfloat16 (activations)B: shape(K, N), quantized to int4 (weights)C: shape(M, K), same 16-bit dtype asA(output)
Each element C[i, k] is computed as:
For output row i and column k, multiply A[i, j] by B[k, j] for
each input position j from zero through N - 1, then add the products.
With quantization, B[k, j] is represented as:
Reconstruct B[k, j] by multiplying its unsigned four-bit code by the
stored scale for row k and group g, then adding that group’s stored bias.
Here, g = floor(j / G) is the group index.
Substituting:
For each group g and position j' within it, multiply activation
A[i, g*G + j'] by the reconstructed weight at row k and that position;
sum those products over all N/G groups and all G positions per group.
Rearranging:
Equivalently, within each group, multiply the stored scale by the sum of
activation-times-code products, and add the stored bias multiplied by the sum
of those same activations. Add the group results to obtain C[i, k].
The scale and bias are constant within a group, so the computation can reuse them across all values in that group.
Computation Flow
Input:
A: M × N (bfloat16 activations)
B_quantized: K × (N/8) (uint32, packed weights)
scales: K × (N/G) (bfloat16)
biases: K × (N/G) (bfloat16)
Output:
C: M × K (bfloat16)
For each output element C[i, k]:
sum = 0 # float accumulator
for each group g in 0..(N/G - 1):
scale = scales[k, g]
bias = biases[k, g]
# Process G values in the group (G/8 uint32 packs)
for each pack p in 0..(G/8 - 1):
packed_value = B_quantized[k, g*(G/8) + p]
# Unpack 8 × 4-bit values
for bit_offset in [0, 4, 8, 12, 16, 20, 24, 28]:
quantized = (packed_value >> bit_offset) & 0xF
b_value = quantized * scale + bias
a_value = A[i, g*G + p*8 + bit_offset/4]
sum = sum + a_value * b_value
C[i, k] = bfloat16(sum)
Task 1: Implement Quantized Linear and Embedding
src/tiny_llm/quantize.py
src/tiny_llm/embedding.py
Inspect and reuse the starter’s QuantizedWeights.from_mlx_layer packed-weight
plumbing. Modify these learner-owned functions:
dequantize_weightsandquantized_linearinsrc/tiny_llm/quantize.py;QuantizedEmbedding.__call__andQuantizedEmbedding.as_linearinsrc/tiny_llm/embedding.py.
QuantizedWeights holds a quantized matrix and its dequantization parameters:
| Field | Shape | Description |
|---|---|---|
weight | (K, N/8) uint32 | Packed quantized weights. Each uint32 stores eight consecutive 4-bit values. |
scales | (K, N/G) bfloat16 | Stored signed per-group scale factors for dequantization. The sign determines which endpoint maps to the low codes. |
biases | (K, N/G) bfloat16 | Stored per-group offsets. Code 0 reconstructs to this value. |
group_size | int | Number of consecutive values that share the same scale/bias. For the Qwen3 MLX 4-bit weights used here, this is 128. |
bits | int | Quantization bit width (typically 4, meaning values are in range [0, 15]) |
The supplied from_mlx_layer method extracts these fields from an MLX
quantized layer during model loading. Reuse it rather than introducing a second
loader.
Then implement quantized_linear as a wrapper around quantized_matmul, using
the same input convention as the standard linear function. The next task
makes quantized_matmul runnable.
Keep the token embedding table quantized too. Add a QuantizedEmbedding
wrapper with two call patterns:
embedding(input_ids)performs a row lookup. Gather the matching packed weights, scales, and biases. Unpack eachuint32with shifts and masks, repeat each group’s scale and bias across its 128 values, and computeq * scale + biaswith basicmlx.corearray operations. Do not callmx.dequantize. Put this unpacking logic indequantize_weights(...)so the embedding and its direct tests share one explicit implementation.embedding.as_linear(h)is the tied output projection. Implement this withquantized_linear(h, embedding_weight)so it uses your quantized matmul path instead of materializing the fullvocab_size x hidden_sizetable. This path starts working once the quantized matmul kernel is implemented in the next tasks.
Task 2: Define the Quantized Matmul Primitive
src/extensions/src/tiny_llm_ext.h
src/extensions/bindings.cpp
src/extensions/src/quantized_matmul.cpp
src/extensions/CMakeLists.txt
The declaration, fail-closed source stub, binding, and build registration are
already present. Keep the C++ declarations and definitions in the
tiny_llm_ext namespace, and modify these exact functions:
tiny_llm_ext.h— Read the Week 2 Day 2quantized_matmul(...)declaration andQuantizedMatmulprimitive interface; keep its signature in sync with the binding.bindings.cpp— Verify the existingm.def("quantized_matmul", ...)entry; do not create a second binding.quantized_matmul.cpp— Replace the body oftiny_llm_ext::quantized_matmul(...)to validate inputs, determine the output shape, return a lazymx::array, and reject CPU evaluation explicitly inQuantizedMatmul::eval_cpu(...).CMakeLists.txt— Verify the existingquantized_matmul.cppsource registration; do not add a duplicate.
The extension API lets an mx.array graph node schedule the Metal loop from
the next task. MLX owns the array lifetime and command encoder; your primitive
supplies the quantized multiplication.
Build now to catch declaration, binding, and registration mismatches. The focused test checks the Task 1 Python wrappers. The primitive becomes runnable after you implement its Metal schedules in Task 3:
pdm run build-ext
pdm run test --week 2 --day 2 -- -k task_1
Task 3: Implement Metal Matrix Products
Before writing the Metal kernels, connect their work to Metal’s four nested execution scopes:
- Lane (thread). The smallest unit. Each lane executes the same
instruction stream with its own register file. Lanes within a SIMD group
can share data through
simd_operations. - SIMD group (warp/subgroup). A fixed-size set of lanes (32 on Apple
GPUs) that execute in lockstep.
simd_sum,simd_shuffle, andsimdgroup_matrixoperations work within this scope. A SIMD group cannot directly share registers with another SIMD group in the same threadgroup. - Threadgroup (block). A collection of SIMD groups scheduled together on
one GPU core. Threadgroups share threadgroup memory (explicitly allocated
with
threadgroupaddress space and synchronized withthreadgroup_barrier). The grid is a 1D/2D/3D array of threadgroups. - Grid. The total work dispatched.
dispatchThreadgroupslaunches a grid of threadgroups; the GPU schedules them across available cores. Increasing the grid’s threadgroup count can expose more independent work, but a finer partition can also duplicate reads or require partial-result merging.
Keep two launch knobs separate. Adding SIMD groups within one threadgroup adds threads and can raise register demand; it increases threadgroup-memory use only when the schedule allocates shared storage per group or tile. Either resource can reduce the number of resident threadgroups. Adding threadgroups to the grid changes how output or reduction work is partitioned. Measure both choices; neither guarantees higher throughput.
Begin with the required two-SIMD-group matvec schedule for Qwen. Then benchmark two, four, eight, and sixteen groups per threadgroup as described below. Change the grid partition separately so each measurement isolates one launch knob.
src/extensions/src/quantized_matmul.metal
src/extensions/src/quantized_matmul.cpp
Modify these exact starter functions:
QuantizedMatmul::eval_gpuinquantized_matmul.cpp;quantized_matmul_vanilla_w4a16_g128andquantized_matvec_x4_fast_w4a16_g128inquantized_matmul.metal;quantized_matmul_vanillaandquantized_matvec_custominsrc/tiny_llm/quantize.pyfor the explicit comparison paths.
Write both Metal kernels, then connect eval_gpu to them. On GPU, the Python
quantized_matmul wrapper always dispatches your primitive. The required path
never routes through mx.quantized_matmul.
Work in two measured stages. Both implement the same math, with schedules for different shapes:
- Vanilla matmul: one Metal thread computes one output element. This is the direct GPU translation of the computation flow above and an inspectable bring-up control.
- SIMD matvec: for decode, SIMD lanes cooperate on the reduction for one activation row and calculate several output columns together.
Here, M is the number of activation rows after flattening every leading
dimension. Day 2 uses this explicit dispatch:
| Activation rows | Kernel | Role at this checkpoint |
|---|---|---|
M <= 8 | SIMD matvec | Optimized path for decode and other very small matrix inputs. |
M > 8 | Vanilla matmul | Correctness-first prefill path; Day 3 replaces it with a cooperative tiled kernel. |
The cutoff does not extend the SIMD kernel to larger M. These are separate
schedules: Day 2 optimizes vector-shaped decode and leaves matrix-shaped
prefill visible for the later benchmark to select.
Keep the vanilla function callable as quantized_matmul_vanilla so every
optimization can be compared directly with its readable control.
Stage 1: Vanilla Matmul
Start with a two-dimensional grid over output row i and output column k.
Each thread walks all N input values, unpacks eight int4 weights from each
uint32, applies the group scale and bias, and accumulates one C[i, k] in
float32. This kernel repeats activation loads and does not share work, but its
control flow mirrors the equation and makes it a useful debugging control. The
Python mlx.core equation remains the correctness oracle for both Metal
schedules.
Keep the vanilla kernel for matrix-shaped prefill in this chapter; Day 3 revisits that workload with cooperative tiling.
Stage 2: SIMD Matvec
Decode normally has M = 1, so an 8×8 matrix tile would leave most rows empty.
Instead, let one SIMD group reduce the input dimension and use simd_sum to
combine lane-local partial sums. Begin with two output columns per group as an
inspectable schedule. For the Qwen3-4B checkpoint, evaluate a four-column path
in which each lane loads two adjacent packed words, or 16 activations, and
reuses them across the four outputs.
The optimized path also uses this affine identity to avoid applying the bias separately to every unpacked value:
The left side sums activation a_j times its reconstructed weight
s*q_j + b. The right side computes the same sum by multiplying s by the
activation-times-code sum, then adding b times the activation sum.
The kernel also scales the activations once and reads four packed int4 values through a 16-bit mask, avoiding a shift for every weight and output row. This adds live accumulators, so test it as a complete schedule rather than assuming fewer integer instructions must be faster.
Tune the SIMD Schedule
Treat output width, threadgroup size, and shared-memory reuse as benchmark variables. Begin with this Qwen-focused configuration:
- flatten all leading activation dimensions into
M, - use the custom matvec when
M <= 8and the vanilla matmul whenM > 8, - compute four output columns per SIMD group and load two adjacent packed words per lane,
- launch two SIMD groups, or eight output rows, per threadgroup.
These thresholds are measured starting points rather than mathematical requirements. Keep them visible in the dispatcher and vary one choice at a time. Compare two, four, and eight output columns per SIMD group. More columns increase activation reuse, but also extend accumulator lifetimes and raise register pressure. Compare two, four, eight, and sixteen SIMD groups per threadgroup. More groups expose additional outputs, but may duplicate activation reads and reduce residency.
Evaluate the affine rearrangement as part of the complete schedule. A lower instruction count helps only if the longer-lived activation sum and output accumulators do not reduce occupancy. Select the schedule with a synchronized whole-model decode benchmark, not with an instruction-count estimate.
Define a row-contiguous Python-to-extension contract for scales, biases,
activations, and packed weights. Call mx.contiguous once at that boundary and
validate the layout in the C++ primitive before encoding the kernel. Metal
receives raw buffers rather than implicit array strides, so layout is a
correctness condition as well as a performance condition.
Read activations directly in the kernel. The one-row activation is small and cache-friendly, while staging it in threadgroup memory adds a barrier to every projection. If you test shared staging as an ablation, keep it only when the whole-model result shows that reuse outweighs synchronization.
Kernel Requirements
Implement both required layouts in quantized_matmul.metal:
- First, implement the vanilla one-thread-per-output matrix grid.
- For
M <= 8, assign one SIMD group to an output tile. Cooperatively reduce the input dimension and compute several output columns per group. - For
M > 8, dispatch the vanilla matrix grid. Do not loop over rows with the SIMD matvec schedule; Day 3 introduces the tiled prefill schedule. - The required kernel supports
bfloat16_tinputs and outputs. The Week 2 checkpoint does not add a second model-storage dtype. - Apply the group-wise dequantization loop defined earlier in this chapter:
- Iterate over groups of 128 values.
- Unpack int4 values from each
uint32. - Dequantize each value with
q * scale + bias. - Accumulate products in
float, then cast the result to the kernel dtype.
- Add boundary checks (
i < M,k < K) before writing output.
The custom kernel only needs to support bits = 4 and group_size = 128. Use
the group size to compute groups_per_row and the packed-weight offsets.
Instantiate the required Metal kernel for bfloat16_t and select it in
eval_gpu. If you retain an optional half specialization, keep it out of the
model dispatch in your solution.
GPU Dispatch
Complete eval_gpu in quantized_matmul.cpp with axpby’s GPU dispatch
pattern:
- Get the Metal device and command encoder from the stream.
- Load the quantized matmul kernel matching the output dtype from the Metal library.
- Bind the input and output buffers and the dimension constants (
M,N,K). The buffer order must match the kernel signature. - Select the matrix-vector layout for
M <= 8; otherwise select the vanilla matrix layout. Keep both paths explicit for direct comparisons. Calculate a SIMD-aligned thread-group configuration and tile output columns so packed input values and activations can be reused. Use the four-column, two-packed-word kernel with two SIMD groups. - Dispatch with
dispatchThreadgroups.
Run the focused GPU gate:
pdm run build-ext
pdm run build-ext-ref
pdm run test --week 2 --day 2 -- -k gpu
The direct tests cover matvec at M = 1 and M = 8, the vanilla matmul at
M = 128, and compare them with an MLX oracle. The oracle checks the result;
it is not the implementation under test.
Task 4: Integrate Before Continuing
src/tiny_llm/qwen3_week2.py
Modify Qwen3ModelWeek2.__init__, Qwen3MultiHeadAttention.__call__,
Qwen3MLP.__call__, and Qwen3ModelWeek2.__call__. These are the points that
load quantized weights, replace dense projections, and keep only the requested
logits row.
Now integrate quantized matrix multiplication into the Week 2 Qwen3 model so its linear layers remain quantized throughout inference.
Change the weight type from mx.array to QuantizedWeights for every
attention projection (wq, wk, wv, and wo) and MLP projection (w_gate,
w_up, and w_down). Replace linear(x, w) with quantized_linear(x, w). In
the Week 2 model loader, use QuantizedWeights.from_mlx_layer(...) instead of
materializing a 16-bit matrix. Keep the Week 1 model’s boundary intact; its
layers still expect plain mx.array weights.
For embeddings, wire the QuantizedEmbedding from Task 1 into the loader. Load
embed_tokens with QuantizedWeights.from_mlx_layer(...) and pass it to
QuantizedEmbedding. If the model has a separate lm_head, keep that head as
QuantizedWeights too and apply it with quantized_linear; lm_head is a
projection rather than an embedding lookup.
Normalize each loaded layer’s scales and biases to BF16. Require scales,
biases, and activations to match and return BF16. If the output is nan or
otherwise invalid, check for a dtype mismatch first.
Preserve the quantized layer’s parameters as well. The model should pass
w.group_size and w.bits to the extension, which should validate the course
assumptions: group_size = 128 and bits = 4.
Run the complete gate, then the live model checkpoint:
pdm run test --week 2 --day 2
pdm run main --solution tiny_llm --loader week2 \
--week2-checkpoint quantized-matvec --model qwen3-0.6b --max-tokens 16
Measure that same learner solution with the cached 0.6B model. Keep the model, prompt/output lengths, and last-logit serving mode matched when changing the solution or checkpoint:
pdm run bench --solution tiny_llm --loader week2 \
--week2-checkpoint quantized-matvec --model qwen3-0.6b \
--num-seqs 1 --min-input-len 128 --max-input-len 128 \
--min-output-len 65 --max-output-len 65 --warmup 2 \
--prefill-logits last
Run the same command with --solution tiny_llm_ref to compare it with the
reference solution. For a Day 1 control, change only the checkpoint to
capacity-cache; do not compare different models or prefill-logit modes.
Keep the vanilla matrix product callable as an inspectable Metal control. The
Python mlx.core equation remains the correctness oracle, while decode
integrates only the SIMD matvec.
Verify Quantization in the Complete Model
Before moving on, confirm that model inference actually calls the quantized matvec kernel instead of merely registering and testing it in isolation.
The checkpoint is complete when the model’s projection dispatcher is wired to
your custom primitive. Decode-shaped work must route through
quantized_linear → quantized_matvec_custom → the extension primitive → the
Metal matvec. Matrix-shaped work must route through quantized_linear →
quantized_matmul → the extension primitive → its Metal matrix schedule. The
supplied tests compare public quantized-matvec checkpoint outputs and direct
operator results with MLX controls; they do not assert private routing. Use
the live model command to verify that those pieces compose.
Measure the cumulative model and the real projection shapes:
pdm run bench-week2-progression --offline --solution tiny_llm --suite week2 --repeats 2 \
--variant week2-capacity-cache --variant week2-quantized-matvec --variant mlx \
--model qwen3-0.6b --input-len 128 --output-len 129 --warmup 2 \
--prefill-logits last
pdm run profile-week2-kernels --solution tiny_llm --model qwen3-0.6b \
--case capacity-cache:decode:128 --case quantized-matvec:decode:128 \
--warmup 4 --iterations 12 \
--json-output week2-day2-attribution.json
Keep one cumulative model row and one representative real-shape projection
comparison. In the checked M4 Pro example, packed W4 reduced attributed
projection time by 69.0% and fixed-workload decode rose from 24.38 to 58.90
tokens/s. The re-profile then exposed normalization, position, and activation
at 33.5% of attributed time, pointing to later fusion work. These are
historical Qwen3-4B observations from one machine and two product samples,
not results to expect from the required 0.6B run or portable timing thresholds.
For an optional 4B comparison, first cache the model with
hf download Qwen/Qwen3-4B-MLX-4bit, then rerun every compared row with
--model qwen3-4b; the dense capacity-cache control needs the memory
recommended in the model table.
The complete historical campaign and attribution are in the
performance appendix.
If you need to continue without the custom Day 2 kernels, implement the same
quantized_linear interface with mx.quantized_matmul and leave the rest of
the course model unchanged. This is a local operator off-ramp, not
--solution mlx: the latter runs a separate complete model and does not
exercise your cache or model wiring.
Your feedback is greatly appreciated. Join our Discord community.
Found an issue? Open an issue or pull request at github.com/skyzh/tiny-llm.
tiny-llm-book © 2025 by Alex Chi Z is licensed under CC BY-NC-SA 4.0.