From 4bd61c6bdb69eb80a8ea800ba189f864f7c173a8 Mon Sep 17 00:00:00 2001 From: Sentinel Date: Wed, 23 Sep 2026 12:22:27 -0700 Subject: [PATCH 1/5] book: stage Week 2 Day 1 route AI-Assisted: GPT-6 Sol + Sentinel --- README.md | 20 +- book/src/SUMMARY.md | 18 +- book/src/appendix-performance.md | 8 +- book/src/glossary.md | 14 +- book/src/preface.md | 129 +++----- book/src/week1-07-sampling-prepare.md | 17 +- book/src/week2-01-kv-cache.md | 374 ++++++++++++++++++++--- book/src/week2-02-benchmark-profile.md | 9 +- book/src/week2-03-quantize-model.md | 9 +- book/src/week2-04-fused-model-kernels.md | 9 +- book/src/week2-05-simd-matrix-prefill.md | 9 +- book/src/week2-06-operator-lab.md | 9 +- book/src/week2-07-split-k-prefill.md | 9 +- book/src/week2-advanced-profiling.md | 9 +- book/src/week2-overview.md | 204 ++++--------- 15 files changed, 523 insertions(+), 324 deletions(-) diff --git a/README.md b/README.md index 04f51a3f..5d3e24a4 100644 --- a/README.md +++ b/README.md @@ -27,12 +27,10 @@ The course follows a four-week learning path: - **Week 1: From Matmul to Text.** Build a Qwen3 model directly from `mlx.core` array operations: attention, RoPE, GQA, RMSNorm, the MLP, sampling, and the autoregressive loop. -- **Week 2: A Step Closer to vLLM.** Add a KV cache, establish a - synchronized MLX baseline, and let matched benchmarks choose each - optimization. The causal path moves from quantized decode matvec to fused - model kernels and SIMD-matrix prefill; decode attention is an optional - workload-conditioned lab, and split-K stays only where a measured short - shape supports it. +- **Week 2: A Faster Single Request.** The current Day 1 route adds + `kv-cache` and request-bounded `capacity-cache`, then measures both against + the Week 1 full-prefix control. Later packed-W4, SIMD, fused-primitive, and + tiled-attention lessons are planned; their checkpoints are not Day 1 gates. - **Week 3: Build a Mini vLLM.** Introduce continuous batching and chunked admission, then make paged KV the canonical serving layout. Decode attention and FlashAttention learn to read pages directly so @@ -109,13 +107,7 @@ one explicit byte range through the existing loop. | 1.5 | Load the Model | ✅ | ✅ | ✅ | ✅ | | 1.6 | Generate Responses (aka Decoding) | ✅ | ✅ | ✅ | ✅ | | 1.7 | Sampling | ✅ | ✅ | ✅ | ✅ | -| 2.1 | KV Cache | ✅ | ✅ | ✅ | 🚧 | -| 2.2 | Benchmarking and Profiling | ✅ | ✅ | ✅ | 🚧 | -| 2.3 | Quantize the Model | ✅ | ✅ | ✅ | 🚧 | -| 2.4 | Fused Model Kernels | ✅ | ✅ | ✅ | 🚧 | -| 2.5 | SIMD-Matrix Prefill | ✅ | ✅ | ✅ | 🚧 | -| 2.6 (optional) | Workload-Conditioned Operator Lab | ✅ | ✅ | ✅ | 🚧 | -| 2.7 | Conditional Split-K and Final Decision | ✅ | ✅ | ✅ | 🚧 | +| 2.1 | Cache and Measure (`kv-cache`, `capacity-cache`) | 🚧 | 🚧 | ✅ | 🚧 | | 3.1 | Continuous Batching | ✅ | ✅ | ✅ | 🚧 | | 3.2 | Chunked Prefill | ✅ | ✅ | ✅ | 🚧 | | 3.3 | Paged KV Cache | ✅ | ✅ | ✅ | 🚧 | @@ -133,6 +125,8 @@ one explicit byte range through the existing loop. | 4.8 | Fork, Steer, and Select | ✅ | ✅ | ✅ | 🚧 | | 4.9 | Bound Tool Evidence | ✅ | ✅ | ✅ | 🚧 | +The older Week 2 chapter URLs remain available as [historical material](book/src/week2-02-benchmark-profile.md). Their former Day 2–7 tests and checkpoint commands are not part of the current Day 1 learner route. The Day 1 code and test checks are still being completed. + Other topics not covered include quantized or compressed KV caches, cross-request prefix caching, fine-tuning, and long-context techniques. diff --git a/book/src/SUMMARY.md b/book/src/SUMMARY.md index 68705cde..be1578ed 100644 --- a/book/src/SUMMARY.md +++ b/book/src/SUMMARY.md @@ -13,15 +13,15 @@ - [The Qwen3 Model](./week1-05-qwen3-model.md) - [Generating the Response](./week1-06-generate-response.md) - [Sampling and Preparing for Week 2](./week1-07-sampling-prepare.md) -- [🚧 Week 2: A Step Closer to vLLM](./week2-overview.md) - - [🚧 KV Cache](./week2-01-kv-cache.md) - - [🚧 Benchmarking and Profiling](./week2-02-benchmark-profile.md) - - [🚧 Optional: Metal Profiling](./week2-advanced-profiling.md) - - [🚧 Quantize the Model](./week2-03-quantize-model.md) - - [🚧 Fused Model Kernels](./week2-04-fused-model-kernels.md) - - [🚧 SIMD-Matrix Prefill](./week2-05-simd-matrix-prefill.md) - - [🚧 Day 6 (Optional): Workload-Conditioned Operator Lab](./week2-06-operator-lab.md) - - [🚧 Conditional Split-K and Final Decision](./week2-07-split-k-prefill.md) +- [🚧 Week 2: A Faster Single Request](./week2-overview.md) + - [🚧 Day 1: Cache and Measure](./week2-01-kv-cache.md) + - [Historical Week 2 lesson addresses](./week2-02-benchmark-profile.md) + - [Earlier quantization lesson](./week2-03-quantize-model.md) + - [Earlier fused-kernel lesson](./week2-04-fused-model-kernels.md) + - [Earlier SIMD-prefill lesson](./week2-05-simd-matrix-prefill.md) + - [Earlier bounded-decode lab](./week2-06-operator-lab.md) + - [Earlier Split-K lab](./week2-07-split-k-prefill.md) + - [Earlier optional macOS capture lab](./week2-advanced-profiling.md) - [🚧 Week 3: Build a Mini vLLM](./week3-overview.md) - [🚧 Continuous Batching](./week3-01-continuous-batching.md) - [🚧 Chunked Prefill](./week3-02-chunked-prefill.md) diff --git a/book/src/appendix-performance.md b/book/src/appendix-performance.md index 827b7770..2040bf14 100644 --- a/book/src/appendix-performance.md +++ b/book/src/appendix-performance.md @@ -1,8 +1,10 @@ # 🚧 Appendix: Performance Evidence Ledger -> **Status: Experimental, single-machine evidence.** See the -> [Week 2 verification matrix](./week2-overview.md#verification-status) before -> treating a correctness, integration, or performance result as broader proof. +> **Historical evidence from an earlier full Week 2 course state.** The +> [current Week 2 route](./week2-overview.md) ships Day 1 only. Later +> checkpoint labels, commands, and measured results below describe the older +> source tree; they are not runnable gates or performance results for this +> checkout. This appendix records the measurements that determined the course order. The numbers are not additive promises: after one bottleneck shrinks, every other diff --git a/book/src/glossary.md b/book/src/glossary.md index 35e3dd24..71822084 100644 --- a/book/src/glossary.md +++ b/book/src/glossary.md @@ -14,13 +14,13 @@ - [Qwen3 Transformer Block](./week1-05-qwen3-model.md) - [Week 1 Qwen3 Model](./week1-05-qwen3-model.md) - [dequantize_linear](./week1-05-qwen3-model.md) -- [KV Cache](./week2-01-kv-cache.md) -- [Benchmarking and Profiling](./week2-02-benchmark-profile.md) -- [Quantize the Model](./week2-03-quantize-model.md) -- [Fused Model Kernels](./week2-04-fused-model-kernels.md) -- [Fused Decode Attention](./week2-05-decode-attention.md) -- [SIMD-Matrix Prefill](./week2-06-simd-matrix-prefill.md) -- [Split-K Prefill](./week2-07-split-k-prefill.md) +- [KV Cache and Request-Bounded Capacity](./week2-01-kv-cache.md) +- [Benchmarking, Profiling, and Decode Roofline](./week2-01-kv-cache.md#benchmark-the-cached-model) +- [Historical: Packed W4 Quantization](./week2-03-quantize-model.md) +- [Historical: Fused Model Kernels](./week2-04-fused-model-kernels.md) +- [Historical: SIMD-Matrix Prefill](./week2-05-simd-matrix-prefill.md) +- [Historical: Bounded Decode Attention](./week2-06-operator-lab.md) +- [Historical: Split-K Prefill](./week2-07-split-k-prefill.md) - [Flash Attention](./week3-05-flash-attention.md) - [Paged Attention](./week3-04-paged-attention-part2.md) diff --git a/book/src/preface.md b/book/src/preface.md index 38739af1..74144971 100644 --- a/book/src/preface.md +++ b/book/src/preface.md @@ -31,8 +31,8 @@ This course is divided into four weeks. We will serve Qwen3 MLX models, optimize small coding agent. - Week 1: Serve Qwen3 using array and matrix operations written in Python. -- Week 2: Measure the cached model, implement the selected C++ and Metal - kernels, and re-profile after each change. +- Week 2: The current Day 1 route caches a request prefix, bounds its dense + storage, and measures matched requests. Later kernel lessons will follow. - Week 3: Add further optimizations and batch requests for high-throughput serving. - Week 4: Reuse the serving stack in a local coding agent with tools, sessions, and evaluation. @@ -42,70 +42,31 @@ The course supports two different goals: **implementing** the cumulative serving stack, or **studying and running** a later checkpoint without completing all earlier exercises. These are not the same path. - - -
- Tiny-LLM roadmap. The cumulative interface and state path runs from Week 1 through the seven Week 2 days and Week 3 into Week 4. Week 2 Days 3 through 7 show optional MLX operator off-ramps that preserve the course interfaces; they are different from the full-MLX model baseline. Week 4 keeps the course prerequisite of setup plus Weeks 1 through 3, while its deterministic scripted-model tests for Days 1 through 7 can run after setup. Day 8 joins the scripted sequence to the real-model path, and Day 9 continues the Week 4 sequence. -
- -

On a narrow screen, scroll the roadmap -horizontally; when the roadmap is focused, the left and right arrow keys move -through it without changing chapters. Its labels stay at their readable desktop -size.

- - - -Solid arrows in the diagram are **interface and state prerequisites**. They do -not mean that you must hand-write every earlier optimization. A dashed border -marks a custom operator that you may replace locally with its MLX equivalent -while keeping the surrounding course interface. The reference and full-MLX -lanes let you observe a completed system, but they do not fill in unfinished -functions in `src/tiny_llm`. +The current route runs from Week 1 to [Week 2 Day 1](./week2-01-kv-cache.md): +first `kv-cache`, then `capacity-cache`. The later Week 2 checkpoints are +planned, so a Day 1 checkout does not offer a completed Week 2 → Week 3 learner +path. The [earlier full-course roadmap diagram](./course-roadmap.svg) is +retained as historical context; its seven-day Week 2 order is not this +checkout's navigation. + +The reference and full-MLX models are useful controls, but they do not fill +unfinished functions in `src/tiny_llm`. A later custom kernel may have an MLX +operator substitution at the same interface; Day 1's cache state has no such +operator shortcut. | Your goal | Start here | What earlier implementation is required? | | --- | --- | --- | -| Build the whole serving system | Week 1, then follow the solid arrows | Each week uses interfaces and mechanisms established by the previous week. | -| Skip a Week 2 kernel optimization | Keep that day's course interface and wire the corresponding MLX operator at the seam | The earlier model, state, and interface work still needs to exist. This is a local code choice, not a CLI flag. | +| Build the currently shipped cache path | Week 1, then Week 2 Day 1 | Implement both cache checkpoints and their matched measurement. | +| Study a later Week 2 kernel | Read its historical page while waiting for the staged lesson | Its commands and checkpoint are not a Day 1 gate. | | Read or experiment with a later week | Open that chapter and use `tiny_llm_ref` | None in your learner tree. Run the supplied reference tests or reference loader. | | Compare with the production-library baseline | Use `--solution mlx` | None, but this runs the full MLX model and bypasses the course implementation. | | Run the Week 4 Days 1–7 deterministic tests before finishing the serving stack | After setup, run the supplied scripted-model tests | The tests do not need a working serving implementation. The course still assumes setup plus Weeks 1–3 before Week 4; follow the Week 4 days in order, and Day 8's real-model bridge needs the Week 3 model/tokenizer/KV-cache boundary. | The cumulative dependencies are deliberate: -- **Week 1 → Week 2:** Week 2 starts from the readable Qwen3 model and replaces - costs one measured mechanism at a time: first the generation algorithm and - KV cache, then quantized and fused kernels. Days 1–2 establish state and a - repeatable measurement; Days 3–5 follow the dominant cost, Day 6 is an - optional operator lab, and Day 7 closes with a conditional schedule decision. +- **Week 1 → Week 2:** Day 1 keeps the readable model, adds request-owned + dense K/V reuse and bounded storage, and compares the same request at both + cache checkpoints. Later packed-weight and custom-kernel work is planned. - **Week 2 → Week 3:** Week 3 selects MLX quantized projections, but it keeps course-owned normalization, activation, cache, attention, paging, batching, and scheduling. This is an explicit operator seam, not “use the MLX model for @@ -116,7 +77,7 @@ The cumulative dependencies are deliberate: harness to the real tokenizer and KV cache, so that checkpoint needs a working Week 3 path. -> **Is Week 2 required for Week 3? The Week 2 interfaces are; every Week 2 +> **Is Week 2 required for Week 3? Its interfaces are; every future > optimization is not.** The current Week 3 starter reuses the Week 2 model > shell, dense-cache contract, packed-weight plumbing, normalization, > activation, attention, and matrix-fragment interfaces. You may preserve @@ -127,27 +88,15 @@ The cumulative dependencies are deliberate: > the entire Week 2 implementation would require a supplied hybrid starting > checkpoint; that checkpoint does not exist today. -### Week 2 operator off-ramps - -Week 2 separates the mechanism you need later from the kernel you are invited -to optimize. If your goal is to continue into Week 3 rather than implement -every Metal kernel, you can make these explicit local substitutions: +### Later Week 2 operator off-ramps -| Week 2 day | Keep in the course stack | Optional MLX substitution | -| --- | --- | --- | -| Days 1–2 | Dense KV-cache state, the Week 2 model boundary, and the matched measurement method | None; these are state and methodology rather than replaceable operators. | -| Day 3 | Packed-weight containers, quantized embedding/model wiring, and the `quantized_linear` interface | Route projections through `mx.quantized_matmul` instead of the custom matrix-vector kernel. | -| Day 4 | The Week 2 norm, position, and activation call sites | Use the corresponding MLX RMSNorm/RoPE operators and an MLX SiLU-based SwiGLU composition instead of the custom fused kernels. | -| Day 5 | The quantized-projection interface and matrix-shaped dispatch boundary | Keep using the Day 3 MLX projection seam instead of implementing the SIMD-matrix schedule. | -| Day 6 (optional) | The dense-cache attention interface and its shape/mask adapter | Use `mx.fast.scaled_dot_product_attention` instead of the supplied bounded decode-attention branch. | -| Day 7 | The Day 5 unsplit projection fallback and measured dispatch boundary | Keep the unsplit path rather than implementing Split-K where your measurement does not support it. | - -Only the quantized-projection seam is already selected by canonical Week 3. -The Day 4 and optional Day 6 alternatives require you to wire the MLX call at the -existing course interface; there is no `--use-mlx-for-day` command. These -off-ramps let you study later mechanisms, but they do not complete the skipped -day's custom-kernel exercises, implementation-specific tests, or performance -claims. +Day 1 requires cache state and matched measurement; it has no replaceable +custom kernel. The earlier full-course book describes optional MLX substitutions +for later operators, but those are not shipped Day 1 checkpoints. Their +mechanisms and old addresses remain in the +[historical Week 2 pages](./week2-02-benchmark-profile.md). Selecting +`--solution mlx` runs a separate complete model, not a hybrid that completes +learner cache TODOs. To run a completed checkpoint without solving it first: @@ -156,16 +105,16 @@ To run a completed checkpoint without solving it first: pdm run build-ext-ref # Run one supplied reference test group. -pdm run test-refsol --week 3 --day 1 +pdm run test-refsol --week 2 --day 1 # Run a completed course model. -pdm run main --solution ref --loader week3 +pdm run main --solution tiny_llm_ref --loader week2 --week2-checkpoint kv-cache # Run the separate full-MLX baseline. pdm run main --solution mlx ``` -`--solution ref` runs the supplied implementation end to end. `--solution mlx` +`--solution tiny_llm_ref` runs the supplied implementation end to end. `--solution mlx` runs MLX end to end. Neither command composes “earlier weeks from the reference or MLX, this week's TODOs from my learner tree.” Per-operator substitution is a manual code edit that preserves the course interface; it is not a third @@ -183,10 +132,10 @@ default batch settings. | Unified memory | Week 1 | Week 2 | Week 3 | Week 4 | | --- | --- | --- | --- | --- | -| 8 GB | 0.6B / 0.6B | 0.6B / 1.7B[^week2-dense] | 0.6B / 1.7B | 0.6B / 1.7B | -| 16 GB | 0.6B / 1.7B | 4B / 8B[^week2-dense] | 4B / 8B | 4B / 8B | -| 18 GB | 0.6B / 1.7B | 4B / 8B[^week2-dense] | 4B / 8B | 4B / 8B | -| 24 GB | 0.6B / 1.7B | 4B / 8B[^week2-dense] | 4B / 8B | 4B / 8B | +| 8 GB | 0.6B / 0.6B | 0.6B / 0.6B[^week2-dense] | 0.6B / 1.7B | 0.6B / 1.7B | +| 16 GB | 0.6B / 1.7B | 0.6B / 1.7B[^week2-dense] | 4B / 8B | 4B / 8B | +| 18 GB | 0.6B / 1.7B | 0.6B / 1.7B[^week2-dense] | 4B / 8B | 4B / 8B | +| 24 GB | 0.6B / 1.7B | 0.6B / 1.7B[^week2-dense] | 4B / 8B | 4B / 8B | | 32 GB | 4B / 8B | 4B / 8B | 4B / 30B-A3B[^moe] | 4B / 30B-A3B[^moe] | | 36 GB | 4B / 8B | 4B / 8B | 4B / 30B-A3B[^moe] | 4B / 30B-A3B[^moe] | | 48 GB | 4B / 8B | 4B / 8B | 4B / 30B-A3B[^moe] | 4B / 30B-A3B[^moe] | @@ -194,9 +143,9 @@ default batch settings. Week 1 reads an official 4-bit checkpoint but materializes its linear and embedding weights in BF16. On an 8 GB Mac, keep the required path at 0.6B. On a 16–24 GB Mac, use 0.6B for the required work and treat 1.7B as an upper-end experiment. -Week 2 Days 1–2 -retain that dense BF16 model; Day 3 keeps weights packed for the quantized-matvec checkpoint. Weeks 3 -and 4 inherit that packed path. More memory still helps after reaching the largest +Week 2 Day 1 retains that dense BF16 model. A later lesson will keep +weights packed for the quantized-matvec checkpoint; the Week 3 and 4 paths +expect that later interface. More memory still helps after reaching the largest supported model because prompt length, batch size, KV caches, compilation, macOS, and other applications all share the same pool. These ceilings are therefore planning guidance, not a guarantee that every workload will avoid memory pressure. @@ -206,8 +155,8 @@ pressure. [M4 Mac mini](https://support.apple.com/en-us/121555), and [M5 MacBook Air](https://support.apple.com/en-us/126320) specifications. Higher-memory configurations are outside this table. -[^week2-dense]: Week 2 Days 1–2 use the dense Week 1 loader, so keep using the Week 1 recommendation - until the packed quantized-matvec path is complete on Day 3. The larger Week 2 entries apply after that checkpoint. +[^week2-dense]: Week 2 Day 1 uses dense BF16 weights and the same model-size + guidance as Week 1. Larger packed-model options belong to a later lesson. [^moe]: 30B-A3B requires the optional Week 3 MoE implementation. In Week 4, select the Week 3 loader. Use batch size one and a short context when approaching this ceiling; 4B remains the required-course target. diff --git a/book/src/week1-07-sampling-prepare.md b/book/src/week1-07-sampling-prepare.md index c8c19aeb..c84a0d66 100644 --- a/book/src/week1-07-sampling-prepare.md +++ b/book/src/week1-07-sampling-prepare.md @@ -95,10 +95,12 @@ Larger models remain optional and are not a Day 7 completion gate. ## Task 2: Prepare for Week 2 -Week 2 Days 1 and 2 introduce KV caching in Python, so you can begin them before -the custom-extension toolchain is ready. Starting on Day 3, the C++ and Metal -work requires full Xcode, its command-line tools, the Metal compiler, and CMake -3.27 or newer. +Day 1's KV-cache implementation is in Python, but its supplied tests import +the Week 2 native extension when they collect. Prepare and build that extension +before the first Day 1 test. This toolchain requires full Xcode, its +command-line tools, the Metal compiler, and CMake 3.27 or newer. A later +packed-W4 lesson will use Metal for the work itself; that lesson is not part +of the current Day 1 route. 1. **Install Xcode:** @@ -177,13 +179,14 @@ The other exported extension names are fail-closed starter stubs labeled with the Week 2 or Week 3 checkpoint that implements them; this setup check calls only `axpby`. -If you are new to C++ or Metal, try a few small exercises before the custom -kernel work on Day 3. For example, implement element-wise operations such as +If you are new to C++ or Metal, try a few small exercises before the later +custom-kernel lessons. For example, implement element-wise operations such as `exp`, `sin`, and `cos`, then use them in place of the corresponding MLX operations in your model implementation. That completes Week 1: you now have a single-request Python inference loop that loads Qwen3, computes logits, samples tokens, and streams a response. Week 2 -first adds KV caching in Python, then begins the custom Metal kernel path. +first adds KV caching in Python with the extension built for test collection. +The custom Metal kernel path follows in a later release. {{#include copyright.md}} diff --git a/book/src/week2-01-kv-cache.md b/book/src/week2-01-kv-cache.md index d17dbb4d..871dabcc 100644 --- a/book/src/week2-01-kv-cache.md +++ b/book/src/week2-01-kv-cache.md @@ -1,7 +1,8 @@ -# 🚧 Week 2 Day 1: KV Cache +# 🚧 Week 2 Day 1: Reuse the Prefix, Then Bound the Cache Your Week 1 Qwen model already generates by rerunning the full prefix. Day 1 -keeps that path intact while you complete four separate Week 2 shells: +keeps that path intact while you first complete the `kv-cache` checkpoint, then +bound the request-owned storage for `capacity-cache`. The four initial shells are: - `src/tiny_llm/kv_cache.py::TinyKvFullCache` stores one layer's dense K/V; - `src/tiny_llm/qwen3_week2.py::Qwen3ModelWeek2` threads cache state and @@ -11,15 +12,20 @@ keeps that path intact while you complete four separate Week 2 shells: Together, these pieces make prefill populate the cache and make decode send only the new token. The starter already supplies the Week 1 operators and the -model-loading boundary. Start with the focused learner gate: +model-loading boundary. Complete the Week 1 Day 7 toolchain setup first. The +Day 1 tests import the native Week 2 extension even though this first cache +change is in Python, so build it before running the focused learner gate: ```bash -pdm run test --week 2 --day 1 +pdm run build-ext +pdm run test --week 2 --day 1 -- -k 'not capacity' ``` -When it passes, run the `kv-cache` checkpoint shown in Task 4. That live call -puts the cache into the generation loop instead of exercising it only as an -isolated data structure. +The three selected tests cover the initial `kv-cache` work, not the later +capacity checkpoint. They may fail at the learner-owned seams until you finish +Tasks 1–3. Rerun this focused gate after connecting the serving loop in Task +4, then run the `kv-cache` product command there. The whole-day gate follows +the capacity checkpoint below. Each attention layer can then reuse the keys and values from previous tokens instead of recomputing the entire prefix at every step. @@ -72,7 +78,7 @@ Assume that each attention head has dimension `D = 4`: ``` L = 3 -Q x K^T = +Q x K^T = 1 1 1 1 1 2 3 1x1 -inf -inf 2 2 2 2 1 2 3 2x1 2x2 -inf 3 3 3 3 1 2 3 3x1 3x2 3x3 @@ -102,17 +108,17 @@ K in cache: [a b c d] represent cached values L = 1, S = 3 -Q x K^T = +Q x K^T = (⬇️ is K not transposed) - [1 1 1 1] - [2 2 2 2] + [1 1 1 1] + [2 2 2 2] 3 3 3 3 3 3 3 3 3x1 3x2 3x3 L = 1, S = 4 -Q x K^T = +Q x K^T = (⬇️ is K not transposed) - [1 1 1 1] - [2 2 2 2] + [1 1 1 1] + [2 2 2 2] [3 3 3 3] 4 4 4 4 4 4 4 4 4x1 4x2 4x3 4x4 ``` @@ -161,8 +167,9 @@ Keep this first cache deliberately simple and dense. Each `mx.concat` allocates a larger buffer and copies the previous K/V contents. Across a token-by-token decode of length `S`, those copies add up to `O(S²)` bytes even though the cache avoids `O(S²)` prefix recomputation. The reference cache records that traffic -as `growth_copy_bytes` so the profiler can separate it from attention. Week 3 -replaces repeated concatenation with preallocated pages for serving. +as `growth_copy_bytes` so the profiler can separate it from attention. The +second checkpoint in this lesson replaces repeated concatenation with a +request-bounded allocation. Week 3 later introduces pages for serving. ## Task 2: Build the Cached Week 2 Model @@ -224,8 +231,8 @@ Implement `create_kv_cache` so each request receives one cache handle per Transformer layer. Pass the matching cache through each block, and keep the caller's offset equal to the cache's logical length. -The Day 1 test checks this request-scoped lifecycle together with the cache and -model work from the earlier tasks. +The first Day 1 gate checks this request-scoped lifecycle together with the +cache and model work from the earlier tasks. ## Task 4: Connect the Serving Loop @@ -252,37 +259,334 @@ You can test your solution with: ```bash pdm run main --solution tiny_llm --loader week2 \ - --week2-checkpoint kv-cache --model qwen3-4b + --week2-checkpoint kv-cache --model qwen3-0.6b --max-tokens 16 ``` You can also run the same loop with the reference solution: ```bash pdm run main --solution tiny_llm_ref --loader week2 \ - --week2-checkpoint kv-cache --model qwen3-4b + --week2-checkpoint kv-cache --model qwen3-0.6b --max-tokens 16 +``` + +## Measure the First Cache Change + +Before adding capacity, compare the Week 1 full-prefix loop with this +`kv-cache` checkpoint. The progression runner starts each variant in a fresh +process, uses the same Qwen3-4B 128-token prompt and 129-token output, and +records the workload and results in JSON: + +```bash +pdm run bench-week2-progression --offline --solution tiny_llm --suite week2 \ + --repeats 2 --variant week1 --variant week2-kv-cache \ + --model qwen3-4b --input-len 128 --output-len 129 --warmup 2 \ + --prefill-logits all --json-output week2-day1-cache.json +``` + +Keep `week2-day1-cache.json` as the baseline for the capacity checkpoint +below. Week 1 requires `--prefill-logits all`; this is a matched algorithm +comparison, not a serving-only prefill measurement. Record the observation +and workload identity without assuming its speedup holds on another machine. + + +## Second checkpoint: bound the request cache + +The dense cache has removed repeated model computation, but its concatenation +still copies the old K/V prefix on each append. Physical capacity and logical +length are different: only the logical prefix is visible to attention. The +`capacity-cache` checkpoint keeps the same serving loop and readable operators. +It allocates storage from the request's prompt-plus-output bound, writes new +K/V into a slice, and returns a view of the active prefix. Capacity can raise +peak memory for some shapes while removing repeated prefix copies. + +## First Diagnostic: Hide Unused Capacity + +```bash +pdm run test --week 2 --day 1 -- -k logical_prefix +``` + +The expected first failure points at `_logical_key_values` or the capacity +branch of `update_and_fetch`. For physical arrays shaped `B, H, capacity, D`, +attention must see only `:offset`: + +```text +physical storage: [token 0][token 1][unused][unused] +logical prefix: [token 0][token 1] +offset = 2, capacity = 4 +``` + +Do not infer logical length from the backing array's shape. + +## Allocate from the Request Bound + +The generation loop knows the prompt length and maximum number of new tokens. +Use that request-local bound when it creates each layer cache. Capacity is not a +global maximum and must not grow beyond the request's declared budget. + +On the first append, allocate K/V storage for the full capacity. On every +append: + +1. compute `end = offset + L_new`; +2. reject `end > capacity` before changing storage, offset, or counters; +3. use `mx.slice_update` on sequence axis 2; +4. advance `offset` only after the write is valid; +5. return `storage[:, :, :offset, :]` for both K and V. + +That ordering makes overflow transactional. A rejected append leaves the old +logical cache usable. + +## Make Movement Observable + +The cache exposes three counter categories: + +| Counter | Meaning | Expected capacity behavior | +|---|---|---| +| `logical_copy_bytes` | old logical K/V copied by concatenation | zero | +| `physical_growth_copy_bytes` | old K/V copied while growing storage | zero | +| `slice_write_bytes` | newly written K/V bytes | increases by each append's K/V size | +| `growth_copy_bytes` | legacy total for growth copies | zero | + +Counters are mechanism evidence. They explain which bytes moved; they do not by +themselves establish lower complete-request latency or peak memory. + +## Reset and Rewind Without Leaking a Suffix + +`rewind(n)` shortens the logical length and rejects negative or oversized +rewinds. The next append may reuse the abandoned physical slots, but attention +must not see values beyond the new offset. `reset()` returns the logical cache +to length zero. It clears storage for the unbounded fallback but retains the +request-bounded allocation. All four movement counters remain +lifetime-cumulative across rewind and reset, so measure their deltas when you +need per-request evidence. + +Run the state-transition witness before the product: + +```bash +pdm run test --week 2 --day 1 -- -k 'rewind or overflow' +``` + +Test the sequence append → rewind → append as well as an overflow after valid +data. Those cases catch implementations that expose physical capacity as +logical state or mutate before validation. + +## Complete the `capacity-cache` Checkpoint + +```bash +pdm run test --week 2 --day 1 +pdm run main --solution tiny_llm --loader week2 \ + --week2-checkpoint capacity-cache --model qwen3-0.6b --max-tokens 16 +``` + +The predecessor fallback is `kv-cache`: it keeps the same generation algorithm +and readable model but uses concatenation. If bounded allocation cannot be +established, return to that checkpoint rather than exposing unused storage. + + +## Compare the two cache checkpoints + +After the focused gates, use identical prompts and output bounds for the two +public checkpoints. These coarse commands include process startup; use the +supplied progression runner below when you need separated prefill/decode +measurements and fresh-process repeats. + +```bash +/usr/bin/time -p pdm run main --solution tiny_llm --loader week2 \ + --week2-checkpoint kv-cache --model qwen3-0.6b --max-tokens 16 +/usr/bin/time -p pdm run main --solution tiny_llm --loader week2 \ + --week2-checkpoint capacity-cache --model qwen3-0.6b --max-tokens 16 ``` -## Integrate and Measure +The cache counters distinguish old logical-prefix copies, physical growth +copies, and new slice writes. A counter difference establishes the mechanism; +it does not alone prove a complete-request speedup. A historical exact-mechanism +run observed a **+88.0 MiB / +2.276%** temporal 2K/512 peak-memory tradeoff. +That is historical evidence, not a result from this checkout. -Finish Day 1 with a matched Week 1 versus cached Week 2 observation. The runner -uses fresh processes, applies the same Qwen3-4B 128×129 workload to both rows, -and writes the configuration beside the result: +Extend the first JSON baseline with the new capacity checkpoint using the +same model, lengths, warmups, and `all`-logit workload. Save a second JSON +file so the original Week 1 versus `kv-cache` observation remains available: ```bash -pdm run bench-week2-progression --offline --solution tiny_llm --repeats 2 \ - --variant week1 --variant week2-kv-cache \ +pdm run bench-week2-progression --offline --solution tiny_llm --suite week2 \ + --repeats 2 --variant week1 --variant week2-kv-cache \ + --variant week2-capacity-cache --model qwen3-4b \ + --input-len 128 --output-len 129 --warmup 2 \ + --prefill-logits all --json-output week2-day1-cache-ladder.json +``` + +Compare these rows with `week2-day1-cache.json` only when the recorded +workload and device match. The serving comparisons below use +`--prefill-logits last`, so keep their results separate from this ladder. + +## Benchmark the Cached Model + +Before changing the model, make the comparison trustworthy. Prefill processes +many prompt tokens at once, while decode usually processes one token per +request. At this checkpoint, decode repeatedly reads dense BF16 projection +weights. Because a change can help one phase while hurting the other, +`benches/bench.py` reports them separately: + +- prefill tokens per second: prompt tokens divided by prefill time; +- decode tokens per second: generated tokens after the first token divided by + decode time. + +The first generated token is part of prefill. Leaving it out of decode keeps +prompt length from distorting the decode number. + +Decide what prefill should return before comparing implementations. Prompt +scoring needs logits for every position; serving needs only the final prompt +logit. Use `--prefill-logits all` for the former and +`--prefill-logits last` for the latter. The runner applies one choice to your +solution and MLX alike, so the two rows do the same work. + +Keep the Week 2 generation algorithm matched too. Both sides use a KV cache: +prefill the prompt once, then pass only the newly generated token on each +decode step. A cached MLX baseline against a full-prefix solution would compare +two different algorithms instead of locating the next optimization target. + +### Record a Matched Baseline + +Use the same model, prompt length, output length, device, and warmup count for +your solution and MLX: + +```bash +pdm run bench --solution tiny_llm --loader week2 \ + --week2-checkpoint capacity-cache --model qwen3-4b \ + --num-seqs 1 --min-input-len 128 --max-input-len 128 \ + --min-output-len 65 --max-output-len 65 --warmup 2 \ + --prefill-logits last + +pdm run bench --solution mlx --loader week2 --model qwen3-4b \ + --num-seqs 1 --min-input-len 128 --max-input-len 128 \ + --min-output-len 65 --max-output-len 65 --warmup 2 \ + --prefill-logits last +``` + +Use `--solution tiny_llm_ref` with the same arguments when you want to compare +your solution with the reference solution instead of MLX. + +Or run the cumulative ladder in fresh processes: + +```bash +pdm run bench-week2-progression --offline --repeats 2 \ + --solution tiny_llm --suite week2 \ + --variant week2-capacity-cache --variant mlx \ --model qwen3-4b --input-len 128 --output-len 129 --warmup 2 \ - --json-output week2-day1-cache.json + --prefill-logits last --json-output week2-day1-baseline.json ``` -Keep this JSON as Day 2's baseline. Its useful result is the matched observation -and recorded workload identity, not a speedup claim for another model, prompt -length, output length, or device. +Benchmark on an otherwise idle machine. Stop other CPU- and GPU-intensive +workloads, keep power mode and ambient conditions fixed, and wait for a stable +temperature before comparing runs. Repeat each command, report the median, and +record the hardware, MLX and mlx-lm versions, prefill-logit mode, and exact +model. After a dependency upgrade, remeasure MLX instead of carrying the old +baseline forward. + +### Synchronize Lazy Work + +MLX builds computation graphs lazily. Timing only the Python call measures +graph construction instead of GPU execution, so every timed iteration must +evaluate its output: + +```python +start = perf_counter() +output = function() +mx.eval(output) +elapsed = perf_counter() - start +``` + +The benchmark must also call the cache release hook after warmups and timed +runs. That lets caches return owned or shared resources even when a run fails; +the supplied benchmark-lifecycle test covers both paths. + +## Attribute the Cached Model + +Next, attribute the same cached-decode workload. Keep the learner solution, +model, decode phase, and 128-token context fixed: + +```bash +pdm run profile-week2-kernels --solution tiny_llm --model qwen3-4b \ + --case capacity-cache:decode:128 --warmup 4 --iterations 12 \ + --json-output week2-day1-attribution.json +``` + +The result identifies its source, checkpoint, phase, token count, prompt rule, +software, host, category medians, and category shares without depending on a +private function name or Metal symbol. An earlier checked M4 Pro run attributed +81.5% of cached-decode time to dense projections at the then-current `kv-cache` +checkpoint. This is a +historical example, not a measurement of your bounded-cache checkout; another +device or shape may point elsewhere. + +Turn the observation into a decision with three sentences: + +1. “Dense projections dominate this exact cached-decode workload.” +2. “Packing W4 weights and changing only the selected projection path should + reduce that category and improve matched decode.” +3. “I will reject or revise the hypothesis if projection time does not fall or + complete-model decode regresses under the same workload.” + +Substitute the category you observed for the checked example. Your required +work ends with the benchmark, attribution, and decision record. The +[earlier macOS 27 capture lab](./week2-advanced-profiling.md) is historical; no trace, +`gpudebug` output, screenshot, or device-specific counter gates the W4 lesson. + +## Why Quantize: The Decode Roofline + +The measurement now has a hardware reason to test. LLM decode is typically +**memory-bandwidth bound**: each token reads the model's weights while doing +relatively little work with them. Use the dimensions in the official +[Qwen3-4B configuration](https://huggingface.co/Qwen/Qwen3-4B/blob/main/config.json) +to calculate the ideal bound: + +```plain +Qwen3-4B dimensions: + hidden size h = 2,560 + MLP size i = 9,728 + query width q = 4,096 + key/value width kv = 1,024 + layers L = 36 + vocabulary V = 151,936 + +Projection weights per layer: + Q and O: 2 × h × q = 20,971,520 + K and V: 2 × h × kv = 5,242,880 + MLP: 3 × h × i = 74,711,040 + total per layer = 100,925,440 + +All transformer layers: L × 100,925,440 = 3,633,315,840 +Tied vocabulary head: V × h = 388,956,160 +Total streamed weights: 4,022,272,000 + +FLOPs per token: 2 × 4,022,272,000 = 8.045 GFLOPs +``` + +Count the tied embedding matrix once as the vocabulary projection. The +single-row embedding lookup, normalization weights, activations, KV reads, and +attention work are omitted, so the result is an upper bound for linear layers +rather than a prediction of complete-model throughput. A dense FP16 or BF16 +weight occupies two bytes: + +```plain +4,022,272,000 weights × 2 bytes = 8.045 GB per token +arithmetic intensity = 8.045 GFLOPs / 8.045 GB = 1.0 FLOP/byte +``` -Day 1 changes the generation algorithm by removing full-prefix recomputation, -so measure it with the end-to-end benchmark rather than inventing a -shader-level limiter from a GPU trace. On Day 2, attribute this exact cached -workload and turn the observation into a falsifiable next change. Begin Day 3 -only after that evidence names dense projections. +FP16 and BF16 divide their 16 bits differently: FP16 gives more bits to the +significand, while BF16 gives more bits to the exponent. That affects numerical +range and precision, but not this bandwidth calculation. The course uses BF16 +for activations and outputs. + +| Dense weight format | Bits per weight | Bytes per weight | Streamed weight bytes per token | Weight arithmetic intensity | +|---|---:|---:|---:|---:| +| FP16 | 16 | 2 | 8.045 GB | 1.0 FLOP/byte | +| BF16 | 16 | 2 | 8.045 GB | 1.0 FLOP/byte | + +This is the baseline to improve: both dense formats must stream roughly 8 GB +of projection weights to generate one token. Save the matched benchmark result, +and keep it for comparison when the packed-W4 Day 2 lesson ships. The +[earlier quantization lesson](./week2-03-quantize-model.md) preserves that +mechanism, but its old checkpoint commands are historical in this Day 1 +checkout. {{#include copyright.md}} diff --git a/book/src/week2-02-benchmark-profile.md b/book/src/week2-02-benchmark-profile.md index bd4309e3..9d6daaf3 100644 --- a/book/src/week2-02-benchmark-profile.md +++ b/book/src/week2-02-benchmark-profile.md @@ -1,4 +1,11 @@ -# 🚧 Week 2 Day 2: Benchmarking and Profiling +# Historical Week 2: Benchmarking And Profiling + +> **Earlier lesson address.** This page preserves the benchmarking and profiling +> explanation and its original links. The current Week 2 learner route +> ships [Day 1: Cache and Measure](./week2-01-kv-cache.md) only. +> Later checkpoints, day numbers, tests, and commands below belong to an +> earlier all-days course state; do not use them as gates for this checkout. + Day 1 leaves you with a cached BF16 model and a working `kv-cache` checkpoint. Day 2 does not add another model operator. The supplied benchmark and portable diff --git a/book/src/week2-03-quantize-model.md b/book/src/week2-03-quantize-model.md index cf88067e..9f46582f 100644 --- a/book/src/week2-03-quantize-model.md +++ b/book/src/week2-03-quantize-model.md @@ -1,4 +1,11 @@ -# 🚧 Week 2 Day 3: Quantize the Model +# Historical Week 2: Packed W4 Quantization + +> **Earlier lesson address.** This page preserves the packed W4 quantization +> explanation and its original links. The current Week 2 learner route +> ships [Day 1: Cache and Measure](./week2-01-kv-cache.md) only. +> Later checkpoints, day numbers, tests, and commands below belong to an +> earlier all-days course state; do not use them as gates for this checkout. + Day 2 leaves you with a synchronized dense BF16 baseline. The Day 3 starter already supplies the packed-weight container and its diff --git a/book/src/week2-04-fused-model-kernels.md b/book/src/week2-04-fused-model-kernels.md index bf6f7d9b..6fd13b68 100644 --- a/book/src/week2-04-fused-model-kernels.md +++ b/book/src/week2-04-fused-model-kernels.md @@ -1,4 +1,11 @@ -# 🚧 Week 2 Day 4: Fused Model Kernels +# Historical Week 2: Fused Model Kernels + +> **Earlier lesson address.** This page preserves the fused model kernels +> explanation and its original links. The current Week 2 learner route +> ships [Day 1: Cache and Measure](./week2-01-kv-cache.md) only. +> Later checkpoints, day numbers, tests, and commands below belong to an +> earlier all-days course state; do not use them as gates for this checkout. + Day 3 leaves the cached model using packed projections. Day 4 keeps the Week 1 Python equations as readable oracles and completes three separate extension diff --git a/book/src/week2-05-simd-matrix-prefill.md b/book/src/week2-05-simd-matrix-prefill.md index 6ad1b78d..96e6a1c8 100644 --- a/book/src/week2-05-simd-matrix-prefill.md +++ b/book/src/week2-05-simd-matrix-prefill.md @@ -1,4 +1,11 @@ -# 🚧 Week 2 Day 5: SIMD-Matrix Prefill +# Historical Week 2: Simd Matrix Prefill + +> **Earlier lesson address.** This page preserves the SIMD matrix prefill +> explanation and its original links. The current Week 2 learner route +> ships [Day 1: Cache and Measure](./week2-01-kv-cache.md) only. +> Later checkpoints, day numbers, tests, and commands below belong to an +> earlier all-days course state; do not use them as gates for this checkout. + Day 4 ends with a decision, not a predetermined kernel. Re-profile the fixed 128-token prefill and name the dominant category before changing code. On the diff --git a/book/src/week2-06-operator-lab.md b/book/src/week2-06-operator-lab.md index 1f4724a7..e4b86b66 100644 --- a/book/src/week2-06-operator-lab.md +++ b/book/src/week2-06-operator-lab.md @@ -1,4 +1,11 @@ -# 🚧 Week 2 Day 6 (Optional): Workload-Conditioned Operator Lab +# Historical Week 2: Bounded Decode-Attention Operator Lab + +> **Earlier lesson address.** This page preserves the bounded decode-attention operator lab +> explanation and its original links. The current Week 2 learner route +> ships [Day 1: Cache and Measure](./week2-01-kv-cache.md) only. +> Later checkpoints, day numbers, tests, and commands below belong to an +> earlier all-days course state; do not use them as gates for this checkout. + Day 5 restores the matrix-shaped projection path selected by the fixed 128-token prefill profile. Day 6 asks a different question: can a secondary diff --git a/book/src/week2-07-split-k-prefill.md b/book/src/week2-07-split-k-prefill.md index 50bffd44..7d394db2 100644 --- a/book/src/week2-07-split-k-prefill.md +++ b/book/src/week2-07-split-k-prefill.md @@ -1,4 +1,11 @@ -# 🚧 Week 2 Day 7: Conditional Split-K and Final Decision +# Historical Week 2: Conditional Split-K + +> **Earlier lesson address.** This page preserves the conditional Split-K +> explanation and its original links. The current Week 2 learner route +> ships [Day 1: Cache and Measure](./week2-01-kv-cache.md) only. +> Later checkpoints, day numbers, tests, and commands below belong to an +> earlier all-days course state; do not use them as gates for this checkout. + Day 5 leaves a reusable 32×32×32 SIMD-matrix projection and an exact unsplit fallback. Day 6 is an optional branch and is not inherited here: the `split-k` diff --git a/book/src/week2-advanced-profiling.md b/book/src/week2-advanced-profiling.md index 9bc4f437..eef8df15 100644 --- a/book/src/week2-advanced-profiling.md +++ b/book/src/week2-advanced-profiling.md @@ -1,4 +1,11 @@ -# Optional: Inspect a Week 2 Capture on macOS 27 +# Historical Week 2: Optional Macos Capture + +> **Earlier lesson address.** This page preserves the optional macOS capture +> explanation and its original links. The current Week 2 learner route +> ships [Day 1: Cache and Measure](./week2-01-kv-cache.md) only. +> Later checkpoints, day numbers, tests, and commands below belong to an +> earlier all-days course state; do not use them as gates for this checkout. + The synchronized product benchmark and portable operator-attribution runner are sufficient for every required Week 2 checkpoint. This page is an optional diff --git a/book/src/week2-overview.md b/book/src/week2-overview.md index a6b4e9a6..eed665a3 100644 --- a/book/src/week2-overview.md +++ b/book/src/week2-overview.md @@ -2,158 +2,56 @@ tiny-llm-book © 2022-2026 by Alex Chi Z is licensed under CC BY-NC-SA 4.0 --> -# 🚧 Week 2: A Step Closer to vLLM - -Week 1 leaves you with a readable Qwen3 model that can generate text. This week, -you will turn it into a measured single-request serving path. After each change, -you will rerun the same synchronized workload, find where time now goes, and use -that evidence to choose what to change next. Instead of collecting unrelated -kernels, you will build an optimization story you can explain. - -Days 1–5 form the main route. Day 6 is an optional lab for an operator that -matters to a workload you choose. On Day 7, you will test Split-K where a -short-shape measurement suggests it may help, then make a final keep-or-reject -decision on the fixed product workload. - -Begin with [Day 1: KV Cache](./week2-01-kv-cache.md), where you will stop -recomputing the entire prefix for every generated token. - -> ⏱️ **Time commitment.** Days 3–5 introduce custom Metal kernels and may take -> substantially longer than Week 1 Days 6–7. You may skip Day 6. On Day 7, your -> Split-K implementation must be correct, but the course does not require an -> absolute speed or a universal crossover point. - -Week 2 keeps BF16 for dense weights, quantization scales and biases, -activations, projections, KV-cache entries, and model-facing kernel outputs. -Packed W4 weight codes use `uint32`. Reductions, dot products, and -online-softmax state accumulate in FP32 before returning BF16. Week 3 inherits -these interfaces and precision boundaries. - -## Measure, Change, and Measure Again - -Use the same four-part loop at each checkpoint: - -1. **Check correctness.** Run the focused supplied test before timing. -2. **Describe the workload.** Record model, checkpoint, phase, token counts, - prefill-logit mode, warmups, iterations, software, and device. -3. **Find the next cost.** Name the dominant operator category and make one - bounded hypothesis before editing. -4. **Rerun and decide.** Repeat the identical product and attribution workload, - then record `keep`, `reject`, or `inconclusive`, along with evidence that - would change your conclusion. - -The checked example makes that loop visible. Cached decode begins with -projection work dominant; packed W4 exposes the pointwise category; fused -model kernels leave projections as the next target; and the SIMD schedule -shrinks prefill projection time. Split-K helps the measured 32-token shape but -does not improve the fixed 128-token product control. - -![Stacked attribution bars for the checked Week 2 checkpoints in chronological order: cached decode on Day 2, packed W4 on Day 3, fused model kernels on Day 4, the pre-SIMD control immediately before the SIMD rows on Day 5, the optional decode-attention lab on Day 6, and Split-K on Day 7.](./week2-kernel-profile.svg) - -The decisions below use different metrics and denominators, so read each card -as a bounded comparison rather than adding the percentages together. Packed -W4, fused pointwise kernels, and SIMD prefill are kept for their measured -controls; decode attention stays optional and inconclusive; Split-K is -conditional at 32 tokens and rejected for the fixed 128-token workload. - -![Decision cards for one M4 Pro, macOS 27, Qwen3-4B, MLX 0.32.0 observation: keep packed W4 decode, fused pointwise kernels, and SIMD-matrix prefill; treat decode attention as optional and inconclusive; retain Split-K conditionally at 32 tokens but reject it for the fixed 128-token product workload.](./week2-performance-summary.svg) - -You can complete this loop with the synchronized benchmark and portable -attribution runner. Apple GPU capture and `gpudebug` appear only in the -[optional macOS 27 lab](./week2-advanced-profiling.md); neither is a -prerequisite. - -## Daily Checkpoints - -1. **KV cache:** make decode incremental and compare it with the Week 1 model on - a matched workload. -2. **Discover:** learn to synchronize a measurement, attribute the cached - model, and choose one bounded optimization. In the checked run, dense - projections became the next target. -3. **Packed W4 matvec:** keep weights packed while you optimize decode - projections, then re-profile. In the checked run, normalization, position, - and activation work became visible next. -4. **Fused model kernels:** implement RMSNorm, RoPE, and SwiGLU one at a time. - Keep each change only after a matched measurement, then re-profile prefill - before choosing Day 5. -5. **SIMD-matrix prefill:** replace the matrix-shaped projection schedule chosen - from the fixed 128-token prefill attribution. -6. **Optional operator lab:** choose a secondary operator category for one - explicit workload. The supplied branch studies bounded decode attention, - but an equivalent evidence-led operator experiment is also valid. -7. **Conditional Split-K and final decision:** try an under-filled 32-token - projection while preserving the unsplit Day 5 fallback. Finish by rerunning - the fixed 128×129 product workload and deciding whether to keep the change. - -## What Is Supplied and What You Own - -The starter gives you model loading, the extension build system, benchmark and -attribution runners, correctness tests, Python reference equations, stable -checkpoint interfaces, and a compact checked M4 Pro evidence file. You will -build the cache transition, integrate packed weights, implement the custom -operators, and turn each measurement into a decision record. - -The completed course path uses your implementations, not MLX replacements, for -the operators it asks you to build. If you want to reach the later serving -mechanisms without implementing one custom kernel, keep the course interface -and connect the corresponding MLX operator locally. This substitution stays -inside your course model. It is different from `--solution mlx`, which runs the -separate full-MLX model. - -## Check Your Progress - -Run the canonical selector after each day: - -| Course day | Test command | -|---|---| -| Day 1 | `pdm run test --week 2 --day 1` | -| Day 2 | `pdm run test --week 2 --day 2` | -| Day 3 | `pdm run test --week 2 --day 3` | -| Day 4 | `pdm run test --week 2 --day 4` | -| Day 5 | `pdm run test --week 2 --day 5` | -| Day 6 (optional) | `pdm run test --week 2 --day 6` | -| Day 7 | `pdm run test --week 2 --day 7` | - -When a command runs a model, benchmark, profile, capture, or reducer, pass -`--solution tiny_llm` exactly as the chapter shows. Some command-line tools -otherwise default to the completed reference, so omitting it may measure code -you did not write. - -### Bring Forward Work from the Earlier Day Order - -An earlier course order put decode attention on Day 5 and SIMD-matrix prefill on -Day 6. If your checkout contains work from that order, you can keep it. First -complete the current Day 5 SIMD gate, then use the optional Day 6 gate to check -your retained attention implementation. The ordinary `--week 2 --day 5` and -`--week 2 --day 6` commands above are the only selectors you need. Old Day 5 -and Day 6 bookmarks now redirect to the corresponding canonical lessons. - -## What the Gates Check - -Most required gates check public behavior: checkpoint and workload identity, -operator results, fallbacks, synchronized output, and the decision-record -schema. You may organize most internals differently. Course-ownership and -extension-integration witnesses intentionally preserve explicit source and -header seams. The gates do not grade device timings or require the optional -`gpudebug` tooling. - -Read the checked example as one machine's optimization story, not a portable -speed claim. Its absolute measurements come from one M4 Pro running macOS 27 -with Qwen3-4B, a fixed 128-token prompt and 129-output-token product control, -and `n=2` balanced product samples. Six of eight captures exposed full -shader/counter detail. The pre-SIMD prefill capture did not expose a shader -ranking, while the Split-K capture exposed only static dispatch. The example -marks the missing data unavailable instead of guessing. - -## Continue to Week 3 - -By the end of Week 2, your model decodes one token at a time from a dense KV -cache, chooses separate prefill and decode projection schedules, and keeps its -weights quantized. Week 3 keeps these model, cache, precision, and operator -interfaces while adding paging and batching. You do not need the optional Day -6 attention branch to continue, and Day 7 begins from Day 5's unsplit SIMD path. - -The [performance evidence ledger](./appendix-performance.md) shows the checked -causal example and its limits. +# 🚧 Week 2: A Faster Single Request + +Week 1 leaves you with a readable Qwen3 model that regenerates from the full +prefix. The **current Week 2 route contains Day 1 only**: reuse previous keys +and values, then give their dense storage a request bound. The two cumulative +checkpoints are `kv-cache` and `capacity-cache`. + +Begin with [Day 1: Cache and Measure](./week2-01-kv-cache.md). Its first +feedback loop builds the extension required for test collection, runs the +focused KV test, and sends only the new token during decode. Its second loop +bounds storage without exposing unused capacity to attention. After both +checkpoints, compare the same request across Week 1, `kv-cache`, and +`capacity-cache`; keep the serving-only comparison separate from the +all-logit algorithm comparison. + +## Day 1 route + +| Step | What you own | Feedback | +|---|---|---| +| Prepare | Build the Week 2 extension after [Week 1 Day 7](./week1-07-sampling-prepare.md) | `pdm run build-ext` | +| Cache the prefix | Implement dense K/V reuse in the model and generation loop | `pdm run test --week 2 --day 1 -- -k 'not capacity'` | +| Bound the cache | Allocate from the request limit, expose only the logical prefix, and preserve reset/rewind/overflow behavior | `pdm run test --week 2 --day 1` after the capacity work | +| Measure | Keep the workload and prefill-logit mode matched | [Day 1 measurement loop](./week2-01-kv-cache.md#measure-the-first-cache-change) | + +The starter supplies model loading, the extension build, test entrypoints, +and benchmark and attribution helpers. You own the cache state, its model +wiring, and the serving loop. The reference solution and full MLX model are +separate controls; they do not fill your learner TODOs. A cache counter shows +which bytes moved, while a synchronized complete-request comparison shows +whether that mechanism helped the chosen workload. + +## Later lessons + +The reviewed five-day design continues with packed W4, SIMD matrix prefill, +model primitives, and tiled dense prefill attention. Those **Day 2–5 +checkpoints are planned, not shipped in this Day 1 route**. Their commands and +selectors are not Day 1 gates. Week 3 uses later Week 2 interfaces; Day 1 +alone does not supply every prerequisite for its learner exercises. + +The earlier seven-day book remains available at its old addresses as +[historical Week 2 material](./week2-02-benchmark-profile.md). It preserves +benchmark method, W4 derivation, Apple M1–M4 bandwidth and roofline +calculations, fusion/SIMD mechanisms, optional capture, and the old +bounded-decode and Split-K experiments. Its +[operator-attribution diagram](./week2-kernel-profile.svg) and +[decision diagram](./week2-performance-summary.svg) are also historical +evidence, not diagrams of the Day 1 checkout. Those pages describe a different +checkpoint order and may show commands unavailable in this partial branch. +Use Day 1 above for the current learner workflow. The +[performance evidence ledger](./appendix-performance.md) is likewise +historical context, not a performance claim for this checkout. {{#include copyright.md}} From 88a4e5184a4f88960dd9f73df096db8d1ace77fe Mon Sep 17 00:00:00 2001 From: Forge Date: Wed, 23 Sep 2026 12:41:49 -0700 Subject: [PATCH 2/5] Integrate Week 2 Day 1 cache lesson and shipped-day gates --- .github/workflows/macos.yml | 6 +- README.md | 2 +- benches/bench.py | 20 +- benches/bench_course_progression.py | 50 +---- benches/profile_week2_kernels.py | 8 +- benches/week2_gpudebug.py | 8 +- main.py | 20 +- pyproject.toml | 1 - scripts/dev-tools.py | 3 + scripts/week2_shipped_day_ci.json | 177 ++++++++++++++++ scripts/week2_shipped_day_ci.py | 77 +++++++ src/tiny_llm/generate.py | 1 + src/tiny_llm/kv_cache.py | 28 ++- src/tiny_llm/qwen3_week2.py | 53 +---- src/tiny_llm_ref/generate.py | 24 ++- src/tiny_llm_ref/kv_cache.py | 67 +++++- src/tiny_llm_ref/qwen3_week2.py | 52 ++--- tests_refsol/test_dense_kv_capacity.py | 254 ++++++++++++++++++++++ tests_refsol/test_week_2_day_1.py | 106 +++++++++- tests_refsol/test_week_2_day_2.py | 76 ------- tests_refsol/test_week_2_day_3.py | 196 ----------------- tests_refsol/test_week_2_day_4.py | 120 ----------- tests_refsol/test_week_2_day_5.py | 282 ------------------------- tests_refsol/test_week_2_day_6.py | 171 --------------- tests_refsol/test_week_2_day_7.py | 114 ---------- 25 files changed, 772 insertions(+), 1144 deletions(-) create mode 100644 scripts/week2_shipped_day_ci.json create mode 100644 scripts/week2_shipped_day_ci.py create mode 100644 tests_refsol/test_dense_kv_capacity.py delete mode 100644 tests_refsol/test_week_2_day_2.py delete mode 100644 tests_refsol/test_week_2_day_3.py delete mode 100644 tests_refsol/test_week_2_day_4.py delete mode 100644 tests_refsol/test_week_2_day_5.py delete mode 100644 tests_refsol/test_week_2_day_6.py delete mode 100644 tests_refsol/test_week_2_day_7.py diff --git a/.github/workflows/macos.yml b/.github/workflows/macos.yml index bb1d683e..01adf17e 100644 --- a/.github/workflows/macos.yml +++ b/.github/workflows/macos.yml @@ -38,12 +38,10 @@ jobs: - name: Try building extensions run: | pdm run build-ext - pdm run build-ext-test - - - run: pdm run build-ext-ref - run: cargo install mdbook-toc - run: cargo install mdbook-katex --version 0.10.0-alpha - uses: taiki-e/install-action@mdbook - - run: pdm run test-refsol + - name: Run printed shipped-day reference CI manifest + run: pdm run python scripts/week2_shipped_day_ci.py diff --git a/README.md b/README.md index 5d3e24a4..6235540e 100644 --- a/README.md +++ b/README.md @@ -125,7 +125,7 @@ one explicit byte range through the existing loop. | 4.8 | Fork, Steer, and Select | ✅ | ✅ | ✅ | 🚧 | | 4.9 | Bound Tool Evidence | ✅ | ✅ | ✅ | 🚧 | -The older Week 2 chapter URLs remain available as [historical material](book/src/week2-02-benchmark-profile.md). Their former Day 2–7 tests and checkpoint commands are not part of the current Day 1 learner route. The Day 1 code and test checks are still being completed. +The older Week 2 chapter URLs remain available as [historical material](book/src/week2-02-benchmark-profile.md). Their former Day 2–7 tests and checkpoint commands are not part of the current Day 1 learner route. Other topics not covered include quantized or compressed KV caches, cross-request prefix caching, fine-tuning, and long-context techniques. diff --git a/benches/bench.py b/benches/bench.py index d3f4398f..36f66560 100644 --- a/benches/bench.py +++ b/benches/bench.py @@ -86,16 +86,7 @@ def parse_args() -> argparse.Namespace: ) parser.add_argument( "--week2-checkpoint", - choices=( - "kv-cache", - "quantized-matvec", - "rmsnorm", - "rope", - "swiglu", - "simd-matmul", - "decode-attention", - "split-k", - ), + choices=("kv-cache", "capacity-cache"), help="run one cumulative Week 2 end-to-end checkpoint", ) parser.add_argument("--device", type=str, default="gpu", choices=["cpu", "gpu"]) @@ -165,12 +156,15 @@ def validate_args(args: argparse.Namespace) -> None: and args.device != "gpu" and ( args.loader == "week3" - or (args.loader == "week2" and args.week2_checkpoint != "kv-cache") + or ( + args.loader == "week2" + and args.week2_checkpoint not in (None, "kv-cache", "capacity-cache") + ) ) ): raise ValueError( - "The completed Week 2 and Week 3 custom-kernel models are GPU-only; " - "use the Week 2 kv-cache checkpoint for the readable pre-kernel path" + "Week 3 custom-kernel models are GPU-only; " + "Day 1 Week 2 checkpoints use the readable path" ) if args.disable_paged_attention and args.loader != "week3": raise ValueError("--disable-paged-attention requires --loader week3") diff --git a/benches/bench_course_progression.py b/benches/bench_course_progression.py index c1ab5212..6b0d21d7 100644 --- a/benches/bench_course_progression.py +++ b/benches/bench_course_progression.py @@ -42,59 +42,17 @@ class Throughput: WEEK1_VARIANT, Variant( "week2-kv-cache", - "2.1 KV cache", + "2.1 Reuse the prefix", "ref", "week2", ("--week2-checkpoint", "kv-cache"), ), Variant( - "week2-quantized-matvec", - "2.3 Quantized matvec", + "week2-capacity-cache", + "2.1 + Bound KV-cache movement", "ref", "week2", - ("--week2-checkpoint", "quantized-matvec"), - ), - Variant( - "week2-rmsnorm", - "2.4 Fast RMSNorm", - "ref", - "week2", - ("--week2-checkpoint", "rmsnorm"), - ), - Variant( - "week2-rope", - "2.4 + Fast RoPE", - "ref", - "week2", - ("--week2-checkpoint", "rope"), - ), - Variant( - "week2-swiglu", - "2.4 + Fused SwiGLU", - "ref", - "week2", - ("--week2-checkpoint", "swiglu"), - ), - Variant( - "week2-simd-matmul", - "2.5 SIMD matrix prefill", - "ref", - "week2", - ("--week2-checkpoint", "simd-matmul"), - ), - Variant( - "week2-decode-attention", - "2.6 Optional decode attention", - "ref", - "week2", - ("--week2-checkpoint", "decode-attention"), - ), - Variant( - "week2-split-k", - "2.7 Split-K prefill", - "ref", - "week2", - ("--week2-checkpoint", "split-k"), + ("--week2-checkpoint", "capacity-cache"), ), MLX_VARIANT, ) diff --git a/benches/profile_week2_kernels.py b/benches/profile_week2_kernels.py index a399089e..6c9fd0ec 100644 --- a/benches/profile_week2_kernels.py +++ b/benches/profile_week2_kernels.py @@ -25,13 +25,7 @@ DEFAULT_CASES = ( "kv-cache:decode:128", - "quantized-matvec:decode:128", - "swiglu:decode:128", - "simd-matmul:prefill:128", - "simd-matmul:prefill:32", - "decode-attention:decode:128", - "decode-attention:prefill:128", - "split-k:prefill:32", + "capacity-cache:decode:128", ) PROMPT_RULE = "synthetic-token-ids" PREFILL_LOGITS = "all" diff --git a/benches/week2_gpudebug.py b/benches/week2_gpudebug.py index fc22e093..99af6d50 100644 --- a/benches/week2_gpudebug.py +++ b/benches/week2_gpudebug.py @@ -19,13 +19,7 @@ PREFILL_LOGITS = "all" KNOWN_CHECKPOINTS = ( "kv-cache", - "quantized-matvec", - "rmsnorm", - "rope", - "swiglu", - "simd-matmul", - "decode-attention", - "split-k", + "capacity-cache", ) diff --git a/main.py b/main.py index 7316a655..e7db9692 100644 --- a/main.py +++ b/main.py @@ -34,16 +34,7 @@ ) parser.add_argument( "--week2-checkpoint", - choices=( - "kv-cache", - "quantized-matvec", - "rmsnorm", - "rope", - "swiglu", - "simd-matmul", - "decode-attention", - "split-k", - ), + choices=("kv-cache", "capacity-cache"), help="run one cumulative Week 2 model checkpoint", ) @@ -58,12 +49,15 @@ and args.device != "gpu" and ( args.loader == "week3" - or (args.loader == "week2" and args.week2_checkpoint != "kv-cache") + or ( + args.loader == "week2" + and args.week2_checkpoint not in (None, "kv-cache", "capacity-cache") + ) ) ): parser.error( - "The completed Week 2 and Week 3 custom-kernel models are GPU-only; " - "use the Week 2 kv-cache checkpoint for the readable pre-kernel path" + "Week 3 custom-kernel models are GPU-only; " + "Day 1 Week 2 checkpoints use the readable path" ) if args.disable_paged_attention and args.loader != "week3": parser.error("--disable-paged-attention requires --loader week3") diff --git a/pyproject.toml b/pyproject.toml index 5732e36c..594b31cb 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -39,7 +39,6 @@ main-week1.cmd = "python main.py --loader week1" main-week2.cmd = "python main.py --loader week2" main-week3.cmd = "python main.py --loader week3" bench.cmd = "python -m benches.bench" -bench-week2-operators.cmd = "python -m benches.bench_week2_operators" profile-week2-kernels.cmd = "python -m benches.profile_week2_kernels" capture-week2.cmd = "python -m benches.week2_gpudebug capture" reduce-week2-gpudebug.cmd = "python -m benches.week2_gpudebug reduce" diff --git a/scripts/dev-tools.py b/scripts/dev-tools.py index 5b94aed7..cd5c488c 100644 --- a/scripts/dev-tools.py +++ b/scripts/dev-tools.py @@ -15,6 +15,9 @@ def validate_week_day(args, required=False): if week_provided and (args.week <= 0 or args.day <= 0): print("Week and day must be positive integers") return False + if week_provided and args.week == 2 and args.day > 1: + print("Week 2 Day 1 is the only shipped day; later tests are deferred") + return False return True diff --git a/scripts/week2_shipped_day_ci.json b/scripts/week2_shipped_day_ci.json new file mode 100644 index 00000000..fd112a90 --- /dev/null +++ b/scripts/week2_shipped_day_ci.json @@ -0,0 +1,177 @@ +{ + "shipped_day": 1, + "included": [ + "tests_refsol/test_dev_tools.py", + "tests_refsol/test_model_names.py", + "tests_refsol/test_rope.py", + "tests_refsol/test_week_1_day_1.py", + "tests_refsol/test_week_1_day_2.py", + "tests_refsol/test_week_1_day_3.py", + "tests_refsol/test_week_1_day_4.py", + "tests_refsol/test_week_1_day_5.py", + "tests_refsol/test_week_1_day_6.py", + "tests_refsol/test_week_1_day_7.py", + "tests_refsol/test_week_2_day_1.py", + "tests_refsol/test_dense_kv_capacity.py", + "tests_refsol/test_week_4_day_1.py", + "benches/test_bench_week3.py" + ], + "deferred": [ + { + "path": "tests_refsol/test_week_2_day_2.py", + "reenable_day": 2, + "reason": "packed W4 and quantized matvec" + }, + { + "path": "tests_refsol/test_week_2_day_3.py", + "reenable_day": 3, + "reason": "SIMD prefill" + }, + { + "path": "tests_refsol/test_week_2_day_4.py", + "reenable_day": 4, + "reason": "RMSNorm, RoPE, and SwiGLU" + }, + { + "path": "tests_refsol/test_week_2_day_5.py", + "reenable_day": 5, + "reason": "tiled attention and selected engine" + }, + { + "path": "tests_refsol/test_extension_interface_sync.py", + "reenable_day": 5, + "reason": "cross-day native interface and chapter labels" + }, + { + "path": "benches/test_bench_course_progression.py", + "reenable_day": 5, + "reason": "cross-day progression and chapter labels" + }, + { + "path": "benches/test_profile_week2_kernels.py", + "reenable_day": 5, + "reason": "cross-day profile assertions" + }, + { + "path": "benches/test_week2_gpudebug.py", + "reenable_day": 5, + "reason": "cross-day capture assertions" + }, + { + "path": "benches/test_bench_week2_operators.py", + "reenable_day": 2, + "reason": "quantized operator bench" + }, + { + "path": "benches/test_quantized_matmul.py", + "reenable_day": 2, + "reason": "quantized native operator" + }, + { + "path": "benches/test_attention.py", + "reenable_day": 5, + "reason": "later attention kernels" + }, + { + "path": "tests_refsol/test_week_3_day_1.py", + "reenable_day": 5, + "reason": "requires reference-native RoPE after the Week 2 extension build" + }, + { + "path": "tests_refsol/test_week_3_day_2.py", + "reenable_day": 5, + "reason": "post-Week-2 regression gate" + }, + { + "path": "tests_refsol/test_week_3_day_3.py", + "reenable_day": 5, + "reason": "post-Week-2 regression gate" + }, + { + "path": "tests_refsol/test_week_3_day_4.py", + "reenable_day": 5, + "reason": "post-Week-2 regression gate" + }, + { + "path": "tests_refsol/test_week_3_day_5.py", + "reenable_day": 5, + "reason": "post-Week-2 regression gate" + }, + { + "path": "tests_refsol/test_week_3_day_6.py", + "reenable_day": 5, + "reason": "post-Week-2 regression gate" + }, + { + "path": "tests_refsol/test_week_3_day_7.py", + "reenable_day": 5, + "reason": "post-Week-2 regression gate" + }, + { + "path": "tests_refsol/test_week_4_day_2.py", + "reenable_day": 5, + "reason": "post-Week-2 regression gate" + }, + { + "path": "tests_refsol/test_week_4_day_3.py", + "reenable_day": 5, + "reason": "post-Week-2 regression gate" + }, + { + "path": "tests_refsol/test_week_4_day_4.py", + "reenable_day": 5, + "reason": "post-Week-2 regression gate" + }, + { + "path": "tests_refsol/test_week_4_day_5.py", + "reenable_day": 5, + "reason": "post-Week-2 regression gate" + }, + { + "path": "tests_refsol/test_week_4_day_6.py", + "reenable_day": 5, + "reason": "post-Week-2 regression gate" + }, + { + "path": "tests_refsol/test_week_4_day_7.py", + "reenable_day": 5, + "reason": "post-Week-2 regression gate" + }, + { + "path": "tests_refsol/test_week_4_day_8.py", + "reenable_day": 5, + "reason": "post-Week-2 regression gate" + }, + { + "path": "tests_refsol/test_week_4_day_9.py", + "reenable_day": 5, + "reason": "post-Week-2 regression gate" + }, + { + "path": "tests_refsol/test_week_4_starter_sync.py", + "reenable_day": 5, + "reason": "post-Week-2 starter regression gate" + }, + { + "path": "src/extensions/test.py", + "reenable_day": 5, + "reason": "native tests span future Week 2 operators" + }, + { + "path": "benches/bench_week2_operators.py", + "reenable_day": 2, + "reason": "operator benchmark alias begins with packed W4 work" + } + ], + "deferred_builds": [ + { + "command": "pdm run build-ext-ref", + "reenable_day": 2, + "reason": "reference-native packed W4 build fails in unchanged cooperative_matrix.h under macOS 27.0, Xcode 27.0, Metal toolchain 27.1.266.1; Day 1 cache paths invoke no native kernel" + } + ], + "retired": [ + "tests_refsol/test_week_2_day_6.py at old seven-day main", + "tests_refsol/test_week_2_day_7.py at old seven-day main" + ] +} diff --git a/scripts/week2_shipped_day_ci.py b/scripts/week2_shipped_day_ci.py new file mode 100644 index 00000000..6fb6c407 --- /dev/null +++ b/scripts/week2_shipped_day_ci.py @@ -0,0 +1,77 @@ +"""Run the reference tests that belong to the shipped Week 2 day. + +The JSON manifest is the reviewable temporary gate. Day 5 removes this +selector and restores the unfiltered reference suite. +""" + +import argparse +import json +import subprocess +import sys +from pathlib import Path + +ROOT = Path(__file__).resolve().parents[1] +MANIFEST = ROOT / "scripts" / "week2_shipped_day_ci.json" + + +def main() -> int: + parser = argparse.ArgumentParser() + parser.add_argument("--collect-only", action="store_true") + args = parser.parse_args() + manifest = json.loads(MANIFEST.read_text()) + included = manifest["included"] + deferred = manifest["deferred"] + declared = set(included) | {item["path"] for item in deferred} + discovered = { + str(path.relative_to(ROOT)) + for folder in ("tests_refsol", "benches") + for path in (ROOT / folder).glob("test*.py") + } + if len(declared) != len(included) + len(deferred): + raise SystemExit( + "duplicate included/deferred test path in shipped-day manifest" + ) + unaccounted = discovered - declared + missing_included = set(included) - discovered + if unaccounted or missing_included: + raise SystemExit( + f"unaccounted test files: {sorted(unaccounted)}; " + f"missing included files: {sorted(missing_included)}" + ) + if manifest["shipped_day"] != 1: + raise SystemExit("this temporary CI gate is bound to shipped_day=1") + print("shipped_day=1", flush=True) + print(f"included_files={len(included)}", flush=True) + for path in included: + print(f"INCLUDE {path}", flush=True) + for item in deferred: + print( + f"DEFER {item['path']} reenable_day={item['reenable_day']} " + f"reason={item['reason']}", + flush=True, + ) + if len(manifest["deferred_builds"]) != 1: + raise SystemExit("Day 1 must declare its one deferred reference-native build") + for item in manifest["deferred_builds"]: + if item["command"] != "pdm run build-ext-ref" or item["reenable_day"] != 2: + raise SystemExit("unexpected deferred native build in Day 1 manifest") + print( + f"DEFER BUILD {item['command']} reenable_day={item['reenable_day']} " + f"reason={item['reason']}", + flush=True, + ) + for path in manifest["retired"]: + print(f"RETIRED {path}", flush=True) + collect = subprocess.run( + [sys.executable, "-m", "pytest", "--collect-only", "-q", *included], + cwd=ROOT, + check=False, + text=True, + ) + if collect.returncode or args.collect_only: + return collect.returncode + return subprocess.call([sys.executable, "-m", "pytest", "-q", *included], cwd=ROOT) + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/src/tiny_llm/generate.py b/src/tiny_llm/generate.py index 74ee5fca..9d0059fc 100644 --- a/src/tiny_llm/generate.py +++ b/src/tiny_llm/generate.py @@ -28,6 +28,7 @@ def simple_generate_with_kv_cache( tokenizer: TokenizerWrapper, prompt: str, max_tokens: int = 256, + use_bounded_kv_capacity: bool | None = None, ) -> str: def _step(model, y, offset, kv_cache): pass diff --git a/src/tiny_llm/kv_cache.py b/src/tiny_llm/kv_cache.py index 3eb827d1..5fbb8f6e 100644 --- a/src/tiny_llm/kv_cache.py +++ b/src/tiny_llm/kv_cache.py @@ -1,8 +1,11 @@ from abc import ABC, abstractmethod -from typing import Optional +from typing import TYPE_CHECKING, Optional import mlx.core as mx +if TYPE_CHECKING: + from .paged_kv_cache import PagedKvMetadata + class TinyKvCache(ABC): @abstractmethod @@ -79,9 +82,26 @@ def remove_request(self, id: int): class TinyKvFullCache(TinyKvCache): - def __init__(self): + def __init__(self, capacity: int | None = None): + if capacity is not None and ( + not isinstance(capacity, int) or isinstance(capacity, bool) or capacity < 0 + ): + raise ValueError("capacity must be a non-negative integer or None") self.key_values = None self.offset = 0 + self.capacity = capacity + self.logical_copy_bytes = 0 + self.physical_growth_copy_bytes = 0 + self.slice_write_bytes = 0 + self.growth_copy_bytes = 0 + + @property + def uses_capacity(self) -> bool: + return self.capacity is not None + + def _logical_key_values(self) -> tuple[mx.array, mx.array]: + """Return only initialized tokens, never the unused physical tail.""" + pass def update_and_fetch( self, @@ -95,5 +115,9 @@ def update_and_fetch( def materialize(self): pass + def reset(self): + """Reset logical length while retaining bounded physical storage.""" + pass + def rewind(self, n: int): pass diff --git a/src/tiny_llm/qwen3_week2.py b/src/tiny_llm/qwen3_week2.py index c4cc3ea4..018bf3a4 100644 --- a/src/tiny_llm/qwen3_week2.py +++ b/src/tiny_llm/qwen3_week2.py @@ -4,19 +4,20 @@ import mlx.core as mx -from .embedding import Embedding +from .embedding import Embedding # noqa: F401 - learner checkpoint dependency from .kv_cache import TinyKvCache -from .quantize import QuantizedWeights, dequantize_linear +from .quantize import QuantizedWeights, dequantize_linear # noqa: F401 from .week2_kernels import ( - FastRMSNorm, - FastRoPE, - scaled_dot_product_attention, - swiglu, + FastRMSNorm, # noqa: F401 - learner checkpoint dependency + FastRoPE, # noqa: F401 - learner checkpoint dependency + scaled_dot_product_attention, # noqa: F401 - learner checkpoint dependency + swiglu, # noqa: F401 - learner checkpoint dependency ) @dataclass(frozen=True) class Week2CheckpointFeatures: + bounded_kv_capacity: bool = False quantized_weights: bool = False fast_rms_norm: bool = False fast_rope: bool = False @@ -29,40 +30,7 @@ class Week2CheckpointFeatures: WEEK2_CHECKPOINT_FEATURES = MappingProxyType( { "kv-cache": Week2CheckpointFeatures(), - "quantized-matvec": Week2CheckpointFeatures(quantized_weights=True), - "rmsnorm": Week2CheckpointFeatures(quantized_weights=True, fast_rms_norm=True), - "rope": Week2CheckpointFeatures( - quantized_weights=True, fast_rms_norm=True, fast_rope=True - ), - "swiglu": Week2CheckpointFeatures( - quantized_weights=True, - fast_rms_norm=True, - fast_rope=True, - fast_swiglu=True, - ), - "simd-matmul": Week2CheckpointFeatures( - quantized_weights=True, - fast_rms_norm=True, - fast_rope=True, - fast_swiglu=True, - simdgroup_matmul=True, - ), - "decode-attention": Week2CheckpointFeatures( - quantized_weights=True, - fast_rms_norm=True, - fast_rope=True, - fast_swiglu=True, - simdgroup_matmul=True, - decode_attention=True, - ), - "split-k": Week2CheckpointFeatures( - quantized_weights=True, - fast_rms_norm=True, - fast_rope=True, - fast_swiglu=True, - simdgroup_matmul=True, - split_k_matmul=True, - ), + "capacity-cache": Week2CheckpointFeatures(bounded_kv_capacity=True), } ) WEEK2_CHECKPOINTS = tuple(WEEK2_CHECKPOINT_FEATURES) @@ -162,13 +130,14 @@ class Qwen3ModelWeek2: def __init__( self, mlx_model: Any, - checkpoint: str = "split-k", + checkpoint: str = "capacity-cache", use_mlx_quantized_linear: bool = False, + use_bounded_kv_capacity: bool | None = None, ): self.num_hidden_layers = mlx_model.args.num_hidden_layers pass - def create_kv_cache(self) -> list[TinyKvCache]: + def create_kv_cache(self, capacity: int | None = None) -> list[TinyKvCache]: pass def __call__( diff --git a/src/tiny_llm_ref/generate.py b/src/tiny_llm_ref/generate.py index 7a48eb3e..0532dc9b 100644 --- a/src/tiny_llm_ref/generate.py +++ b/src/tiny_llm_ref/generate.py @@ -1,6 +1,6 @@ +import inspect import mlx.core as mx from mlx_lm.tokenizer_utils import TokenizerWrapper -from .kv_cache import * from .qwen3_week1 import Qwen3ModelWeek1 from .qwen3_week2 import Qwen3ModelWeek2 from typing import Callable @@ -69,25 +69,41 @@ def simple_generate_with_kv_cache( tokenizer: TokenizerWrapper, prompt: str, max_tokens: int = 256, + use_bounded_kv_capacity: bool | None = None, ) -> str: _validate_max_tokens(max_tokens) + if use_bounded_kv_capacity is not None and not isinstance( + use_bounded_kv_capacity, bool + ): + raise ValueError("use_bounded_kv_capacity must be a bool or None") if max_tokens == 0: return "" - kv_cache = model.create_kv_cache() def _step(model, y, offset, kv_cache): logits = model(y[None], offset, kv_cache, logits_to_keep=1) logits = logits[:, -1, :] logprobs = logits - mx.logsumexp(logits, keepdims=True) - sampler = lambda x: mx.argmax(x, axis=-1).astype(mx.int32) - y = sampler(logprobs) + y = mx.argmax(logprobs, axis=-1).astype(mx.int32) return y, logprobs.squeeze(0) + kv_cache = None try: # prefill with the prompt tokens = mx.array( tokenizer.encode(prompt, add_special_tokens=False), dtype=mx.int32 ) + if use_bounded_kv_capacity is None: + use_bounded_kv_capacity = bool( + getattr(model, "use_bounded_kv_capacity", False) + ) + capacity = int(tokens.size) + max_tokens if use_bounded_kv_capacity else None + if "capacity" in inspect.signature(model.create_kv_cache).parameters: + kv_cache = model.create_kv_cache(capacity=capacity) + elif capacity is None: + # Week 3's paged-cache factory has no dense-capacity argument. + kv_cache = model.create_kv_cache() + else: + raise ValueError("bounded KV capacity requires a capacity-aware cache") detokenizer = tokenizer.detokenizer detokenizer.reset() offset = 0 diff --git a/src/tiny_llm_ref/kv_cache.py b/src/tiny_llm_ref/kv_cache.py index 298a9233..714bb6cd 100644 --- a/src/tiny_llm_ref/kv_cache.py +++ b/src/tiny_llm_ref/kv_cache.py @@ -244,11 +244,27 @@ def remove_request(self, id: int): class TinyKvFullCache(TinyKvCache): - def __init__(self): + def __init__(self, capacity: int | None = None): + if capacity is not None and ( + not isinstance(capacity, int) or isinstance(capacity, bool) or capacity < 0 + ): + raise ValueError("capacity must be a non-negative integer or None") self.key_values = None self.offset = 0 + self.capacity = capacity + self.logical_copy_bytes = 0 + self.physical_growth_copy_bytes = 0 + self.slice_write_bytes = 0 self.growth_copy_bytes = 0 + @property + def uses_capacity(self) -> bool: + return self.capacity is not None + + def _logical_key_values(self) -> tuple[mx.array, mx.array]: + keys, values = self.key_values + return keys[:, :, : self.offset], values[:, :, : self.offset] + def update_and_fetch( self, key: mx.array, @@ -256,19 +272,49 @@ def update_and_fetch( mask_length: int | None = None, mask: mx.array | str | None = None, ) -> tuple[mx.array, mx.array, int, Optional[mx.array]]: + assert key.shape == value.shape + B, H, S, D = key.shape + if self.uses_capacity: + end = self.offset + S + if end > self.capacity: + raise ValueError( + f"KV cache capacity {self.capacity} exceeded by append ending at {end}" + ) + if self.key_values is None: + assert self.offset == 0 + keys = mx.zeros((B, H, self.capacity, D), dtype=key.dtype) + values = mx.zeros((B, H, self.capacity, D), dtype=value.dtype) + self.key_values = (keys, values) + else: + keys, values = self.key_values + assert keys.shape == (B, H, self.capacity, D) + assert values.shape == (B, H, self.capacity, D) + assert keys.dtype == key.dtype + assert values.dtype == value.dtype + if S: + keys, values = self.key_values + start = mx.array([self.offset]) + keys = mx.slice_update(keys, key, start_indices=start, axes=(2,)) + values = mx.slice_update(values, value, start_indices=start, axes=(2,)) + self.key_values = (keys, values) + self.slice_write_bytes += key.nbytes + value.nbytes + self.offset = end + logical_keys, logical_values = self._logical_key_values() + return logical_keys, logical_values, self.offset, mask + if self.key_values is None: assert self.offset == 0 self.key_values = (key, value) - B, H, S, D = key.shape self.offset = S return key, value, self.offset, mask else: - B, H, S, D = key.shape - assert key.shape == value.shape prev_keys, prev_values = self.key_values assert prev_keys.shape == (B, H, self.offset, D) assert prev_values.shape == (B, H, self.offset, D) - self.growth_copy_bytes += prev_keys.nbytes + prev_values.nbytes + copied_bytes = prev_keys.nbytes + prev_values.nbytes + self.logical_copy_bytes += copied_bytes + self.physical_growth_copy_bytes += copied_bytes + self.growth_copy_bytes += copied_bytes new_keys = mx.concat([prev_keys, key], axis=2) new_values = mx.concat([prev_values, value], axis=2) self.key_values = (new_keys, new_values) @@ -279,6 +325,12 @@ def materialize(self): if self.key_values is not None: mx.eval(*self.key_values) + def reset(self): + """Reset the logical request while retaining bounded physical storage.""" + self.offset = 0 + if not self.uses_capacity: + self.key_values = None + def rewind(self, n: int): if not isinstance(n, int) or isinstance(n, bool) or not 0 <= n <= self.offset: raise ValueError("rewind length must be between zero and the cache length") @@ -286,7 +338,10 @@ def rewind(self, n: int): return self.offset -= n if self.offset == 0: - self.key_values = None + if not self.uses_capacity: + self.key_values = None + return + if self.uses_capacity: return self.key_values = ( self.key_values[0][:, :, : self.offset], diff --git a/src/tiny_llm_ref/qwen3_week2.py b/src/tiny_llm_ref/qwen3_week2.py index eec8163c..11502ed1 100644 --- a/src/tiny_llm_ref/qwen3_week2.py +++ b/src/tiny_llm_ref/qwen3_week2.py @@ -21,6 +21,7 @@ @dataclass(frozen=True) class Week2CheckpointFeatures: + bounded_kv_capacity: bool = False quantized_weights: bool = False fast_rms_norm: bool = False fast_rope: bool = False @@ -33,40 +34,7 @@ class Week2CheckpointFeatures: WEEK2_CHECKPOINT_FEATURES = MappingProxyType( { "kv-cache": Week2CheckpointFeatures(), - "quantized-matvec": Week2CheckpointFeatures(quantized_weights=True), - "rmsnorm": Week2CheckpointFeatures(quantized_weights=True, fast_rms_norm=True), - "rope": Week2CheckpointFeatures( - quantized_weights=True, fast_rms_norm=True, fast_rope=True - ), - "swiglu": Week2CheckpointFeatures( - quantized_weights=True, - fast_rms_norm=True, - fast_rope=True, - fast_swiglu=True, - ), - "simd-matmul": Week2CheckpointFeatures( - quantized_weights=True, - fast_rms_norm=True, - fast_rope=True, - fast_swiglu=True, - simdgroup_matmul=True, - ), - "decode-attention": Week2CheckpointFeatures( - quantized_weights=True, - fast_rms_norm=True, - fast_rope=True, - fast_swiglu=True, - simdgroup_matmul=True, - decode_attention=True, - ), - "split-k": Week2CheckpointFeatures( - quantized_weights=True, - fast_rms_norm=True, - fast_rope=True, - fast_swiglu=True, - simdgroup_matmul=True, - split_k_matmul=True, - ), + "capacity-cache": Week2CheckpointFeatures(bounded_kv_capacity=True), } ) WEEK2_CHECKPOINTS = tuple(WEEK2_CHECKPOINT_FEATURES) @@ -295,16 +263,24 @@ class Qwen3ModelWeek2: def __init__( self, mlx_model: Any, - checkpoint: str = "split-k", + checkpoint: str = "capacity-cache", use_mlx_quantized_linear: bool = False, + use_bounded_kv_capacity: bool | None = None, ): if checkpoint not in WEEK2_CHECKPOINTS: raise ValueError( f"unknown Week 2 checkpoint {checkpoint!r}; " f"choose one of {WEEK2_CHECKPOINTS}" ) + if use_bounded_kv_capacity is not None and not isinstance( + use_bounded_kv_capacity, bool + ): + raise ValueError("use_bounded_kv_capacity must be a bool or None") self.checkpoint = checkpoint features = WEEK2_CHECKPOINT_FEATURES[checkpoint] + if use_bounded_kv_capacity is None: + use_bounded_kv_capacity = features.bounded_kv_capacity + self.use_bounded_kv_capacity = use_bounded_kv_capacity use_quantized_weights = features.quantized_weights use_fast_rms_norm = features.fast_rms_norm use_fast_rope = features.fast_rope @@ -393,10 +369,12 @@ def model_weight(layer: Any) -> mx.array | QuantizedWeights: self.w_lm_head = None self.mlx_model = mlx_model - def create_kv_cache(self) -> list[TinyKvCache]: + def create_kv_cache(self, capacity: int | None = None) -> list[TinyKvCache]: from .kv_cache import TinyKvFullCache - return [TinyKvFullCache() for _ in range(self.num_hidden_layers)] + return [ + TinyKvFullCache(capacity=capacity) for _ in range(self.num_hidden_layers) + ] def __call__( self, diff --git a/tests_refsol/test_dense_kv_capacity.py b/tests_refsol/test_dense_kv_capacity.py new file mode 100644 index 00000000..9b4d3dc0 --- /dev/null +++ b/tests_refsol/test_dense_kv_capacity.py @@ -0,0 +1,254 @@ +"""Focused request-bounded dense KV capacity controls.""" + +import mlx.core as mx +import pytest + +from tiny_llm_ref.generate import simple_generate_with_kv_cache +from tiny_llm_ref.kv_cache import TinyKvFullCache +from tiny_llm_ref.qwen3_week2 import Qwen3ModelWeek2 +from .utils import assert_allclose, tiny_qwen3_mlx_model + + +def _chunk(start: int, length: int, *, dtype=mx.float32): + values = mx.arange(start, start + length * 2, dtype=dtype).reshape(1, 1, length, 2) + return values, values + 100 + + +def test_capacity_cache_writes_only_new_slices_and_hides_the_tail(): + cache = TinyKvFullCache(capacity=5) + key_1, value_1 = _chunk(0, 2) + key_2, value_2 = _chunk(4, 1) + + first_key, first_value, offset, _ = cache.update_and_fetch(key_1, value_1) + assert offset == 2 + assert first_key.shape == first_value.shape == (1, 1, 2, 2) + + cached_key, cached_value, offset, _ = cache.update_and_fetch(key_2, value_2) + mx.eval(cached_key, cached_value) + + assert offset == 3 + assert cached_key.shape == cached_value.shape == (1, 1, 3, 2) + assert cache.key_values[0].shape == cache.key_values[1].shape == (1, 1, 5, 2) + assert_allclose(cached_key, mx.concat([key_1, key_2], axis=2), mx.float32) + assert_allclose(cached_value, mx.concat([value_1, value_2], axis=2), mx.float32) + assert cache.key_values[0][:, :, 3:].tolist() == [[[[0.0, 0.0], [0.0, 0.0]]]] + assert cache.key_values[1][:, :, 3:].tolist() == [[[[0.0, 0.0], [0.0, 0.0]]]] + assert cache.logical_copy_bytes == 0 + assert cache.physical_growth_copy_bytes == 0 + assert cache.growth_copy_bytes == 0 + assert ( + cache.slice_write_bytes + == key_1.nbytes + value_1.nbytes + key_2.nbytes + value_2.nbytes + ) + + +def test_disabled_control_keeps_concatenate_accounting(): + cache = TinyKvFullCache() + key_1, value_1 = _chunk(0, 2) + key_2, value_2 = _chunk(4, 1) + + cache.update_and_fetch(key_1, value_1) + cached_key, cached_value, offset, _ = cache.update_and_fetch(key_2, value_2) + + copied = key_1.nbytes + value_1.nbytes + assert offset == 3 + assert_allclose(cached_key, mx.concat([key_1, key_2], axis=2), mx.float32) + assert_allclose(cached_value, mx.concat([value_1, value_2], axis=2), mx.float32) + assert cache.capacity is None + assert cache.logical_copy_bytes == copied + assert cache.physical_growth_copy_bytes == copied + assert cache.growth_copy_bytes == copied + assert cache.slice_write_bytes == 0 + + +def test_capacity_boundary_is_transactional_and_zero_append_is_a_noop(): + for invalid_capacity in (-1, 1.5, True): + with pytest.raises(ValueError, match="non-negative integer"): + TinyKvFullCache(capacity=invalid_capacity) + + cache = TinyKvFullCache(capacity=2) + key, value = _chunk(0, 2) + cache.update_and_fetch(key, value) + mx.eval(*cache.key_values) + before = tuple(array.tolist() for array in cache.key_values) + before_counters = (cache.offset, cache.slice_write_bytes) + + empty_key, empty_value = _chunk(0, 0) + cached_key, cached_value, offset, _ = cache.update_and_fetch(empty_key, empty_value) + assert cached_key.shape == cached_value.shape == (1, 1, 2, 2) + assert offset == 2 + assert (cache.offset, cache.slice_write_bytes) == before_counters + + extra_key, extra_value = _chunk(4, 1) + with pytest.raises(ValueError, match="capacity 2 exceeded"): + cache.update_and_fetch(extra_key, extra_value) + + mx.eval(*cache.key_values) + assert tuple(array.tolist() for array in cache.key_values) == before + assert (cache.offset, cache.slice_write_bytes) == before_counters + + +def test_capacity_rewind_reuses_storage_and_overwrites_the_logical_suffix(): + cache = TinyKvFullCache(capacity=4) + first_key, first_value = _chunk(0, 3) + cache.update_and_fetch(first_key, first_value) + cache.rewind(2) + replacement_key, replacement_value = _chunk(20, 2) + cached_key, cached_value, offset, _ = cache.update_and_fetch( + replacement_key, replacement_value + ) + mx.eval(cached_key, cached_value) + + assert offset == 3 + assert cache.key_values[0].shape == cache.key_values[1].shape == (1, 1, 4, 2) + assert cache.physical_growth_copy_bytes == 0 + assert_allclose( + cached_key, + mx.concat([first_key[:, :, :1], replacement_key], axis=2), + mx.float32, + ) + assert_allclose( + cached_value, + mx.concat([first_value[:, :, :1], replacement_value], axis=2), + mx.float32, + ) + + cache.rewind(3) + assert cache.offset == 0 + assert cache.key_values is not None + + cache.update_and_fetch(replacement_key[:, :, :1], replacement_value[:, :, :1]) + cache.reset() + assert cache.offset == 0 + assert cache.key_values is not None + + +def test_capacity_mode_matches_the_legacy_model_cache(): + model = Qwen3ModelWeek2(tiny_qwen3_mlx_model(), checkpoint="kv-cache") + bounded = model.create_kv_cache(capacity=3) + legacy = model.create_kv_cache() + + for tokens, offset in ( + (mx.array([[1, 2]], dtype=mx.int32), 0), + (mx.array([[3]], dtype=mx.int32), 2), + ): + bounded_output = model(tokens, offset, bounded) + legacy_output = model(tokens, offset, legacy) + assert_allclose(bounded_output, legacy_output, mx.bfloat16) + + for bounded_layer, legacy_layer in zip(bounded, legacy): + bounded_keys, bounded_values = bounded_layer._logical_key_values() + legacy_keys, legacy_values = legacy_layer.key_values + assert_allclose(bounded_keys, legacy_keys, mx.bfloat16) + assert_allclose(bounded_values, legacy_values, mx.bfloat16) + assert bounded_layer.physical_growth_copy_bytes == 0 + assert legacy_layer.physical_growth_copy_bytes > 0 + + +def test_week2_factory_and_generation_bind_the_request_capacity(capsys): + model = object.__new__(Qwen3ModelWeek2) + model.num_hidden_layers = 2 + caches = model.create_kv_cache(capacity=7) + assert [cache.capacity for cache in caches] == [7, 7] + + class Detokenizer: + last_segment = "" + + def reset(self): + self.last_segment = "" + + def add_token(self, token): + self.last_segment = "" + + def finalize(self): + self.last_segment = "" + + class Tokenizer: + eos_token_id = 99 + detokenizer = Detokenizer() + + def encode(self, prompt, add_special_tokens=False): + assert prompt == "hi" + assert add_special_tokens is False + return [4, 5] + + class Model: + def __init__(self): + self.capacities = [] + self.caches = [] + + def create_kv_cache(self, capacity=None): + self.capacities.append(capacity) + cache = [TinyKvFullCache(capacity=capacity)] + self.caches.append(cache) + return cache + + def __call__(self, inputs, offset, cache, logits_to_keep=1): + assert offset == cache[0].offset + values = mx.ones((1, 1, inputs.shape[1], 1), dtype=mx.float32) + cache[0].update_and_fetch(values, values) + return mx.array([[[0.0, 1.0, 0.0]]], dtype=mx.float32) + + enabled = Model() + simple_generate_with_kv_cache( + enabled, + Tokenizer(), + "hi", + max_tokens=2, + use_bounded_kv_capacity=True, + ) + assert enabled.capacities == [4] + assert enabled.caches[0][0].offset == 3 + assert enabled.caches[0][0].slice_write_bytes > 0 + + disabled = Model() + simple_generate_with_kv_cache(disabled, Tokenizer(), "hi", max_tokens=2) + assert disabled.capacities == [None] + assert disabled.caches[0][0].offset == 3 + assert disabled.caches[0][0].logical_copy_bytes > 0 + enabled_keys, enabled_values = enabled.caches[0][0].key_values + disabled_keys, disabled_values = disabled.caches[0][0].key_values + assert_allclose(enabled_keys[:, :, :3], disabled_keys, mx.float32) + assert_allclose(enabled_values[:, :, :3], disabled_values, mx.float32) + assert capsys.readouterr().out == "" + + +def test_generation_preserves_a_no_argument_cache_factory(capsys): + """Week 3 models still expose create_kv_cache() without dense capacity.""" + + class Detokenizer: + last_segment = "" + + def reset(self): + self.last_segment = "" + + def add_token(self, token): + self.last_segment = "" + + def finalize(self): + self.last_segment = "" + + class Tokenizer: + eos_token_id = 99 + detokenizer = Detokenizer() + + def encode(self, prompt, add_special_tokens=False): + return [4, 5] + + class Model: + def __init__(self): + self.calls = 0 + + def create_kv_cache(self): + self.calls += 1 + return [TinyKvFullCache()] + + def __call__(self, inputs, offset, cache, logits_to_keep=1): + values = mx.ones((1, 1, inputs.shape[1], 1), dtype=mx.float32) + cache[0].update_and_fetch(values, values) + return mx.array([[[0.0, 1.0, 0.0]]], dtype=mx.float32) + + model = Model() + simple_generate_with_kv_cache(model, Tokenizer(), "hi", max_tokens=1) + assert model.calls == 1 + assert capsys.readouterr().out == "" diff --git a/tests_refsol/test_week_2_day_1.py b/tests_refsol/test_week_2_day_1.py index 10c9a79d..d627f644 100644 --- a/tests_refsol/test_week_2_day_1.py +++ b/tests_refsol/test_week_2_day_1.py @@ -14,9 +14,11 @@ def test_task_1_full_cache_appends_chunks(): key_2 = mx.random.normal((1, 2, 2, 4)).astype(mx.bfloat16) value_2 = mx.random.normal((1, 2, 2, 4)).astype(mx.bfloat16) - cached_key, cached_value, offset, mask = cache.update_and_fetch( - key_1, value_1, mask="causal" + first_update = cache.update_and_fetch(key_1, value_1, mask="causal") + assert first_update is not None, ( + "implement the TinyKvFullCache.update_and_fetch learner seam" ) + cached_key, cached_value, offset, mask = first_update assert offset == 3 assert mask == "causal" assert_allclose(cached_key, key_1, mx.bfloat16) @@ -35,6 +37,7 @@ def test_tasks_2_and_3_cached_checkpoint_is_runnable_and_readable(): assert isinstance(model.embedding, Embedding) assert isinstance(layer.input_layernorm, RMSNorm) assert isinstance(layer.self_attn.rope, RoPE) + assert not model.use_bounded_kv_capacity assert not layer.self_attn.use_decode_attention assert not layer.mlp.use_fast_swiglu assert len(model.create_kv_cache()) == model.num_hidden_layers @@ -47,3 +50,102 @@ def test_task_3_rejects_a_position_that_disagrees_with_the_cache(): model = Qwen3ModelWeek2(tiny_qwen3_mlx_model(), checkpoint="kv-cache") with pytest.raises(ValueError): model(mx.array([[1]], dtype=mx.int32), 1, model.create_kv_cache()) + + +# Second checkpoint: request-bounded capacity. + + +def _chunk(start: int, length: int): + key = mx.arange(start, start + length * 2, dtype=mx.float32).reshape( + 1, 1, length, 2 + ) + return key, key + 100 + + +def test_capacity_cache_exposes_only_the_logical_prefix(): + cache = TinyKvFullCache(capacity=5) + key_1, value_1 = _chunk(0, 2) + key_2, value_2 = _chunk(4, 1) + + cache.update_and_fetch(key_1, value_1) + cached_key, cached_value, offset, _ = cache.update_and_fetch(key_2, value_2) + mx.eval(cached_key, cached_value) + + assert offset == 3 + assert cached_key.shape == cached_value.shape == (1, 1, 3, 2) + assert cache.key_values[0].shape == cache.key_values[1].shape == (1, 1, 5, 2) + assert_allclose(cached_key, mx.concat([key_1, key_2], axis=2), mx.float32) + assert_allclose(cached_value, mx.concat([value_1, value_2], axis=2), mx.float32) + assert cache.logical_copy_bytes == 0 + assert cache.physical_growth_copy_bytes == 0 + assert cache.slice_write_bytes == ( + key_1.nbytes + value_1.nbytes + key_2.nbytes + value_2.nbytes + ) + + +def test_capacity_rewind_reuses_storage_and_overflow_is_transactional(): + cache = TinyKvFullCache(capacity=3) + key, value = _chunk(0, 2) + cache.update_and_fetch(key, value) + cache.rewind(1) + replacement_key, replacement_value = _chunk(20, 2) + replacement_update = cache.update_and_fetch(replacement_key, replacement_value) + assert replacement_update is not None, ( + "implement the TinyKvFullCache capacity-cache update_and_fetch learner seam" + ) + cached_key, cached_value, offset, _ = replacement_update + mx.eval(cached_key, cached_value) + + assert offset == 3 + assert cache.physical_growth_copy_bytes == 0 + assert_allclose( + cached_key, + mx.concat([key[:, :, :1], replacement_key], axis=2), + mx.float32, + ) + assert_allclose( + cached_value, + mx.concat([value[:, :, :1], replacement_value], axis=2), + mx.float32, + ) + + before = tuple(array.tolist() for array in cache.key_values) + with pytest.raises(ValueError, match="capacity 3 exceeded"): + extra_key, extra_value = _chunk(30, 1) + cache.update_and_fetch(extra_key, extra_value) + mx.eval(*cache.key_values) + assert tuple(array.tolist() for array in cache.key_values) == before + assert cache.offset == 3 + + movement_counters = ( + cache.logical_copy_bytes, + cache.physical_growth_copy_bytes, + cache.slice_write_bytes, + cache.growth_copy_bytes, + ) + cache.reset() + assert cache.offset == 0 + assert cache.key_values is not None + assert ( + cache.logical_copy_bytes, + cache.physical_growth_copy_bytes, + cache.slice_write_bytes, + cache.growth_copy_bytes, + ) == movement_counters + + +def test_capacity_checkpoint_runs_the_week2_engine(): + model = Qwen3ModelWeek2(tiny_qwen3_mlx_model(), checkpoint="capacity-cache") + assert model.use_bounded_kv_capacity + + bounded = model.create_kv_cache(capacity=3) + readable = Qwen3ModelWeek2( + tiny_qwen3_mlx_model(), checkpoint="kv-cache" + ).create_kv_cache() + inputs = mx.array([[1, 2, 3]], dtype=mx.int32) + actual = model(inputs, 0, bounded) + expected = model(inputs, 0, readable) + + assert_allclose(actual, expected, mx.bfloat16) + assert all(cache.capacity == 3 for cache in bounded) + assert all(cache.slice_write_bytes > 0 for cache in bounded) diff --git a/tests_refsol/test_week_2_day_2.py b/tests_refsol/test_week_2_day_2.py deleted file mode 100644 index 548c8581..00000000 --- a/tests_refsol/test_week_2_day_2.py +++ /dev/null @@ -1,76 +0,0 @@ -"""Week 2 Day 2 benchmark-lifecycle tests.""" - -import mlx.core as mx -import pytest - -from benches import bench - - -class TrackingCache: - def __init__(self): - self.released = False - - def release(self): - self.released = True - - -class FakeModel: - def __init__(self): - self.caches = [TrackingCache(), TrackingCache()] - - def create_kv_cache(self): - return self.caches - - -def test_single_request_benchmark_releases_cache(monkeypatch): - model = FakeModel() - monkeypatch.setattr( - bench, - "sample_next_week2", - lambda model, tokens, offset, cache, logits_to_keep=1: mx.array( - [7], dtype=mx.uint32 - ), - ) - - generated, _, _ = bench.run_one_request_week2( - model, - bench.BenchRequest(prompt_token_ids=[1, 2, 3], max_new_tokens=3), - ) - - assert generated == 3 - assert all(cache.released for cache in model.caches) - - -def test_single_request_benchmark_releases_cache_after_failure(monkeypatch): - model = FakeModel() - - def fail(*args, **kwargs): - raise RuntimeError("model failure") - - monkeypatch.setattr(bench, "sample_next_week2", fail) - with pytest.raises(RuntimeError, match="model failure"): - bench.run_one_request_week2( - model, - bench.BenchRequest(prompt_token_ids=[1], max_new_tokens=1), - ) - - assert all(cache.released for cache in model.caches) - - -def test_single_request_benchmark_selects_serving_prefill_logits(monkeypatch): - model = FakeModel() - observed = [] - - def sample(model, tokens, offset, cache, logits_to_keep=1): - observed.append(logits_to_keep) - return mx.array([7], dtype=mx.uint32) - - monkeypatch.setattr(bench, "sample_next_week2", sample) - - bench.run_one_request_week2( - model, - bench.BenchRequest(prompt_token_ids=[1, 2, 3], max_new_tokens=1), - prefill_logits_to_keep=1, - ) - - assert observed == [1] diff --git a/tests_refsol/test_week_2_day_3.py b/tests_refsol/test_week_2_day_3.py deleted file mode 100644 index 93caabcc..00000000 --- a/tests_refsol/test_week_2_day_3.py +++ /dev/null @@ -1,196 +0,0 @@ -"""Week 2 Day 3 quantized-matvec tests.""" - -import importlib -import inspect - -import mlx.core as mx - -from .tiny_llm_base import ( - Qwen3ModelWeek2, - QuantizedEmbedding, - QuantizedWeights, - RMSNorm, - RoPE, - quantized_matmul, - quantized_matmul_vanilla, - quantized_matvec_custom, -) -from .utils import assert_allclose, tiny_qwen3_mlx_model - -embedding_module = importlib.import_module(QuantizedEmbedding.__module__) -quantize_module = importlib.import_module(quantized_matmul.__module__) - - -def test_task_1_quantized_embedding_dequantizes_selected_rows(): - weight = mx.random.normal((7, 256)).astype(mx.bfloat16) - packed, scales, biases = mx.quantize(weight, group_size=128, bits=4) - embedding = QuantizedEmbedding( - 7, 256, QuantizedWeights(scales, biases, 128, 4, packed) - ) - indices = mx.array([[1, 4]]) - - result = embedding(indices) - expected = mx.dequantize( - packed[indices], scales[indices], biases[indices], group_size=128, bits=4 - ) - assert_allclose(result, expected, mx.bfloat16, atol=2e-2, rtol=2e-2) - - -def test_task_1_quantized_embedding_accepts_sampled_uint32_tokens(): - weight = mx.random.normal((7, 256)).astype(mx.bfloat16) - packed, scales, biases = mx.quantize(weight, group_size=128, bits=4) - embedding = QuantizedEmbedding( - 7, 256, QuantizedWeights(scales, biases, 128, 4, packed) - ) - indices = mx.array([[1, 4]], dtype=mx.uint32) - - result = embedding(indices) - expected = mx.dequantize( - packed[indices], scales[indices], biases[indices], group_size=128, bits=4 - ) - assert_allclose(result, expected, mx.bfloat16, atol=2e-2, rtol=2e-2) - - -def test_week2_quantization_path_uses_course_owned_operators(): - source = ( - inspect.getsource(quantize_module.quantized_matmul) - + inspect.getsource(quantize_module.dequantize_weights) - + inspect.getsource(embedding_module.QuantizedEmbedding.__call__) - ) - assert "mx.quantized_matmul" not in source - assert "mx.dequantize" not in source - - -def test_task_4_model_integrates_packed_weights_before_fast_kernels(): - model = Qwen3ModelWeek2(tiny_qwen3_mlx_model(), checkpoint="quantized-matvec") - layer = model.layers_inner[0] - - assert isinstance(model.embedding, QuantizedEmbedding) - assert isinstance(layer.self_attn.wq, QuantizedWeights) - assert isinstance(layer.mlp.w_gate, QuantizedWeights) - assert isinstance(layer.input_layernorm, RMSNorm) - assert isinstance(layer.self_attn.rope, RoPE) - assert not layer.self_attn.use_decode_attention - assert not layer.mlp.use_fast_swiglu - - -def quantized_matmul_helper( - stream: mx.Stream, - precision: mx.Dtype, - identity_matrix: bool, -): - with mx.stream(stream): - group_size = 128 - if identity_matrix: - input = mx.eye(group_size, dtype=precision) - else: - input = mx.random.normal(shape=(3, group_size), dtype=precision) - weight = mx.random.normal(shape=(5, group_size), dtype=precision) - w_q, scales, biases = mx.quantize(weight, group_size=group_size, bits=4) - user_out = quantized_matmul( - scales=scales, - biases=biases, - group_size=group_size, - bits=4, - a=input, - b=w_q, - transpose_b=True, - ) - ref_out = mx.quantized_matmul( - input, - w_q, - scales, - biases, - group_size=group_size, - bits=4, - transpose=True, - ) - assert user_out.dtype == mx.bfloat16 - if identity_matrix: - assert_allclose(user_out, ref_out, precision) - else: - assert_allclose( - user_out, - ref_out, - precision, - atol=5.0e-1, - message=f"quantized matmul {precision} comparison", - ) - - -def test_task_3_quantized_matmul_simple_bf16_gpu(): - quantized_matmul_helper(mx.gpu, mx.bfloat16, True) - - -def test_task_3_quantized_matmul_complex_bf16_gpu(): - quantized_matmul_helper(mx.gpu, mx.bfloat16, False) - - -def test_task_3_optimized_matvec_matches_vanilla_gpu(): - """The scalar baseline must remain callable for a decode-shaped input.""" - with mx.stream(mx.gpu): - input = mx.random.normal((1, 256)).astype(mx.bfloat16) - weight = mx.random.normal((96, 256)).astype(mx.bfloat16) - packed, scales, biases = mx.quantize(weight, group_size=128, bits=4) - optimized = quantized_matvec_custom( - scales, biases, 128, 4, input, packed, transpose_b=True - ) - vanilla = quantized_matmul_vanilla( - scales, biases, 128, 4, input, packed, transpose_b=True - ) - assert_allclose(optimized, vanilla, mx.bfloat16, atol=0.5, rtol=2e-2) - - -def quantized_matvec_custom_helper(num_rows: int): - with mx.stream(mx.gpu): - group_size = 128 - input = mx.random.normal(shape=(num_rows, group_size), dtype=mx.bfloat16) - weight = mx.random.normal(shape=(64, group_size), dtype=mx.bfloat16) - w_q, scales, biases = mx.quantize(weight, group_size=group_size, bits=4) - user_out = quantized_matvec_custom( - scales=scales, - biases=biases, - group_size=group_size, - bits=4, - a=input, - b=w_q, - transpose_b=True, - ) - ref_out = mx.quantized_matmul( - input, - w_q, - scales, - biases, - group_size=group_size, - bits=4, - transpose=True, - ) - assert_allclose(user_out, ref_out, mx.bfloat16, atol=5.0e-1) - - -def test_task_4_quantized_matvec_custom_m1_gpu(): - quantized_matvec_custom_helper(1) - - -def test_task_4_quantized_matvec_custom_m8_gpu(): - quantized_matvec_custom_helper(8) - - -def test_task_4_quantized_matvec_custom_qwen_shape_gpu(): - with mx.stream(mx.gpu): - input = mx.random.normal((1, 2560)).astype(mx.bfloat16) - weight = mx.random.normal((1024, 2560)).astype(mx.bfloat16) - packed, scales, biases = mx.quantize(weight, group_size=128, bits=4) - result = quantized_matvec_custom( - scales, biases, 128, 4, input, packed, transpose_b=True - ) - expected = mx.quantized_matmul( - input, - packed, - scales, - biases, - group_size=128, - bits=4, - transpose=True, - ) - assert_allclose(result, expected, mx.bfloat16, atol=1.5) diff --git a/tests_refsol/test_week_2_day_4.py b/tests_refsol/test_week_2_day_4.py deleted file mode 100644 index 2d909015..00000000 --- a/tests_refsol/test_week_2_day_4.py +++ /dev/null @@ -1,120 +0,0 @@ -"""Week 2 Day 4 fused-model-kernel tests.""" - -import inspect -import importlib - -import mlx.core as mx -import pytest - -from tiny_llm_ref.basics import silu -from tiny_llm_ref.layer_norm import RMSNorm -from tiny_llm_ref.positional_encoding import RoPE -from .tiny_llm_base import FastRMSNorm, FastRoPE, Qwen3ModelWeek2, swiglu -from .utils import assert_allclose, tiny_qwen3_mlx_model - -implementation_package = FastRMSNorm.__module__.split(".")[0] -week2_kernels = importlib.import_module(FastRMSNorm.__module__) -week1_model = importlib.import_module(f"{implementation_package}.qwen3_week1") -implementation_norm = importlib.import_module(f"{implementation_package}.layer_norm") -implementation_rope = importlib.import_module( - f"{implementation_package}.positional_encoding" -) - - -def test_week2_fast_operators_are_course_owned(): - source = inspect.getsource(week2_kernels) - assert "mx.fast" not in source - - -def test_fast_rms_norm_matches_week1_implementation(): - x = mx.random.normal((2, 3, 16)).astype(mx.bfloat16) - weight = mx.random.normal((16,)).astype(mx.bfloat16) - expected = RMSNorm(16, weight, eps=1e-5)(x) - result = FastRMSNorm(16, weight, eps=1e-5)(x) - assert_allclose(result, expected, mx.bfloat16, atol=2e-2, rtol=2e-2) - - -@pytest.mark.parametrize("offsets", [3, [3, 7]]) -def test_fast_rope_matches_week1_implementation(offsets): - batch_size = 1 if isinstance(offsets, int) else len(offsets) - seq_len = 4 - x = mx.random.normal((batch_size, seq_len, 2, 16)).astype(mx.bfloat16) - fast = FastRoPE(16, 32, base=10000) - readable = RoPE(16, 32, base=10000) - readable_offsets = ( - slice(offsets, offsets + seq_len) - if isinstance(offsets, int) - else [slice(offset, offset + seq_len) for offset in offsets] - ) - result = fast(x, offsets) - expected = readable(x, readable_offsets) - assert_allclose(result, expected, mx.bfloat16, atol=2e-2, rtol=2e-2) - - -def test_swiglu_matches_readable_expression(): - gate = mx.random.normal((2, 4, 16)).astype(mx.bfloat16) - up = mx.random.normal((2, 4, 16)).astype(mx.bfloat16) - assert_allclose(swiglu(gate, up), silu(gate) * up, mx.bfloat16) - - -def test_completed_day4_integrates_all_fused_model_kernels(): - model = Qwen3ModelWeek2(tiny_qwen3_mlx_model(), checkpoint="swiglu") - layer = model.layers_inner[0] - - assert isinstance(model.norm, FastRMSNorm) - assert isinstance(layer.input_layernorm, FastRMSNorm) - assert isinstance(layer.self_attn.rope, FastRoPE) - assert not layer.self_attn.use_decode_attention - assert layer.mlp.use_fast_swiglu - - -def test_rmsnorm_checkpoint_does_not_enable_later_fast_kernels(): - model = Qwen3ModelWeek2(tiny_qwen3_mlx_model(), checkpoint="rmsnorm") - layer = model.layers_inner[0] - - assert isinstance(model.norm, FastRMSNorm) - assert isinstance(layer.input_layernorm, FastRMSNorm) - assert type(layer.self_attn.rope) is implementation_rope.RoPE - assert not layer.self_attn.use_decode_attention - assert not layer.mlp.use_fast_swiglu - - -def test_rope_checkpoint_does_not_enable_swiglu_early(): - model = Qwen3ModelWeek2(tiny_qwen3_mlx_model(), checkpoint="rope") - layer = model.layers_inner[0] - - assert isinstance(model.norm, FastRMSNorm) - assert isinstance(layer.input_layernorm, FastRMSNorm) - assert isinstance(layer.self_attn.rope, FastRoPE) - assert not layer.self_attn.use_decode_attention - assert not layer.mlp.use_fast_swiglu - - -def test_week1_keeps_readable_kernels(): - hidden_size = 16 - num_heads = 2 - num_kv_heads = 1 - head_dim = 8 - attention = week1_model.Qwen3MultiHeadAttention( - hidden_size, - num_heads, - num_kv_heads, - head_dim, - mx.zeros((num_heads * head_dim, hidden_size)), - mx.zeros((num_kv_heads * head_dim, hidden_size)), - mx.zeros((num_kv_heads * head_dim, hidden_size)), - mx.zeros((hidden_size, num_heads * head_dim)), - mx.ones((head_dim,)), - mx.ones((head_dim,)), - ) - assert type(attention.rope) is implementation_rope.RoPE - assert type(attention.q_norm) is implementation_norm.RMSNorm - - mlp = week1_model.Qwen3MLP( - hidden_size, - hidden_size * 2, - mx.zeros((hidden_size * 2, hidden_size)), - mx.zeros((hidden_size * 2, hidden_size)), - mx.zeros((hidden_size, hidden_size * 2)), - ) - assert mlp(mx.ones((1, 1, hidden_size))).shape == (1, 1, hidden_size) diff --git a/tests_refsol/test_week_2_day_5.py b/tests_refsol/test_week_2_day_5.py deleted file mode 100644 index 39a043e2..00000000 --- a/tests_refsol/test_week_2_day_5.py +++ /dev/null @@ -1,282 +0,0 @@ -"""Week 2 Day 5 SIMD-matrix prefill tests.""" - -from pathlib import Path - -import mlx.core as mx -import pytest - -from .tiny_llm_base import ( - Qwen3ModelWeek2, - quantized_matmul, - quantized_matmul_vanilla, -) -from .utils import ( - assert_allclose, - qwen3_0_6b_model_exists, - qwen3_1_7b_model_exists, - qwen3_4b_model_exists, - tiny_qwen3_mlx_model, -) -from mlx_lm import load - - -def test_simd_matmul_checkpoint_is_completed_week2_model(): - model = Qwen3ModelWeek2(tiny_qwen3_mlx_model(), checkpoint="simd-matmul") - layer = model.layers_inner[0] - - assert not model.embedding.use_custom_kernel - assert model.embedding.weight.use_simdgroup_matmul - assert layer.self_attn.wq.use_simdgroup_matmul - assert not layer.self_attn.use_decode_attention - - -def test_task_2_simdgroup_matmul_matches_vanilla_gpu(): - with mx.stream(mx.gpu): - inputs = mx.random.normal((128, 256)).astype(mx.bfloat16) - weight = mx.random.normal((96, 256)).astype(mx.bfloat16) - packed, scales, biases = mx.quantize(weight, group_size=128, bits=4) - tiled = quantized_matmul( - scales, - biases, - 128, - 4, - inputs, - packed, - transpose_b=True, - use_simdgroup=True, - ) - vanilla = quantized_matmul_vanilla( - scales, biases, 128, 4, inputs, packed, transpose_b=True - ) - assert_allclose(tiled, vanilla, mx.bfloat16, atol=1.0, rtol=2e-2) - - -def test_task_2_simdgroup_matmul_uses_accurate_partial_tiles_gpu(): - """A non-multiple-of-eight prefill must not accumulate in bfloat16.""" - with mx.stream(mx.gpu): - inputs = mx.random.normal((10, 256)).astype(mx.bfloat16) - weight = mx.random.normal((96, 256)).astype(mx.bfloat16) - packed, scales, biases = mx.quantize(weight, group_size=128, bits=4) - tiled = quantized_matmul( - scales, - biases, - 128, - 4, - inputs, - packed, - transpose_b=True, - use_simdgroup=True, - ) - vanilla = quantized_matmul_vanilla( - scales, biases, 128, 4, inputs, packed, transpose_b=True - ) - assert_allclose(tiled, vanilla, mx.bfloat16, atol=0.25, rtol=1e-2) - - -@pytest.mark.parametrize( - ("rows", "outputs", "input_dim"), - [(9, 33, 128), (31, 65, 256), (33, 97, 512)], -) -@pytest.mark.parametrize("use_split_k", [False, True]) -def test_task_2_course_owned_tiles_cover_matrix_boundaries_gpu( - rows: int, - outputs: int, - input_dim: int, - use_split_k: bool, -): - mx.random.seed(rows + outputs + input_dim) - with mx.stream(mx.gpu): - inputs = mx.random.normal((rows, input_dim)).astype(mx.bfloat16) - weight = mx.random.normal((outputs, input_dim)).astype(mx.bfloat16) - packed, scales, biases = mx.quantize(weight, group_size=128, bits=4) - tiled = quantized_matmul( - scales, - biases, - 128, - 4, - inputs, - packed, - transpose_b=True, - use_simdgroup=True, - use_split_k=use_split_k, - ) - vanilla = quantized_matmul_vanilla( - scales, biases, 128, 4, inputs, packed, transpose_b=True - ) - assert tiled.shape == (rows, outputs) - assert_allclose(tiled, vanilla, mx.bfloat16, atol=0.25, rtol=1e-2) - - -def test_task_2_course_owned_loader_zero_fills_partial_rows_gpu(): - """A full-width tile with partial rows must still use the safe path.""" - extension = ( - "extensions_ref" - if Path(__file__).parent.name == "tests_refsol" - else "extensions" - ) - header = ( - Path(__file__).parents[1] / f"src/{extension}/src/cooperative_matrix.h" - ).read_text() - source = """ - threadgroup T tile[32]; - using Loader = tiny_llm::CooperativeTileLoader; - Loader::load( - inp, 8, tile, thread_position_in_threadgroup.x, 2, 8); - threadgroup_barrier(mem_flags::mem_threadgroup); - for (uint index = thread_position_in_threadgroup.x; - index < 32; - index += 4) { - out[index] = tile[index]; - } - """ - kernel = mx.fast.metal_kernel( - name="probe_course_loader_partial_rows", - input_names=["inp"], - output_names=["out"], - source=source, - header=header, - ) - - valid = (mx.arange(16).reshape(2, 8) + 1).astype(mx.bfloat16) - sentinel = mx.full((1, 8), 37, dtype=mx.bfloat16) - padding = mx.full((1, 8), -11, dtype=mx.bfloat16) - loader_source = mx.concatenate([valid, sentinel, padding], axis=0) - output = kernel( - inputs=[loader_source], - template=[("T", mx.bfloat16)], - grid=(4, 1, 1), - threadgroup=(4, 1, 1), - output_shapes=[(4, 8)], - output_dtypes=[mx.bfloat16], - )[0] - mx.eval(output) - - assert output.shape == (4, 8) - assert mx.array_equal(output[:2], valid).item() - assert mx.array_equal(output[2:], mx.zeros((2, 8), dtype=mx.bfloat16)).item() - - -@pytest.mark.skipif( - not qwen3_0_6b_model_exists(), reason="Qwen3-0.6B-4bit model not found" -) -def test_utils_qwen3_0_6b(): - pass - - -@pytest.mark.skipif(not qwen3_4b_model_exists(), reason="Qwen3-4B-4bit model not found") -def test_utils_qwen3_4b(): - pass - - -@pytest.mark.skipif( - not qwen3_1_7b_model_exists(), reason="Qwen3-1.7B-4bit model not found" -) -def test_utils_qwen3_1_7b(): - pass - - -def helper_test_task_5(model_name: str, iters: int = 10): - mlx_model, tokenizer = load(model_name) - model = Qwen3ModelWeek2(mlx_model, checkpoint="simd-matmul") - assert not model.embedding.use_custom_kernel - assert model.embedding.weight.use_simdgroup_matmul - assert all(layer.self_attn.wq.use_simdgroup_matmul for layer in model.layers_inner) - for iteration in range(iters): - cache = model.create_kv_cache() - input = (mx.arange(10, dtype=mx.int32) + iteration * 10).reshape( - 1, 10 - ) % tokenizer.vocab_size - user_output = model(input, 0, cache) - ref_output = mlx_model(input) - user_output = user_output - mx.logsumexp(user_output, axis=-1, keepdims=True) - ref_output = ref_output - mx.logsumexp(ref_output, axis=-1, keepdims=True) - assert_allclose( - user_output, ref_output, precision=mx.bfloat16, rtol=0.1, atol=2.5 - ) - - -@pytest.mark.skipif( - not qwen3_0_6b_model_exists(), reason="Qwen3-0.6B-4bit model not found" -) -def test_task_5_qwen3_0_6b(): - helper_test_task_5("Qwen/Qwen3-0.6B-MLX-4bit", 5) - - -@pytest.mark.skipif(not qwen3_4b_model_exists(), reason="Qwen3-4B-4bit model not found") -def test_task_5_qwen3_4b(): - helper_test_task_5("Qwen/Qwen3-4B-MLX-4bit", 1) - - -@pytest.mark.skipif( - not qwen3_1_7b_model_exists(), reason="Qwen3-1.7B-4bit model not found" -) -def test_task_5_qwen3_1_7b(): - helper_test_task_5("Qwen/Qwen3-1.7B-MLX-4bit", 3) - - -def helper_test_task_5_incremental( - model_name: str, - seq_len: int, - iters: int = 1, -): - mlx_model, tokenizer = load(model_name) - model = Qwen3ModelWeek2(mlx_model, checkpoint="simd-matmul") - for _ in range(iters): - inputs = mx.random.randint(0, tokenizer.vocab_size, (1, seq_len)) - ref_outputs = mlx_model(inputs) - decode_cache = model.create_kv_cache() - for offset in range(seq_len): - user_out = model( - inputs=inputs[:, offset : offset + 1], - offset=offset, - cache=decode_cache, - ) - ref_out = ref_outputs[:, offset : offset + 1, :] - user_out = user_out - mx.logsumexp(user_out, axis=-1, keepdims=True) - ref_out = ref_out - mx.logsumexp(ref_out, axis=-1, keepdims=True) - assert_allclose( - user_out, ref_out, precision=mx.bfloat16, rtol=0.1, atol=2.5 - ) - - -@pytest.mark.skipif( - not qwen3_0_6b_model_exists(), reason="Qwen3-0.6B-4bit model not found" -) -def test_task_5_incremental_qwen3_0_6b(): - helper_test_task_5_incremental("Qwen/Qwen3-0.6B-MLX-4bit", seq_len=3) - - -@pytest.mark.skipif(not qwen3_4b_model_exists(), reason="Qwen3-4B-4bit model not found") -def test_task_5_incremental_qwen3_4b(): - helper_test_task_5_incremental( - "Qwen/Qwen3-4B-MLX-4bit", - seq_len=3, - ) - - -@pytest.mark.skipif( - not qwen3_1_7b_model_exists(), reason="Qwen3-1.7B-4bit model not found" -) -def test_task_5_incremental_qwen3_1_7b(): - helper_test_task_5_incremental("Qwen/Qwen3-1.7B-MLX-4bit", seq_len=3) - - -class FakeEmbedding: - def __call__(self, inputs): - return mx.stack([inputs, inputs + 1], axis=-1).astype(mx.bfloat16) - - def as_linear(self, hidden): - return hidden - - -@pytest.mark.parametrize("logits_to_keep,expected_length", [(1, 1), (None, 4)]) -def test_task_5_logits_to_keep_controls_output_length(logits_to_keep, expected_length): - model = Qwen3ModelWeek2.__new__(Qwen3ModelWeek2) - model.num_hidden_layers = 0 - model.embedding = FakeEmbedding() - model.layers_inner = [] - model.norm = lambda hidden: hidden - model.w_lm_head = None - inputs = mx.array([[1, 2, 3, 4]]) - result = model(inputs, 0, [], logits_to_keep=logits_to_keep) - assert result.shape == (1, expected_length, 2) diff --git a/tests_refsol/test_week_2_day_6.py b/tests_refsol/test_week_2_day_6.py deleted file mode 100644 index 1f5bae4a..00000000 --- a/tests_refsol/test_week_2_day_6.py +++ /dev/null @@ -1,171 +0,0 @@ -"""Week 2 Day 6 optional decode-attention tests.""" - -from math import prod - -import mlx.core as mx -import pytest - -from tiny_llm_ref.attention import scaled_dot_product_attention_grouped -from .tiny_llm_base import ( - FastRMSNorm, - FastRoPE, - Qwen3ModelWeek2, - decode_attention_custom, - scaled_dot_product_attention, -) - -from .utils import assert_allclose, tiny_qwen3_mlx_model - - -def test_model_integrates_decode_attention_after_fast_kernels(): - model = Qwen3ModelWeek2(tiny_qwen3_mlx_model(), checkpoint="decode-attention") - layer = model.layers_inner[0] - - assert layer.self_attn.use_decode_attention - assert layer.self_attn.wq.use_simdgroup_matmul - assert isinstance(layer.input_layernorm, FastRMSNorm) - assert isinstance(layer.self_attn.rope, FastRoPE) - assert layer.mlp.use_fast_swiglu - - -def test_model_uses_decode_attention_only_through_measured_context(monkeypatch): - module = __import__(Qwen3ModelWeek2.__module__, fromlist=["unused"]) - readable_attention = module.scaled_dot_product_attention_grouped - calls = [] - - def record_custom(query, key, value, scale, mask): - calls.append(("custom", key.shape[-2])) - return mx.zeros_like(query) - - def record_readable(query, key, value, scale, mask): - calls.append(("readable", key.shape[-2])) - return readable_attention(query, key, value, scale, mask) - - monkeypatch.setattr(module, "decode_attention_custom", record_custom) - monkeypatch.setattr(module, "scaled_dot_product_attention_grouped", record_readable) - - cases = ( - (0, 1, "custom", 1), - (29, 2, "custom", 31), - (126, 2, "custom", 128), - (254, 2, "custom", 256), - (256, 1, "readable", 257), - (253, 3, "readable", 256), - (248, 8, "readable", 256), - ) - for prefix_length, query_length, expected_path, expected_context in cases: - model = Qwen3ModelWeek2(tiny_qwen3_mlx_model(), checkpoint="decode-attention") - attention = model.layers_inner[0].self_attn - cache = model.create_kv_cache()[0] - hidden = model.hidden_size - if prefix_length: - mx.eval( - attention( - mx.zeros((1, prefix_length, hidden), dtype=model.precision), - 0, - cache, - ) - ) - calls.clear() - mx.eval( - attention( - mx.zeros((1, query_length, hidden), dtype=model.precision), - prefix_length, - cache, - ) - ) - assert calls == [(expected_path, expected_context)] - - -def test_model_keeps_explicit_masks_on_readable_path(monkeypatch): - model = Qwen3ModelWeek2(tiny_qwen3_mlx_model(), checkpoint="decode-attention") - attention = model.layers_inner[0].self_attn - cache = model.create_kv_cache()[0] - module = __import__(Qwen3ModelWeek2.__module__, fromlist=["unused"]) - readable_attention = module.scaled_dot_product_attention_grouped - calls = [] - - def reject_custom(*args, **kwargs): - pytest.fail("explicit masks must not use the bounded decode kernel") - - def record_readable(query, key, value, scale, mask): - calls.append(key.shape[-2]) - return readable_attention(query, key, value, scale, mask) - - monkeypatch.setattr(module, "decode_attention_custom", reject_custom) - monkeypatch.setattr(module, "scaled_dot_product_attention_grouped", record_readable) - - hidden = mx.zeros((1, 1, model.hidden_size), dtype=model.precision) - mask = mx.zeros((1, 1, 1, 1), dtype=mx.float32) - mx.eval(attention(hidden, 0, cache, mask)) - - assert calls == [1] - - -def test_fast_attention_matches_grouped_attention(): - query = mx.random.normal((2, 4, 3, 16)).astype(mx.bfloat16) - key = mx.random.normal((2, 2, 5, 16)).astype(mx.bfloat16) - value = mx.random.normal((2, 2, 5, 16)).astype(mx.bfloat16) - mask = mx.broadcast_to( - mx.array([0, 0, 0, 0, -mx.inf], dtype=mx.bfloat16), (2, 1, 3, 5) - ) - scale = 16**-0.5 - result = scaled_dot_product_attention(query, key, value, scale, mask) - expected = scaled_dot_product_attention_grouped(query, key, value, scale, mask) - assert result.shape == query.shape - assert result.dtype == mx.bfloat16 - assert_allclose(result, expected, mx.bfloat16, atol=2e-2, rtol=2e-2) - - -def test_custom_metal_attention_matches_qwen_boundary_sweep(): - head_dim = 128 - query_heads = 4 - shapes = ( - *((1, context) for context in (1, 31, 32, 127, 128, 129, 255, 256)), - *((8, context) for context in (8, 31, 32, 127, 128, 129, 255, 256)), - ) - - def fixture(shape, phase): - values = mx.sin( - mx.arange(prod(shape), dtype=mx.float32) * 0.017 + phase - ).reshape(shape) - return values.astype(mx.bfloat16) - - for query_length, context_length in shapes: - for gqa_ratio in (1, 4): - kv_heads = query_heads // gqa_ratio - query = fixture((1, query_heads, query_length, head_dim), 0.1) - key = fixture((1, kv_heads, context_length, head_dim), 0.7) - value = fixture(key.shape, 1.3) - explicit_mask = mx.where( - mx.arange(context_length) % 5 == 0, - mx.array(-2.0, dtype=mx.float32), - mx.array(0.0, dtype=mx.float32), - ).reshape(1, 1, 1, context_length) - - for mask in ("causal", explicit_mask): - result = decode_attention_custom( - query, key, value, head_dim**-0.5, mask - ) - expected = scaled_dot_product_attention_grouped( - query, key, value, head_dim**-0.5, mask - ) - assert result.shape == query.shape - assert_allclose( - result, - expected, - mx.bfloat16, - atol=3e-2, - rtol=3e-2, - message=( - f"L={query_length}, S={context_length}, " - f"GQA={gqa_ratio}, mask={type(mask).__name__}" - ), - ) - - -def test_custom_metal_attention_rejects_unknown_string_mask(): - query = mx.zeros((1, 4, 1, 128), dtype=mx.bfloat16) - key = mx.zeros((1, 1, 1, 128), dtype=mx.bfloat16) - with pytest.raises(ValueError, match="unsupported attention mask"): - decode_attention_custom(query, key, key, 128**-0.5, "sliding") diff --git a/tests_refsol/test_week_2_day_7.py b/tests_refsol/test_week_2_day_7.py deleted file mode 100644 index 1af5845a..00000000 --- a/tests_refsol/test_week_2_day_7.py +++ /dev/null @@ -1,114 +0,0 @@ -"""Week 2 Day 7 split-K quantized-prefill tests.""" - -import mlx.core as mx - -from .tiny_llm_base import Qwen3ModelWeek2, quantized_matmul -from .utils import assert_allclose, tiny_qwen3_mlx_model - - -def test_split_k_checkpoint_inherits_core_without_optional_attention(): - day_5 = Qwen3ModelWeek2(tiny_qwen3_mlx_model(), checkpoint="simd-matmul") - day_6 = Qwen3ModelWeek2(tiny_qwen3_mlx_model(), checkpoint="decode-attention") - day_7 = Qwen3ModelWeek2(tiny_qwen3_mlx_model(), checkpoint="split-k") - - assert day_5.layers_inner[0].self_attn.wk.use_simdgroup_matmul - assert not day_5.layers_inner[0].self_attn.use_decode_attention - assert day_6.layers_inner[0].self_attn.wk.use_simdgroup_matmul - assert not day_6.layers_inner[0].self_attn.wk.use_split_k_matmul - assert day_6.layers_inner[0].self_attn.use_decode_attention - assert day_7.layers_inner[0].self_attn.wk.use_simdgroup_matmul - assert day_7.layers_inner[0].self_attn.wk.use_split_k_matmul - assert not day_7.layers_inner[0].self_attn.use_decode_attention - - -def test_split_k_matches_unsplit_qwen_4b_kv_shape_gpu(): - """The optimized case uses Qwen3-4B's hidden and KV projection sizes.""" - with mx.stream(mx.gpu): - inputs = mx.random.normal((32, 2560)).astype(mx.bfloat16) - weight = mx.random.normal((1024, 2560)).astype(mx.bfloat16) - packed, scales, biases = mx.quantize(weight, group_size=128, bits=4) - expected = mx.quantized_matmul( - inputs, - packed, - scales, - biases, - transpose=True, - group_size=128, - bits=4, - ) - split = quantized_matmul( - scales, - biases, - 128, - 4, - inputs, - packed, - transpose_b=True, - use_simdgroup=True, - use_split_k=True, - ) - # Each K partition is stored in BF16 before the FP32 reduction. Keep - # the tolerance at one output BF16 bin for the extra rounding step. - assert_allclose(split, expected, mx.bfloat16, atol=1.5, rtol=2e-2) - - -def test_split_k_handles_partial_output_tiles_gpu(): - with mx.stream(mx.gpu): - inputs = mx.random.normal((17, 2560)).astype(mx.bfloat16) - weight = mx.random.normal((1032, 2560)).astype(mx.bfloat16) - packed, scales, biases = mx.quantize(weight, group_size=128, bits=4) - expected = mx.quantized_matmul( - inputs, - packed, - scales, - biases, - transpose=True, - group_size=128, - bits=4, - ) - split = quantized_matmul( - scales, - biases, - 128, - 4, - inputs, - packed, - transpose_b=True, - use_simdgroup=True, - use_split_k=True, - ) - # Each K partition is stored in BF16 before the FP32 reduction. Keep - # the tolerance at one output BF16 bin for the extra rounding step. - assert_allclose(split, expected, mx.bfloat16, atol=1.5, rtol=2e-2) - - -def test_split_k_request_falls_back_for_larger_prefill_gpu(): - with mx.stream(mx.gpu): - # Four row tiles by 80 output tiles already fill the target grid, so - # another K partition would add reduction overhead without useful work. - inputs = mx.random.normal((128, 256)).astype(mx.bfloat16) - weight = mx.random.normal((2560, 256)).astype(mx.bfloat16) - packed, scales, biases = mx.quantize(weight, group_size=128, bits=4) - unsplit = quantized_matmul( - scales, - biases, - 128, - 4, - inputs, - packed, - transpose_b=True, - use_simdgroup=True, - ) - requested = quantized_matmul( - scales, - biases, - 128, - 4, - inputs, - packed, - transpose_b=True, - use_simdgroup=True, - use_split_k=True, - ) - mx.eval(requested, unsplit) - assert mx.array_equal(requested, unsplit).item() From a4cc882743676022f23c8b992f3a493ee287964d Mon Sep 17 00:00:00 2001 From: Sentinel Date: Wed, 23 Sep 2026 13:06:36 -0700 Subject: [PATCH 3/5] Repair Day 1 prerequisites and historical diagrams AI-Assisted: GPT-6 Sol + Sentinel --- book/src/course-roadmap.svg | 8 +++--- book/src/preface.md | 9 ++++--- book/src/week2-01-kv-cache.md | 35 ++++++++++++++++++-------- book/src/week2-kernel-profile.svg | 8 +++--- book/src/week2-performance-summary.svg | 8 +++--- 5 files changed, 42 insertions(+), 26 deletions(-) diff --git a/book/src/course-roadmap.svg b/book/src/course-roadmap.svg index 01034c28..ee2f2329 100644 --- a/book/src/course-roadmap.svg +++ b/book/src/course-roadmap.svg @@ -1,6 +1,6 @@ - Tiny-LLM course roadmap, Week 2 operator off-ramps, and entry paths - The cumulative course implementation proceeds from setup and Week 1 through seven Week 2 days and Week 3 into Week 4. Week 2 Days 1 and 2 establish required state and measurement. Days 3 through 7 preserve course interfaces but allow manual substitution of corresponding MLX operators instead of custom kernels. This per-operator path is different from the full-MLX solution, which bypasses the course stack. Week 4 keeps setup plus Weeks 1 through 3 as course prerequisites, while the deterministic scripted-model tests for Days 1 through 7 can run after setup without a working serving implementation. Day 8 reconnects the scripted sequence to the Week 3 tokenizer and KV cache. + Historical Tiny-LLM roadmap: earlier seven-day Week 2 operator routes + This historical roadmap showed a former seven-day Week 2 course route from setup and Week 1 through Week 3 into Week 4. In that route, Week 2 Days 1 and 2 established state and measurement. Days 3 through 7 preserved course interfaces but allowed manual substitution of corresponding MLX operators instead of custom kernels. That per-operator path differed from the full-MLX solution, which bypassed the course stack. Week 4 kept setup plus Weeks 1 through 3 as course prerequisites, while deterministic scripted-model tests for Days 1 through 7 could run after setup without a working serving implementation. Day 8 reconnected the scripted sequence to the Week 3 tokenizer and KV cache. Consult the current book navigation for today's Week 2 sequence.