diff --git a/.github/workflows/macos.yml b/.github/workflows/macos.yml index 7112c453..bb1d683e 100644 --- a/.github/workflows/macos.yml +++ b/.github/workflows/macos.yml @@ -38,11 +38,12 @@ jobs: - name: Try building extensions run: | pdm run build-ext - pdm run build-ext-ref + 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 - - name: Run printed shipped-day reference CI manifest - run: pdm run python scripts/week2_shipped_day_ci.py + - run: pdm run test-refsol diff --git a/README.md b/README.md index e2b6c669..7be21fcf 100644 --- a/README.md +++ b/README.md @@ -31,7 +31,7 @@ The course follows a four-week learning path: request-bounded `capacity-cache`; Day 2 keeps W4 projection weights packed through the `quantized-matvec` checkpoint. Day 3 adds SIMD matrix prefill at `simd-matmul`. Day 4 adds cumulative RMSNorm, RoPE, and SwiGLU checkpoints. - Tiled dense prefill attention remains a later lesson. + Day 5 adds `tiled-prefill` and runs the completed `selected` model. - **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 @@ -112,6 +112,7 @@ one explicit byte range through the existing loop. | 2.2 | Keep W4 Packed (`quantized-matvec`) | 🚧 | 🚧 | ✅ | 🚧 | | 2.3 | SIMD Matrix Prefill (`simd-matmul`) | 🚧 | 🚧 | ✅ | 🚧 | | 2.4 | Fused Model Primitives (`rmsnorm`, `rope`, `swiglu`) | 🚧 | 🚧 | ✅ | 🚧 | +| 2.5 | Tiled Dense Prefill Attention (`tiled-prefill`, `selected`) | 🚧 | 🚧 | ✅ | 🚧 | | 3.1 | Continuous Batching | ✅ | ✅ | ✅ | 🚧 | | 3.2 | Chunked Prefill | ✅ | ✅ | ✅ | 🚧 | | 3.3 | Paged KV Cache | ✅ | ✅ | ✅ | 🚧 | @@ -129,7 +130,11 @@ 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 test and checkpoint order is separate from the current Day 1 → Day 2 → Day 3 → Day 4 route. +Earlier Week 2 chapter URLs not reused by active Days 1–5 remain available as +[historical material](book/src/week2-02-benchmark-profile.md). The former Day 4 +address now serves the [active fused-primitives lesson](book/src/week2-04-fused-model-kernels.md). +Those historical pages describe an earlier test and checkpoint order, separate +from the current Day 1 → Day 2 → Day 3 → Day 4 → Day 5 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 df38ff0e..e08ac8f0 100644 --- a/benches/bench.py +++ b/benches/bench.py @@ -94,6 +94,8 @@ def parse_args() -> argparse.Namespace: "rmsnorm", "rope", "swiglu", + "tiled-prefill", + "selected", ), help="run one cumulative Week 2 end-to-end checkpoint", ) @@ -166,7 +168,7 @@ def validate_args(args: argparse.Namespace) -> None: args.loader == "week3" or ( args.loader == "week2" - and (args.week2_checkpoint or "swiglu") + and (args.week2_checkpoint or "selected") not in ("kv-cache", "capacity-cache") ) ) diff --git a/benches/bench_course_progression.py b/benches/bench_course_progression.py index cbffc297..d9821ce7 100644 --- a/benches/bench_course_progression.py +++ b/benches/bench_course_progression.py @@ -89,6 +89,20 @@ class Throughput: "week2", ("--week2-checkpoint", "swiglu"), ), + Variant( + "week2-tiled-prefill", + "2.5 Tiled dense prefill attention", + "ref", + "week2", + ("--week2-checkpoint", "tiled-prefill"), + ), + Variant( + "week2-selected", + "2.5 + Run the selected inference engine", + "ref", + "week2", + ("--week2-checkpoint", "selected"), + ), MLX_VARIANT, ) VARIANTS_BY_KEY = { diff --git a/benches/bench_week2_operators.py b/benches/bench_week2_operators.py index cb2494fd..32061b3a 100644 --- a/benches/bench_week2_operators.py +++ b/benches/bench_week2_operators.py @@ -28,6 +28,8 @@ "embedding", "decode-projections", "prefill-projections", + "model-kernels", + "attention", ) @@ -170,7 +172,7 @@ def parse_args() -> argparse.Namespace: parser.add_argument( "--include-split-k", action="store_true", - help="also benchmark the Day 7 split-K path", + help="also benchmark the historical Split-K experiment (not a current checkpoint)", ) parser.add_argument( "--json-output", diff --git a/benches/profile_week2_kernels.py b/benches/profile_week2_kernels.py index fb9a3dba..52af8c4d 100644 --- a/benches/profile_week2_kernels.py +++ b/benches/profile_week2_kernels.py @@ -31,6 +31,8 @@ "rmsnorm:decode:128", "rope:decode:128", "swiglu:decode:128", + "tiled-prefill:prefill:128", + "selected:prefill:128", ) PROMPT_RULE = "synthetic-token-ids" PREFILL_LOGITS = "all" @@ -64,6 +66,8 @@ class KernelImplementation: swiglu: Callable[[mx.array, mx.array], mx.array] decode_attention_max_query: int decode_attention_max_context: int + tiled_attention: Callable[..., mx.array] | None = None + tiled_attention_min_query: int = 9 def load_implementation(name: str) -> KernelImplementation: @@ -85,6 +89,8 @@ def load_implementation(name: str) -> KernelImplementation: swiglu=kernels.swiglu, decode_attention_max_query=getattr(model, "DECODE_ATTENTION_MAX_QUERY", 0), decode_attention_max_context=getattr(model, "DECODE_ATTENTION_MAX_CONTEXT", 0), + tiled_attention=kernels.dense_prefill_attention_mma, + tiled_attention_min_query=kernels.DENSE_PREFILL_MIN_QUERY, ) @@ -370,9 +376,23 @@ def attention(self) -> list[mx.array]: mask = "causal" if self.phase == "prefill" else None for layer in self.model.layers_inner: attention = layer.self_attn - if should_use_decode_attention( + if ( + getattr(attention, "use_tiled_prefill_attention", False) + and self.rows + >= getattr(self.implementation, "tiled_attention_min_query", 9) + and self.dtype == mx.bfloat16 + and attention.head_dim == 128 + ): + output = self.implementation.tiled_attention( + self.query, + self.key, + self.value, + scale=attention.scale, + mask=mask, + ) + elif should_use_decode_attention( self.implementation, - attention.use_decode_attention, + getattr(attention, "use_decode_attention", False), self.rows, self.context, mask, diff --git a/benches/test_bench_course_progression.py b/benches/test_bench_course_progression.py index a65aaeeb..86d9842a 100644 --- a/benches/test_bench_course_progression.py +++ b/benches/test_bench_course_progression.py @@ -1,11 +1,13 @@ import json import re +import sys import unicodedata from pathlib import Path from types import SimpleNamespace import pytest +from benches import bench as benchmark from benches import bench_course_progression as progression @@ -96,54 +98,109 @@ def _assert_required_progression_is_profile_free(chapter: str, day: int) -> None assert optional_blocks, f"Day {day} must label its optional evidence" -def test_week2_live_labels_follow_the_seven_day_book(): +def test_week2_live_labels_follow_the_five_day_book(): labels = {variant.key: variant.label for variant in WEEK2_VARIANTS} assert labels == { "week1": "Week 1 readable", - "week2-kv-cache": "2.1 KV cache", - "week2-quantized-matvec": "2.3 Quantized matvec", - "week2-rmsnorm": "2.4 Fast RMSNorm", - "week2-rope": "2.4 + Fast RoPE", - "week2-swiglu": "2.4 + Fused SwiGLU", - "week2-simd-matmul": "2.5 SIMD matrix prefill", - "week2-decode-attention": "2.6 Optional decode attention", - "week2-split-k": "2.7 Split-K prefill", + "week2-kv-cache": "2.1 Reuse the prefix", + "week2-capacity-cache": "2.1 + Bound KV-cache movement", + "week2-quantized-matvec": "2.2 Keep W4 packed", + "week2-simd-matmul": "2.3 SIMD matrix prefill", + "week2-rmsnorm": "2.4 Compact RMSNorm", + "week2-rope": "2.4 + Compact RoPE", + "week2-swiglu": "2.4 + Compact SwiGLU", + "week2-tiled-prefill": "2.5 Tiled dense prefill attention", + "week2-selected": "2.5 + Run the selected inference engine", "mlx": "MLX", } - readme = (ROOT / "README.md").read_text() summary = (ROOT / "book/src/SUMMARY.md").read_text() - assert "| 2.2 | Benchmarking and Profiling |" in readme - assert "| 2.3 | Quantize the Model |" in readme - assert "| 2.4 | Fused Model Kernels |" in readme - assert "| 2.5 | SIMD-Matrix Prefill |" in readme - assert "| 2.6 (optional) | Workload-Conditioned Operator Lab |" in readme - assert "| 2.7 | Conditional Split-K and Final Decision |" in readme - assert "./week2-02-benchmark-profile.md" in summary - assert "./week2-03-quantize-model.md" in summary - assert "./week2-04-fused-model-kernels.md" in summary - assert "./week2-05-simd-matrix-prefill.md" in summary - assert "./week2-06-operator-lab.md" in summary - assert "./week2-07-split-k-prefill.md" in summary + week2_summary = summary.split("Week 2:", 1)[1].split("Week 3:", 1)[0] + for title in ( + "Day 1: Cache and Measure", + "Day 2: Keep W4 Packed", + "Day 3: SIMD Matrix Prefill", + "Day 4: Fused Model Primitives", + ): + assert title in week2_summary + day5_doc = ROOT / "book/src/week2-05-tiled-prefill-attention.md" + if day5_doc.exists(): + assert "Day 5: Tiled Dense Prefill Attention" in week2_summary + assert len(re.findall(r"\[🚧 Day \d:", week2_summary)) == 4 + day5_doc.exists() chapter_headings = { - "week2-02-benchmark-profile.md": ( - "# 🚧 Week 2 Day 2: Benchmarking and Profiling" - ), - "week2-03-quantize-model.md": "# 🚧 Week 2 Day 3: Quantize the Model", - "week2-04-fused-model-kernels.md": "# 🚧 Week 2 Day 4: Fused Model Kernels", - "week2-05-simd-matrix-prefill.md": "# 🚧 Week 2 Day 5: SIMD-Matrix Prefill", - "week2-06-operator-lab.md": ( - "# 🚧 Week 2 Day 6 (Optional): Workload-Conditioned Operator Lab" - ), - "week2-07-split-k-prefill.md": ( - "# 🚧 Week 2 Day 7: Conditional Split-K and Final Decision" - ), + "week2-01-kv-cache.md": "# 🚧 Week 2 Day 1: Reuse the Prefix, Then Bound the Cache", + "week2-02-quantize-model.md": "# 🚧 Week 2 Day 2: Keep W4 Packed", + "week2-03-simd-matrix-prefill.md": "# 🚧 Week 2 Day 3: SIMD Matrix Prefill", + "week2-04-fused-model-kernels.md": "# 🚧 Week 2 Day 4: Fused Model Primitives", } for filename, expected_heading in chapter_headings.items(): heading = (ROOT / "book/src" / filename).read_text().splitlines()[0] assert heading == expected_heading + if day5_doc.exists(): + assert ( + day5_doc.read_text().splitlines()[0] + == "# 🚧 Week 2 Day 5: Tiled Dense Prefill Attention" + ) + + readme = (ROOT / "README.md").read_text() + week2_overview = readme.split("- **Week 2:", 1)[1].split("- **Week 3:", 1)[0] + assert "decode attention" not in week2_overview.lower() + assert "split-k" not in week2_overview.lower() + for row in ( + "| 2.1 | Cache and Measure", + "| 2.2 | Keep W4 Packed", + "| 2.3 | SIMD Matrix Prefill", + "| 2.4 | Fused Model Primitives", + ): + assert row in readme + if day5_doc.exists(): + assert "| 2.5 | Tiled Dense Prefill Attention" in readme + assert "| 2.6 |" not in readme + assert "| 2.7 |" not in readme + + +def test_week2_progression_checkpoints_parse_through_public_bench(monkeypatch): + checkpoints = tuple( + variant.extra_args[1] + for variant in WEEK2_VARIANTS + if variant.loader == "week2" and variant.extra_args + ) + assert checkpoints == ( + "kv-cache", + "capacity-cache", + "quantized-matvec", + "simd-matmul", + "rmsnorm", + "rope", + "swiglu", + "tiled-prefill", + "selected", + ) + + monkeypatch.setattr( + benchmark, + "load", + lambda *_args, **_kwargs: pytest.fail("model work must not begin"), + ) + for checkpoint in checkpoints: + monkeypatch.setattr( + sys, + "argv", + ["bench", "--loader", "week2", "--week2-checkpoint", checkpoint], + ) + assert benchmark.parse_args().week2_checkpoint == checkpoint + + for retired in ("decode-attention", "split-k"): + monkeypatch.setattr( + sys, + "argv", + ["bench", "--loader", "week2", "--week2-checkpoint", retired], + ) + with pytest.raises(SystemExit): + benchmark.parse_args() + def test_benchmark_refuses_existing_json_before_host_or_model_work( tmp_path, monkeypatch @@ -164,73 +221,64 @@ def test_benchmark_refuses_existing_json_before_host_or_model_work( progression.main() -def test_week2_profile_boundary_is_optional_and_quantization_is_day_3(): - day2 = (ROOT / "book/src/week2-02-benchmark-profile.md").read_text() - day3 = (ROOT / "book/src/week2-03-quantize-model.md").read_text() - appendix = (ROOT / "book/src/week2-advanced-profiling.md").read_text() - - assert "pdm run bench" in day2 - assert "capture is optional" in day2 - assert "pdm run profile-week2-kernels --solution tiny_llm" in day2 - assert "macOS 27" in day2 - assert "pdm run test --week 2 --day 3" in day3 - assert "Week 2 Day 2: Benchmark, Profile, and Quantize" not in day3 - - removed_workflow_tokens = ("capture-week2-shader", "MLX_METAL_DEBUG") - live_week2 = "\n".join( - (ROOT / "book/src" / f"week2-0{day}-{name}.md").read_text() - for day, name in ( - (2, "benchmark-profile"), - (3, "quantize-model"), - (4, "fused-model-kernels"), - (5, "simd-matrix-prefill"), - (6, "operator-lab"), - (7, "split-k-prefill"), - ) - ) - assert not any(token in live_week2 for token in removed_workflow_tokens) +def test_week2_profile_boundary_is_optional_and_quantization_is_day_2(): + day1 = (ROOT / "book/src/week2-01-kv-cache.md").read_text() + day2 = (ROOT / "book/src/week2-02-quantize-model.md").read_text() + optional_capture = (ROOT / "book/src/week2-advanced-profiling.md").read_text() + overview = (ROOT / "book/src/week2-overview.md").read_text() + + assert "--week2-checkpoint kv-cache" in day1 + assert "--week2-checkpoint capacity-cache" in day1 + assert "pdm run test --week 2 --day 1" in day1 + assert "pdm run test --week 2 --day 2" in day2 + assert "--week2-checkpoint quantized-matvec" in day2 + assert "optional capture" in overview + assert "It is never an acceptance gate." in optional_capture scripts = (ROOT / "pyproject.toml").read_text() assert "capture-week2" in scripts assert "reduce-week2-gpudebug" in scripts - assert "never an acceptance gate" in appendix - assert "macOS 27" in appendix -def test_required_week2_progression_uses_portable_attribution_not_local_capture(): +def test_required_week2_progression_uses_portable_checks_not_local_capture(): days = { day: (ROOT / "book/src" / filename).read_text() for day, filename in { - 3: "week2-03-quantize-model.md", + 1: "week2-01-kv-cache.md", + 2: "week2-02-quantize-model.md", + 3: "week2-03-simd-matrix-prefill.md", 4: "week2-04-fused-model-kernels.md", - 5: "week2-05-simd-matrix-prefill.md", - 6: "week2-06-operator-lab.md", + **( + {5: "week2-05-tiled-prefill-attention.md"} + if (ROOT / "book/src/week2-05-tiled-prefill-attention.md").exists() + else {} + ), }.items() } - - expected_cases = { - 3: ("kv-cache:decode:128", "quantized-matvec:decode:128"), - 4: ("quantized-matvec:decode:128", "swiglu:decode:128"), - 5: ("swiglu:prefill:128", "simd-matmul:prefill:128"), - 6: ("simd-matmul:decode:128", "decode-attention:decode:128"), + expected_checkpoints = { + 1: ("kv-cache", "capacity-cache"), + 2: ("quantized-matvec",), + 3: ("simd-matmul",), + 4: ("rmsnorm", "rope", "swiglu"), + 5: ("tiled-prefill", "selected"), } - local_capture_tokens = ( - "Xcode GPU capture", - "Metal System Trace", - ".gputrace", - "gpudebug", - "screenshot", - "GPU duration", - ) - for day, chapter in days.items(): - assert "pdm run profile-week2-kernels --solution tiny_llm" in chapter - assert all(case in chapter for case in expected_cases[day]) - assert not any(token in chapter for token in local_capture_tokens) - - assert "The re-profile then exposed normalization" in days[3] - assert "Re-profiling then placed" in days[4] - assert "Rerun the exact commands from the baseline section" in days[5] - assert "Continue to [Day 7]" in days[6] + assert all(checkpoint in chapter for checkpoint in expected_checkpoints[day]) + assert f"pdm run test --week 2 --day {day}" in chapter + assert "--week2-checkpoint decode-attention" not in chapter + assert "--week2-checkpoint split-k" not in chapter + assert "`gpudebug` output, screenshot, or device-specific counter gates" in days[1] + for day in days.keys() - {1}: + assert not any( + token in days[day] + for token in ( + "Xcode GPU capture", + "Metal System Trace", + ".gputrace", + "gpudebug", + "screenshot", + "GPU duration", + ) + ) @pytest.mark.parametrize( diff --git a/benches/test_profile_week2_kernels.py b/benches/test_profile_week2_kernels.py index 4ebeaf2c..a986fad9 100644 --- a/benches/test_profile_week2_kernels.py +++ b/benches/test_profile_week2_kernels.py @@ -1,9 +1,11 @@ +from dataclasses import replace from types import SimpleNamespace import mlx.core as mx import pytest from benches import profile_week2_kernels as profile +from tests_refsol.utils import tiny_qwen3_mlx_model def test_kernel_group_profile_rotates_every_group_through_each_position(): @@ -124,8 +126,27 @@ def attention(query, _key, _value, *, scale, mask): def test_student_and_reference_profiles_share_the_production_guard(): for name in ("tiny_llm", "tiny_llm_ref"): implementation = profile.load_implementation(name) - assert implementation.decode_attention_max_query == 2 - assert implementation.decode_attention_max_context == 256 + assert implementation.decode_attention_max_query == 0 + assert implementation.decode_attention_max_context == 0 + + +def test_selected_prefill_attribution_uses_the_tiled_operator(): + implementation = profile.load_implementation("tiny_llm_ref") + model = implementation.model_type( + tiny_qwen3_mlx_model(head_dim=128), checkpoint="selected" + ) + calls = [] + + def record(query, _key, _value, *, scale, mask): + calls.append((query.shape[-2], scale, mask)) + return query + + replay = profile.KernelReplay( + replace(implementation, tiled_attention=record), model, "prefill", 9 + ) + outputs = replay.attention() + assert len(outputs) == len(model.layers_inner) + assert calls == [(9, model.layers_inner[0].self_attn.scale, "causal")] def test_decision_requires_exact_source_solution_model_and_workload_identity(): @@ -210,21 +231,32 @@ def test_profile_workload_identity_covers_every_workload_field(monkeypatch): def test_default_attribution_and_model_expose_only_canonical_checkpoints(): cases = [profile.parse_case(value) for value in profile.DEFAULT_CASES] - checkpoints = [case.checkpoint for case in cases] - assert checkpoints.index("simd-matmul") < checkpoints.index("decode-attention") - assert checkpoints[-1] == "split-k" + checkpoints = tuple(case.checkpoint for case in cases) + assert checkpoints == ( + "kv-cache", + "capacity-cache", + "quantized-matvec", + "simd-matmul", + "rmsnorm", + "rope", + "swiglu", + "tiled-prefill", + "selected", + ) + assert {"decode-attention", "split-k"}.isdisjoint(checkpoints) for implementation_name in ("tiny_llm", "tiny_llm_ref"): implementation = profile.load_implementation(implementation_name) assert implementation.checkpoints == ( "kv-cache", + "capacity-cache", "quantized-matvec", + "simd-matmul", "rmsnorm", "rope", "swiglu", - "simd-matmul", - "decode-attention", - "split-k", + "tiled-prefill", + "selected", ) diff --git a/benches/test_week2_gpudebug.py b/benches/test_week2_gpudebug.py index b61f155a..516fa9a0 100644 --- a/benches/test_week2_gpudebug.py +++ b/benches/test_week2_gpudebug.py @@ -1,5 +1,6 @@ import hashlib import json +import sys from argparse import Namespace import pytest @@ -10,6 +11,54 @@ ROOT = gpu.ROOT +def test_capture_parser_exposes_only_current_week2_checkpoints(monkeypatch, capsys): + assert gpu.KNOWN_CHECKPOINTS == ( + "kv-cache", + "capacity-cache", + "quantized-matvec", + "simd-matmul", + "rmsnorm", + "rope", + "swiglu", + "tiled-prefill", + "selected", + ) + + def argv(checkpoint): + return [ + "capture-week2", + "capture", + "--solution", + "tiny_llm", + "--model", + "qwen3-0.6b", + "--checkpoint", + checkpoint, + "--phase", + "decode", + "--tokens", + "1", + "--trace", + "trace.gputrace", + "--metadata", + "capture.json", + "--manifest", + "trace.sha256", + ] + + for checkpoint in gpu.KNOWN_CHECKPOINTS: + monkeypatch.setattr(sys, "argv", argv(checkpoint)) + args = gpu.build_parser().parse_args() + assert args.checkpoint == checkpoint + assert args.handler is gpu.capture + + for retired in ("decode-attention", "split-k"): + monkeypatch.setattr(sys, "argv", argv(retired)) + with pytest.raises(SystemExit): + gpu.build_parser().parse_args() + capsys.readouterr() + + def _identity() -> dict: workload = gpu.workload_record("swiglu", "decode", 128) return { diff --git a/benches/week2_gpudebug.py b/benches/week2_gpudebug.py index 6f6b6181..f06b8ad5 100644 --- a/benches/week2_gpudebug.py +++ b/benches/week2_gpudebug.py @@ -25,6 +25,8 @@ "rmsnorm", "rope", "swiglu", + "tiled-prefill", + "selected", ) diff --git a/book/src/SUMMARY.md b/book/src/SUMMARY.md index a0d49800..39073c2f 100644 --- a/book/src/SUMMARY.md +++ b/book/src/SUMMARY.md @@ -18,6 +18,7 @@ - [🚧 Day 2: Keep W4 Packed](./week2-02-quantize-model.md) - [🚧 Day 3: SIMD Matrix Prefill](./week2-03-simd-matrix-prefill.md) - [🚧 Day 4: Fused Model Primitives](./week2-04-fused-model-kernels.md) + - [🚧 Day 5: Tiled Dense Prefill Attention](./week2-05-tiled-prefill-attention.md) - [Historical Week 2 lesson addresses](./week2-02-benchmark-profile.md) - [Earlier quantization lesson](./week2-03-quantize-model.md) - [Earlier SIMD-prefill lesson](./week2-05-simd-matrix-prefill.md) diff --git a/book/src/appendix-performance.md b/book/src/appendix-performance.md index f00faad2..858515a0 100644 --- a/book/src/appendix-performance.md +++ b/book/src/appendix-performance.md @@ -1,7 +1,8 @@ # 🚧 Appendix: Performance Evidence Ledger > **Historical evidence from an earlier full Week 2 course state.** The -> [current Week 2 route](./week2-overview.md) ships Days 1–4. The +> [current Week 2 route](./week2-overview.md) ships Days 1–5, ending at +> [tiled dense attention and `selected`](./week2-05-tiled-prefill-attention.md). The > checkpoint labels, commands, and measured results below describe the older > source tree; they are not runnable gates or performance results for this > checkout. diff --git a/book/src/glossary.md b/book/src/glossary.md index ad0ba4a3..ca3a86a4 100644 --- a/book/src/glossary.md +++ b/book/src/glossary.md @@ -18,6 +18,7 @@ - [Benchmarking, Profiling, and Decode Roofline](./week2-01-kv-cache.md#benchmark-the-cached-model) - [Packed W4 Quantization](./week2-02-quantize-model.md) - [Fused RMSNorm, RoPE, and SwiGLU](./week2-04-fused-model-kernels.md) +- [Tiled Dense Prefill Attention and Selected Model](./week2-05-tiled-prefill-attention.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) diff --git a/book/src/preface.md b/book/src/preface.md index 8eaf8e62..71fc7239 100644 --- a/book/src/preface.md +++ b/book/src/preface.md @@ -34,7 +34,8 @@ small coding agent. - Week 2: Day 1 caches a request prefix and bounds its dense storage. Day 2 keeps W4 projection weights packed in the cached model. Day 3 adds a SIMD matrix prefill path. Day 4 integrates fused RMSNorm, RoPE, and SwiGLU one - checkpoint at a time. Tiled prefill attention remains a later lesson. + checkpoint at a time. Day 5 adds tiled dense prefill attention and runs the + completed `selected` single-request model. - 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. @@ -46,10 +47,11 @@ all earlier exercises. These are not the same path. The current route runs from Week 1 through [Week 2 Day 1](./week2-01-kv-cache.md), [Day 2](./week2-02-quantize-model.md), [Day 3](./week2-03-simd-matrix-prefill.md), -then [Day 4](./week2-04-fused-model-kernels.md): `kv-cache`, `capacity-cache`, -`quantized-matvec`, `simd-matmul`, `rmsnorm`, `rope`, then `swiglu`. -The tiled-prefill Day 5 checkpoint is planned, so this four-day checkout does -not offer a completed Week 2 → Week 3 learner path. The +[Day 4](./week2-04-fused-model-kernels.md), then +[Day 5](./week2-05-tiled-prefill-attention.md): `kv-cache`, `capacity-cache`, +`quantized-matvec`, `simd-matmul`, `rmsnorm`, `rope`, `swiglu`, +`tiled-prefill`, then `selected`. The five-day route supplies the dense-model +interface that Week 3 extends with paging and batching. 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. @@ -65,6 +67,7 @@ operator shortcut. | Build the current packed-weight path | Complete Day 1, then Week 2 Day 2 | Keep the capacity cache and implement the packed operator and model wiring. | | Build the current SIMD prefill path | Complete Days 1 and 2, then Week 2 Day 3 | Keep the bounded cache and packed weights; implement the SIMD matrix tile and model wiring. | | Build the current fused-primitives path | Complete Days 1–3, then Week 2 Day 4 | Keep the cached packed model and integrate RMSNorm, RoPE, and SwiGLU in order. | +| Build the current tiled-attention path | Complete Days 1–4, then Week 2 Day 5 | Keep the cached packed model, implement supported BF16/D128 prefill attention, retain readable decode, and run `selected`. | | Study an older Week 2 experiment | Read its historical page | Its former day numbers and commands are not current gates. | | 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. | @@ -77,6 +80,8 @@ The cumulative dependencies are deliberate: cache checkpoints. Day 2 keeps W4 weights packed through the live model. Day 3 adds a SIMD matrix path for prefill while keeping that model state. Day 4 adds three cumulative fused primitives around those projections. + Day 5 tiles dense attention for eligible prompt rows and names the complete + single-request product `selected`. - **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 @@ -87,7 +92,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? Its interfaces are; every future +> **Is Week 2 required for Week 3? Its interfaces are; every custom > 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 @@ -98,15 +103,17 @@ The cumulative dependencies are deliberate: > the entire Week 2 implementation would require a supplied hybrid starting > checkpoint; that checkpoint does not exist today. -### Later Week 2 operator off-ramps +### Week 2 operator off-ramps Day 1 requires cache state and matched measurement; it has no replaceable custom kernel. Days 2 and 3 have an optional `mx.quantized_matmul` substitution at the projection operator boundary; both still need the course-owned cache and model wiring. Day 4 can substitute the equivalent MLX operator or equation at one RMSNorm, RoPE, or SwiGLU boundary while preserving the other -course-owned paths. The earlier full-course book retains additional -mechanisms and old addresses in the +course-owned paths. Day 5 can substitute equivalent MLX attention at the +dense prefill boundary while keeping the cache and shape/mask adapter. The +earlier full-course book retains additional mechanisms at former URLs not +reused for active Days 1–5 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. @@ -125,7 +132,7 @@ pdm run main --solution mlx ``` The Day 1 reference tests do not require `pdm run build-ext-ref`. Build the -reference extension for Days 2–4's native tests, alongside +reference extension for Days 2–5's native tests, alongside the learner extension; neither reference solution fills your learner TODOs. `--solution tiny_llm_ref` runs the supplied implementation end to end. `--solution mlx` @@ -160,7 +167,8 @@ keep the required path at 0.6B. On a 16–24 GB Mac, use 0.6B for the required w Week 2 Day 1 retains that dense BF16 model. In this checkout, Day 2 keeps projection weights packed for `quantized-matvec`, and Day 3 reuses those weights for `simd-matmul` prefill. Day 4 keeps the packed path while adding -RMSNorm, RoPE, and SwiGLU; the Week 3 and 4 paths expect the packed interface. +RMSNorm, RoPE, and SwiGLU. Day 5 keeps that model and adds tiled attention for +supported prefill; the Week 3 and 4 paths expect the packed 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 @@ -172,8 +180,8 @@ pressure. [M5 MacBook Air](https://support.apple.com/en-us/126320) specifications. Higher-memory configurations are outside this table. [^week2-dense]: These conservative Week 2 model-size choices cover Day 1's - dense BF16 checkpoint. Days 2–4 retain the cache and keep projection - weights packed. Use 0.6B for the required Day 3 and Day 4 comparisons. + dense BF16 checkpoint. Days 2–5 retain the cache and keep projection + weights packed. Use 0.6B for the required Day 3–5 comparisons. The optional packed 4B comparison is described in [Day 2](./week2-02-quantize-model.md); its matched Day 1 control still needs enough memory for the dense `capacity-cache` model. diff --git a/book/src/sitemap.txt b/book/src/sitemap.txt index 9316c4c8..18a46e46 100644 --- a/book/src/sitemap.txt +++ b/book/src/sitemap.txt @@ -20,6 +20,7 @@ https://skyzh.github.io/tiny-llm/week2-03-simd-matrix-prefill https://skyzh.github.io/tiny-llm/week2-04-fused-model-kernels https://skyzh.github.io/tiny-llm/week2-05-decode-attention https://skyzh.github.io/tiny-llm/week2-05-simd-matrix-prefill +https://skyzh.github.io/tiny-llm/week2-05-tiled-prefill-attention https://skyzh.github.io/tiny-llm/week2-06-operator-lab https://skyzh.github.io/tiny-llm/week2-06-simd-matrix-prefill https://skyzh.github.io/tiny-llm/week2-07-split-k-prefill diff --git a/book/src/sitemap.xml b/book/src/sitemap.xml index 9ac3834f..17e5f6a0 100644 --- a/book/src/sitemap.xml +++ b/book/src/sitemap.xml @@ -66,6 +66,9 @@ https://skyzh.github.io/tiny-llm/week2-05-simd-matrix-prefill + + https://skyzh.github.io/tiny-llm/week2-05-tiled-prefill-attention + https://skyzh.github.io/tiny-llm/week2-06-operator-lab diff --git a/book/src/week2-02-benchmark-profile.md b/book/src/week2-02-benchmark-profile.md index f6a28384..dc2987da 100644 --- a/book/src/week2-02-benchmark-profile.md +++ b/book/src/week2-02-benchmark-profile.md @@ -4,8 +4,10 @@ > explanation and its original links. The current Week 2 learner route > ships [Day 1: Cache and Measure](./week2-01-kv-cache.md), > [Day 2: Keep W4 Packed](./week2-02-quantize-model.md), -> [Day 3: SIMD Matrix Prefill](./week2-03-simd-matrix-prefill.md), and -> [Day 4: Fused Model Primitives](./week2-04-fused-model-kernels.md). The commands below +> [Day 3: SIMD Matrix Prefill](./week2-03-simd-matrix-prefill.md), +> [Day 4: Fused Model Primitives](./week2-04-fused-model-kernels.md), then +> [Day 5: Tiled Dense Prefill Attention](./week2-05-tiled-prefill-attention.md). +> The commands below > remain part of the earlier course order. > 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. diff --git a/book/src/week2-02-quantize-model.md b/book/src/week2-02-quantize-model.md index d8dfb2e5..ea27188f 100644 --- a/book/src/week2-02-quantize-model.md +++ b/book/src/week2-02-quantize-model.md @@ -712,8 +712,9 @@ your custom primitive. Decode-shaped work must route through `quantized_linear` → `quantized_matvec_custom` → the extension primitive → the Metal matvec. Matrix-shaped work must route through `quantized_linear` → `quantized_matmul` → the extension primitive → its Metal matrix schedule. The -supplied tests validate packed model state and the direct operators. Use the -live model command to verify that those pieces compose. +supplied tests compare public `quantized-matvec` checkpoint outputs and direct +operator results with MLX controls; they do not assert private routing. Use +the live model command to verify that those pieces compose. Measure the cumulative model and the real projection shapes: diff --git a/book/src/week2-03-quantize-model.md b/book/src/week2-03-quantize-model.md index 2f1a5618..8e274c51 100644 --- a/book/src/week2-03-quantize-model.md +++ b/book/src/week2-03-quantize-model.md @@ -4,8 +4,10 @@ > explanation and its original links. The current Week 2 learner route > ships [Day 1: Cache and Measure](./week2-01-kv-cache.md), the revised > [Day 2: Keep W4 Packed](./week2-02-quantize-model.md), -> [Day 3: SIMD Matrix Prefill](./week2-03-simd-matrix-prefill.md), and -> [Day 4: Fused Model Primitives](./week2-04-fused-model-kernels.md). The old Day 3 body +> [Day 3: SIMD Matrix Prefill](./week2-03-simd-matrix-prefill.md), +> [Day 4: Fused Model Primitives](./week2-04-fused-model-kernels.md), then +> [Day 5: Tiled Dense Prefill Attention](./week2-05-tiled-prefill-attention.md). +> The old Day 3 body > below remains for its original address and historical context. > 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. diff --git a/book/src/week2-03-simd-matrix-prefill.md b/book/src/week2-03-simd-matrix-prefill.md index a7c3d7ae..e1018558 100644 --- a/book/src/week2-03-simd-matrix-prefill.md +++ b/book/src/week2-03-simd-matrix-prefill.md @@ -202,7 +202,8 @@ revise the schedule. The [performance appendix](./appendix-performance.md) keeps older hardware observations separate from this exercise. Continue with [Day 4: Fused Model Primitives](./week2-04-fused-model-kernels.md), keeping this `simd-matmul` result as its pre-edit control. Tiled dense prefill attention -remains a future checkpoint; the older seven-day pages remain +now follows on [Day 5](./week2-05-tiled-prefill-attention.md). Former seven-day +pages at URLs not reused by active Days 1–5 remain [historical material](./week2-02-benchmark-profile.md). {{#include copyright.md}} diff --git a/book/src/week2-04-fused-model-kernels.md b/book/src/week2-04-fused-model-kernels.md index 9171b423..09deae54 100644 --- a/book/src/week2-04-fused-model-kernels.md +++ b/book/src/week2-04-fused-model-kernels.md @@ -292,7 +292,8 @@ at 4B rather than mixing model sizes. If you continue without writing one of these kernels, keep its public course interface and substitute the equivalent MLX operator or equation only at that boundary. The cached Week 2 model and other course-owned operators still run; -`--solution mlx` runs a separate complete model. Day 5's tiled dense prefill -attention remains a future learner checkpoint in this checkout. +`--solution mlx` runs a separate complete model. Continue to +[Day 5: Tiled Dense Prefill Attention](./week2-05-tiled-prefill-attention.md) +with the `swiglu` measurement as its pre-edit control. {{#include copyright.md}} diff --git a/book/src/week2-05-simd-matrix-prefill.md b/book/src/week2-05-simd-matrix-prefill.md index 942653a7..001e7932 100644 --- a/book/src/week2-05-simd-matrix-prefill.md +++ b/book/src/week2-05-simd-matrix-prefill.md @@ -4,8 +4,9 @@ > explanation and its original links. The current Week 2 learner route > ships [Day 1: Cache and Measure](./week2-01-kv-cache.md), > [Day 2: Keep W4 Packed](./week2-02-quantize-model.md), -> [Day 3: SIMD Matrix Prefill](./week2-03-simd-matrix-prefill.md), and -> [Day 4: Fused Model Primitives](./week2-04-fused-model-kernels.md). +> [Day 3: SIMD Matrix Prefill](./week2-03-simd-matrix-prefill.md), +> [Day 4: Fused Model Primitives](./week2-04-fused-model-kernels.md), then +> [Day 5: Tiled Dense Prefill Attention](./week2-05-tiled-prefill-attention.md). > 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. diff --git a/book/src/week2-05-tiled-prefill-attention.md b/book/src/week2-05-tiled-prefill-attention.md new file mode 100644 index 00000000..f6dce299 --- /dev/null +++ b/book/src/week2-05-tiled-prefill-attention.md @@ -0,0 +1,237 @@ +# 🚧 Week 2 Day 5: Tiled Dense Prefill Attention + +Complete [Day 4: Fused Model Primitives](./week2-04-fused-model-kernels.md) +first. You have a cached, packed-W4 Qwen3 model whose `swiglu` checkpoint +already includes SIMD matrix prefill, RMSNorm, RoPE, and SwiGLU. Day 5 changes +the dense attention computation for a prompt with at least nine query tokens. +One-token decode keeps readable grouped attention. + +Keep the same cached 0.6B model and request while you work. Save the Day 4 +`swiglu` product result before editing attention; later compare it with +`tiled-prefill` and the completed `selected` model under that same workload. +The progression runner uses fresh processes and balanced order. Its `--offline` +mode needs the model files cached already: + +```bash +pdm run build-ext +pdm run bench-week2-progression --offline --solution tiny_llm --suite week2 --repeats 2 \ + --variant week2-swiglu --variant mlx \ + --model qwen3-0.6b --input-len 128 --output-len 129 --warmup 2 \ + --prefill-logits last --json-output week2-day5-control.json +``` + +The `mlx` row is a separate full-model baseline. Your course model and its +cache remain the subject of the Day 5 work. Build the supplied reference +extension and run its completed Day 5 checks to see the target behavior; the +learner extension retains its own TODOs: + +```bash +pdm run build-ext-ref +pdm run test-refsol --week 2 --day 5 +``` + +## Task 1: Make Tiled Attention Match the Readable Result + +The supplied `dense_prefill_attention_mma` adapter in +`src/tiny_llm/week2_kernels.py` prepares the model-facing arrays and calls +the learner-owned private `tiny_llm_ext::_dense_attention_prefill_mma` binding. +Implement that binding and `Week2DensePrefillMMA::eval_gpu` in +`src/extensions/src/week2_kernels.cpp`, plus +`week2_dense_prefill_mma_bf16_d128` in +`src/extensions/src/week2_kernels.metal`. Keep the supplied Python adapter and +header interface. Replace the `Week2DensePrefillMMA::eval_cpu` starter TODO in +`src/extensions/src/week2_kernels.cpp` with a GPU-only error. The adapter +validates the shape, prepares contiguous inputs +and an optional mask, invokes the native primitive, and restores the +model-facing shape. + +The operator accepts grouped-query attention in this layout: + +```text +Q: B, Hq, L, 128 +K, V: B, Hkv, S, 128 +output: B, Hq, L, 128 +``` + +Here `L` is the new prompt length and `S` includes the cached source prefix. +Require BF16 Q/K/V, equal K/V shapes, matching batch and head dimensions, +and `Hq % Hkv == 0`. For query head `h`, use KV head +`floor(h / (Hq / Hkv))`; copying K/V into `Hq` heads would spend the bandwidth +the grouped layout is meant to save. The model selects the tiled path for a +supported BF16/D128 prefill with `L >= 9`. The direct operator rejects shorter +queries; model attention routes them to the readable fallback. + +Start with the supplied causal GQA witness. Its 33 query rows and 47 source +positions cross both tile boundaries, so a result that only handles complete +tiles cannot pass: + +```bash +pdm run build-ext +pdm run test --week 2 --day 5 -- -k task_1 +``` + +The direct attention comparison checks BF16 output against the readable +result with `atol=0.02, rtol=0.02`; passing it does not measure speed. + +The intended computation is the same scaled dot-product attention as the +readable control. Process **BQ32** query rows and **BK16** source positions at a +time with four 32-lane SIMD groups. Cooperatively load a 32×128 Q tile and a +16×128 K tile, accumulate score fragments, then reuse the K/V storage for the +matching V tile. The source layout uses 128 threads and 12,928 bytes of Q and +K/V threadgroup storage; those sizes describe the implementation, not its +speed. Two padding elements per shared-tile row account for the storage above +the raw Q plus K/V payload. Guard both partial tails before any load or output +write. + +For each query row, keep an FP32 running maximum `m`, exponential sum `l`, and +weighted value accumulator `a`. If the next source tile contributes maximum +`m_b`, sum `l_b`, and accumulator `a_b`, merge them as follows: + +$$ +m' = \max(m,m_b), \qquad +l' = e^{m-m'}l + e^{m_b-m'}l_b, +$$ + +$$ +a' = e^{m-m'}a + e^{m_b-m'}a_b, +\qquad \mathrm{output} = a'/l'. +$$ + +In words, choose the larger of the old and tile maxima. Multiply the old sum +and accumulator by the exponential of old maximum minus new maximum; multiply +the tile sum and accumulator by the exponential of tile maximum minus new +maximum. Add each pair, then divide the new accumulator by the new sum. +Rescale the old numerator **and** denominator when a later tile raises the +maximum. Padded score cells contribute no probability mass. Return BF16 only +after the final FP32 normalization. This avoids an allocated `L × S` score +workspace, but it still reads dense K/V; a caller-supplied additive mask can +itself have `L × S` entries. + +## Task 2: Preserve Masks and the Readable Fallback + +Accept no mask, the string `"causal"`, or a broadcastable additive array mask. +Causality is aligned to the end of the cached prefix: query row `i` may see +source positions through `S - L + i`. For `L = 3` and `S = 7`, the first new +query can see positions 0–4. Applying a triangle as if the source length were +three would erase useful cached context. Broadcast an explicit mask over batch +and query-head dimensions, then apply it before the online-softmax update. + +A fully masked row has no probability mass. Skip tiles with no valid scores +and make the row's output finite and all zero; do not divide by zero or +propagate `NaN`. The focused mask witness uses +nine BF16 query rows, 17 source positions, and an all-negative-infinity +additive mask: + +```bash +pdm run build-ext +pdm run test --week 2 --day 5 -- -k task_2 +``` + +In `Qwen3MultiHeadAttention.__call__`, keep readable grouped attention for +one-token decode, `L <= 8`, and any model shape outside the tiled BF16/D128 +boundary. The readable path must receive the same cache update, scale, mask, +query-to-KV-head mapping, and output dtype. Keep separate diagnostic counts +for `tiled_prefill`, `tiled_prefill_fallback`, and `readable`: the fallback count +records that the tiled feature was enabled but ineligible, while `readable` +records the actual computation. These are reference diagnostics to inspect; +the supplied learner tests check public output rather than internal counter +values. A green long-query operator comparison alone does not verify decode. + +If you choose the operator off-ramp, keep this model and cache interface and +substitute the equivalent MLX attention operator or readable grouped equation +at this boundary. `--solution mlx` instead runs a separate complete model. + +## Task 3: Run `tiled-prefill` in the Model + +Wire the eligible attention call through `Qwen3ModelWeek2` and its transformer +layers. The `tiled-prefill` checkpoint retains every Day 4 operator, bounded +cache, and packed projection; it changes only supported dense prefill +attention. Check its focused model behavior and run the live model: + +```bash +pdm run build-ext +pdm run test --week 2 --day 5 -- -k task_3 +pdm run main --solution tiny_llm --loader week2 \ + --week2-checkpoint tiled-prefill --model qwen3-0.6b --max-tokens 16 +``` + +The supplied three-token model witness guards the readable branch. The +ten-token BF16/D128 model witness is eligible for tiled attention; the 33×47 +operator witness also crosses both tile tails. Inspect the reference fallback +and tiled dispatch counts separately. The matched 128-token product workload exercises +the tiled branch in the real model; a matching output alone cannot tell you +which kernel ran. + +## Finish the Cumulative `selected` Model + +`selected` is the named final Week 2 checkpoint, not an additional Metal +kernel. It keeps the `tiled-prefill` feature set: request-bounded KV capacity, +register-cached RMSNorm, the tiled prefill selector, packed W4 projections, +SIMD matrix prefill, RoPE, and SwiGLU. Each of the capacity, RMSNorm, and tiled +attention controls can also be switched off independently for a matched +counterfactual. The short attention fallback and wider-row RMSNorm fallback +remain part of this model. + +First verify the supplied selected-model controls, then run the same completed +course model through the public CLI: + +```bash +pdm run test --week 2 --day 5 -- -k selected +pdm run main --solution tiny_llm --loader week2 \ + --week2-checkpoint selected --model qwen3-0.6b --max-tokens 16 +``` + +The test checks the cumulative model and its independently disabled controls; +the model comparison uses `atol=0.75, rtol=0.05`. Its observable assertions +do not prove a particular private dispatch. The CLI run shows that the named +final checkpoint can answer a request. Then run the complete Day 5 learner gate: + +```bash +pdm run test --week 2 --day 5 +``` + +Once complete, `selected` is the +default Week 2 checkpoint, but name it explicitly in measurements so the +comparison remains legible. + +## Measure the Matched Product + +Compare the saved pre-edit `swiglu` row with a new `swiglu` row first. Then +compare `swiglu`, `tiled-prefill`, `selected`, and the full MLX model under one +cached 0.6B request, device, prefill-logit mode, and warmup count: + +```bash +pdm run bench-week2-progression --offline --solution tiny_llm --suite week2 --repeats 2 \ + --variant week2-swiglu --variant week2-tiled-prefill \ + --variant week2-selected --variant mlx \ + --model qwen3-0.6b --input-len 128 --output-len 129 --warmup 2 \ + --prefill-logits last --json-output week2-day5-product.json +``` + +A 128-token prompt crosses the `L >= 9` selector boundary. Report the prefill +and complete-request results separately; the `selected` and `tiled-prefill` +rows name the same default feature set, so do not count any difference between +them as a second mechanism gain. Inspect which kernel category dominates under +the same 0.6B model and 128-token prefill shape: + +```bash +pdm run profile-week2-kernels --solution tiny_llm --model qwen3-0.6b \ + --case swiglu:prefill:128 --case tiled-prefill:prefill:128 \ + --case selected:prefill:128 --warmup 4 --iterations 12 \ + --json-output week2-day5-attribution.json +``` + +This attribution replays kernel groups on synthetic token IDs; it does not +replace the complete-request comparison. Record whether the tiled path ran +and what later category dominates. An operator result or a historical 4B number +cannot establish a speedup for this 0.6B checkout. The +[performance appendix](./appendix-performance.md) keeps earlier measurements +and unavailable rows as historical evidence. A matched 4B follow-up is +optional after caching that model; repeat every compared row at 4B rather than +mixing model sizes. + +Week 2 now supplies a tested single-request model boundary for +[Week 3](./week3-overview.md), where paging and batching add new state and +scheduling work. It does not turn one request into a production serving policy. + +{{#include copyright.md}} diff --git a/book/src/week2-06-operator-lab.md b/book/src/week2-06-operator-lab.md index 308a508f..c7236f9e 100644 --- a/book/src/week2-06-operator-lab.md +++ b/book/src/week2-06-operator-lab.md @@ -4,8 +4,9 @@ > explanation and its original links. The current Week 2 learner route > ships [Day 1: Cache and Measure](./week2-01-kv-cache.md), > [Day 2: Keep W4 Packed](./week2-02-quantize-model.md), -> [Day 3: SIMD Matrix Prefill](./week2-03-simd-matrix-prefill.md), and -> [Day 4: Fused Model Primitives](./week2-04-fused-model-kernels.md). +> [Day 3: SIMD Matrix Prefill](./week2-03-simd-matrix-prefill.md), +> [Day 4: Fused Model Primitives](./week2-04-fused-model-kernels.md), then +> [Day 5: Tiled Dense Prefill Attention](./week2-05-tiled-prefill-attention.md). > 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. diff --git a/book/src/week2-07-split-k-prefill.md b/book/src/week2-07-split-k-prefill.md index 1fd71b65..be7b4d17 100644 --- a/book/src/week2-07-split-k-prefill.md +++ b/book/src/week2-07-split-k-prefill.md @@ -4,8 +4,9 @@ > explanation and its original links. The current Week 2 learner route > ships [Day 1: Cache and Measure](./week2-01-kv-cache.md), > [Day 2: Keep W4 Packed](./week2-02-quantize-model.md), -> [Day 3: SIMD Matrix Prefill](./week2-03-simd-matrix-prefill.md), and -> [Day 4: Fused Model Primitives](./week2-04-fused-model-kernels.md). +> [Day 3: SIMD Matrix Prefill](./week2-03-simd-matrix-prefill.md), +> [Day 4: Fused Model Primitives](./week2-04-fused-model-kernels.md), then +> [Day 5: Tiled Dense Prefill Attention](./week2-05-tiled-prefill-attention.md). > 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. diff --git a/book/src/week2-advanced-profiling.md b/book/src/week2-advanced-profiling.md index 0996c84e..a648060d 100644 --- a/book/src/week2-advanced-profiling.md +++ b/book/src/week2-advanced-profiling.md @@ -4,8 +4,9 @@ > explanation and its original links. The current Week 2 learner route > ships [Day 1: Cache and Measure](./week2-01-kv-cache.md), > [Day 2: Keep W4 Packed](./week2-02-quantize-model.md), -> [Day 3: SIMD Matrix Prefill](./week2-03-simd-matrix-prefill.md), and -> [Day 4: Fused Model Primitives](./week2-04-fused-model-kernels.md). +> [Day 3: SIMD Matrix Prefill](./week2-03-simd-matrix-prefill.md), +> [Day 4: Fused Model Primitives](./week2-04-fused-model-kernels.md), then +> [Day 5: Tiled Dense Prefill Attention](./week2-05-tiled-prefill-attention.md). > 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. diff --git a/book/src/week2-overview.md b/book/src/week2-overview.md index 6dd2be77..070982db 100644 --- a/book/src/week2-overview.md +++ b/book/src/week2-overview.md @@ -5,12 +5,13 @@ # 🚧 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 Days 1–4**: reuse previous keys +prefix. The **current Week 2 route contains Days 1–5**: reuse previous keys and values, bound their dense storage, keep projection weights packed, then reuse W4 and activation tiles during matrix-shaped prefill. Day 4 integrates -RMSNorm, RoPE, and SwiGLU as three separate model checkpoints. Its seven +RMSNorm, RoPE, and SwiGLU as three separate model checkpoints. Day 5 adds +tiled dense prefill attention and runs the final selected model. Its nine cumulative checkpoints are `kv-cache`, `capacity-cache`, `quantized-matvec`, -`simd-matmul`, `rmsnorm`, `rope`, and `swiglu`. +`simd-matmul`, `rmsnorm`, `rope`, `swiglu`, `tiled-prefill`, and `selected`. 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 @@ -38,6 +39,13 @@ RMSNorm, RoPE over the model's head layout, and fused SwiGLU in order. Compare each operator with its readable equation, run its cumulative checkpoint, then measure all four model variants under one matched 0.6B workload. +Finish with [Day 5: Tiled Dense Prefill Attention](./week2-05-tiled-prefill-attention.md). +Keep `swiglu` as the pre-edit product control. Build the BF16/D128 tiled path, +preserve causal and additive masks, GQA, and the short-query fallback, then +run `tiled-prefill` and the completed `selected` model on the same cached +0.6B request. The selected checkpoint names the final cumulative feature set; +it adds no second attention kernel. + ## Day 1 route | Step | What you own | Feedback | @@ -80,34 +88,52 @@ measure all four model variants under one matched 0.6B workload. +## Day 5 route + + + + + + + + + + + + +
StepWhat you ownFeedback
PrepareKeep the Day 4 swiglu model and save its matched 0.6B product controlDay 5 baseline
TileImplement BF16/D128 grouped attention with BQ32/BK16 tiles and online softmaxFocused causal GQA and partial-tail check
PreserveApply causal/additive masks, return zero for fully masked rows, and keep readable short-query attentionMask and fallback checks
IntegrateRun the cumulative tiled-prefill model, then the named final selected modelComplete Day 5 gate and live model commands
MeasureCompare Day 4 control, tiled prefill, selected, and MLX on one cached 0.6B requestDay 5 product loop
+ The Day 1 starter supplies model loading, test entrypoints, and benchmark and attribution helpers; you own the cache state and serving loop. Day 2 adds the packed-weight container and native operator boundaries, but you implement the embedding, kernels, and model wiring. Day 3 adds the SIMD matrix path behind that packed operator. The reference solution and full MLX model are separate controls; they do not fill your learner TODOs. Day 4 adds -three native primitives to the same cached, packed model. A cache +three native primitives to the same cached, packed model. Day 5 adds the +tiled attention operator and model selector. A cache counter shows which bytes moved, while a synchronized complete-request comparison shows whether a mechanism helped the chosen workload. -## Later lessons +## Historical lessons and Week 3 -The reviewed five-day design continues after the fused model primitives with -tiled dense prefill attention. Its **Day 5 checkpoints are planned, not shipped -in this four-day route**. Their commands and selectors are not current gates. -Week 3 uses later Week 2 interfaces; Days 1–4 alone do not supply every -prerequisite for its learner exercises. +The five active days end at `selected`. [Week 3](./week3-overview.md) reuses +the model and dense-cache interfaces while adding paging, batching, and +serving policy; its learner work remains separate from the single-request +Week 2 route. -The earlier seven-day book remains available at its old addresses as -[historical Week 2 material](./week2-02-benchmark-profile.md). It preserves +Earlier seven-day lessons whose URLs are not reused for active Days 1–5 remain +at their old addresses as +[historical Week 2 material](./week2-02-benchmark-profile.md). The former Day 4 +address now serves the [active fused-primitives lesson](./week2-04-fused-model-kernels.md). +The historical pages preserve 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 +bounded-decode and Split-K experiments. Their [operator-attribution diagram](./week2-kernel-profile.svg) and [decision diagram](./week2-performance-summary.svg) are also historical evidence, not diagrams of the current checkout. Those pages describe a different checkpoint order and may show commands unavailable in this partial branch. -Use Days 1–4 above for the current learner workflow. The +Use Days 1–5 above for the current learner workflow. The [performance evidence ledger](./appendix-performance.md) is likewise historical context, not a performance claim for this checkout. diff --git a/book/src/week3-01-continuous-batching.md b/book/src/week3-01-continuous-batching.md index 2b3910db..a27bbcdc 100644 --- a/book/src/week3-01-continuous-batching.md +++ b/book/src/week3-01-continuous-batching.md @@ -1,9 +1,10 @@ # 🚧 Week 3 Day 1: Continuous Batching -You begin with the completed Week 2 single-request model: multi-offset RoPE and -causal masking already have stable interfaces, and each request can own a dense -KV cache. The Day 1 starter leaves four learner-owned slices behind those -interfaces: +You begin with the completed +[Week 2 Day 5 single-request model](./week2-05-tiled-prefill-attention.md): +multi-offset RoPE and causal masking already have stable interfaces, and each +request can own a dense KV cache. The Day 1 starter leaves four learner-owned +slices behind those interfaces: - dense batch assembly and masking in `BatchingKvCache`; - `mlx_quantized_linear` plus the per-weight selector and the explicit @@ -115,8 +116,9 @@ src/tiny_llm/qwen3_week2.py::Qwen3ModelWeek2.__init__ src/tiny_llm/models.py::dispatch_week3_batch_model ``` -Week 2 ends with a course-owned quantized matmul so you can inspect its loader, -SIMD-matrix operations, and Split-K policy. Week 3 teaches serving mechanisms, +Week 2 ends with course-owned quantized projections and tiled dense attention. +The [Day 3 projection](./week2-03-simd-matrix-prefill.md) exposes its loader and +SIMD-matrix operations. Week 3 teaches serving mechanisms, so it should not make every cache and scheduler measurement depend on that teaching kernel's remaining projection overhead. diff --git a/book/src/week3-04-paged-attention-part2.md b/book/src/week3-04-paged-attention-part2.md index 1b90c91b..b65a9cf9 100644 --- a/book/src/week3-04-paged-attention-part2.md +++ b/book/src/week3-04-paged-attention-part2.md @@ -10,8 +10,9 @@ kernel may serve every query shape; separate decode and prefill kernels are an optimization choice. Day 5 replaces the supported BF16 long-prefill hot case with a tiled implementation. -> **Prerequisite:** Complete Week 3 Day 3's paged storage and Week 2 Day 5's -> online-softmax attention. The new concept here is translating logical K/V +> **Prerequisite:** Complete Week 3 Day 3's paged storage and +> [Week 2 Day 5's online-softmax attention](./week2-05-tiled-prefill-attention.md). +> The new concept here is translating logical K/V > positions through a block table. Tiled FlashAttention comes only after this > direct path works. diff --git a/book/src/week3-05-flash-attention.md b/book/src/week3-05-flash-attention.md index b3a39994..bbcd2119 100644 --- a/book/src/week3-05-flash-attention.md +++ b/book/src/week3-05-flash-attention.md @@ -16,18 +16,19 @@ preserving the same observable behavior. The manual control-flow trace at the end of the chapter is the completion feedback for that performance change; it is guidance, not a hidden grading requirement. -This is a required chapter. FlashAttention belongs here rather than in Week 2 -because the serving model's real K/V source is now the page pool. Building a -dense-only kernel first would create a second attention path and then require -students to relearn its memory schedule around page translation. +This is a required chapter. Week 2's dense tiled kernel teaches the +online-softmax update, but the serving model's K/V source is now the page +pool. Reuse that arithmetic while translating each K/V tile through the +block table; the dense loading schedule cannot read paged state directly. ## Prerequisites This chapter combines four prerequisites: -- The Week 2 decode-attention lab introduced the online-softmax recurrence. -- The Week 2 SIMD-matrix prefill lesson introduced the cooperative 32×32 tile built from BF16 8×8 - SIMD-matrix fragments. +- [Week 2 Day 5](./week2-05-tiled-prefill-attention.md) introduced the + online-softmax recurrence for tiled dense prefill attention. +- [Week 2 Day 3](./week2-03-simd-matrix-prefill.md) introduced cooperative + SIMD-matrix fragments for quantized projection prefill. - Week 3 Day 3 introduced physical pages and block tables. - Week 3 Day 4 introduced direct page-walking attention and the decode schedule. diff --git a/book/src/week3-overview.md b/book/src/week3-overview.md index dd919a21..37db7fd9 100644 --- a/book/src/week3-overview.md +++ b/book/src/week3-overview.md @@ -10,9 +10,10 @@ model uses one page-aware attention interface. A correct direct page-walking implementation may serve every query shape; the completed reference adds a tiled schedule for the supported BF16 long-prefill hot case. -Week 2's course-owned quantized projections remain the inspectable endpoint of -that week's kernel lessons. Week 3 deliberately switches dense-model -projections to `mx.quantized_matmul` at model construction, while retaining the +The [Week 2 `selected` model](./week2-05-tiled-prefill-attention.md) completes +the single-request route with tiled dense attention; its course-owned quantized +projections remain an inspectable kernel lesson. Week 3 deliberately switches +dense-model projections to `mx.quantized_matmul` at model construction, while retaining the course-owned normalization, activation, cache, attention, paging, and scheduler paths. This keeps Week 3 focused on serving-system mechanisms rather than carrying the teaching kernel's projection cost through every benchmark. diff --git a/main.py b/main.py index 86e81c87..590b3e91 100644 --- a/main.py +++ b/main.py @@ -42,6 +42,8 @@ "rmsnorm", "rope", "swiglu", + "tiled-prefill", + "selected", ), help="run one cumulative Week 2 model checkpoint", ) @@ -59,7 +61,7 @@ args.loader == "week3" or ( args.loader == "week2" - and (args.week2_checkpoint or "swiglu") + and (args.week2_checkpoint or "selected") not in ("kv-cache", "capacity-cache") ) ) diff --git a/scripts/dev-tools.py b/scripts/dev-tools.py index eebcf749..5b94aed7 100644 --- a/scripts/dev-tools.py +++ b/scripts/dev-tools.py @@ -15,9 +15,6 @@ 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 > 4: - print("Week 2 Days 1–4 are shipped; later tests are deferred") - return False return True diff --git a/scripts/week2_shipped_day_ci.json b/scripts/week2_shipped_day_ci.json deleted file mode 100644 index 212fbd5b..00000000 --- a/scripts/week2_shipped_day_ci.json +++ /dev/null @@ -1,147 +0,0 @@ -{ - "shipped_day": 4, - "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", - "benches/test_bench_week2_operators.py", - "benches/test_quantized_matmul.py", - "tests_refsol/test_week_2_day_2.py", - "tests_refsol/test_week_2_day_3.py", - "tests_refsol/test_week_2_day_4.py", - "tests_refsol/test_rmsnorm_register_cache.py", - "tests_refsol/test_extension_interface_sync.py" - ], - "deferred": [ - { - "path": "tests_refsol/test_week_2_day_5.py", - "reenable_day": 5, - "reason": "tiled attention and selected engine" - }, - { - "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_attention.py", - "reenable_day": 5, - "reason": "later attention kernels" - }, - { - "path": "tests_refsol/test_week_3_day_1.py", - "reenable_day": 5, - "reason": "post-Week-2 regression gate" - }, - { - "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" - } - ], - "deferred_builds": [], - "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" - ], - "required_builds": [ - "pdm run build-ext", - "pdm run build-ext-ref" - ] -} diff --git a/scripts/week2_shipped_day_ci.py b/scripts/week2_shipped_day_ci.py deleted file mode 100644 index 1b70244b..00000000 --- a/scripts/week2_shipped_day_ci.py +++ /dev/null @@ -1,73 +0,0 @@ -"""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"] != 4: - raise SystemExit("this temporary CI gate is bound to shipped_day=4") - print("shipped_day=4", 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 manifest["deferred_builds"]: - raise SystemExit("Day 3 requires both extension builds") - if manifest["required_builds"] != ["pdm run build-ext", "pdm run build-ext-ref"]: - raise SystemExit("Day 3 native build gates changed") - for command in manifest["required_builds"]: - print(f"BUILD {command}", 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/extensions/bindings.cpp b/src/extensions/bindings.cpp index 9922a8aa..958eea34 100644 --- a/src/extensions/bindings.cpp +++ b/src/extensions/bindings.cpp @@ -46,9 +46,8 @@ NB_MODULE(_ext, m) { "stream"_a = nb::none()); m.def("swiglu", &tiny_llm_ext::swiglu, "gate"_a, "up"_a, "stream"_a = nb::none()); - // Week 2, Day 6. - m.def("decode_attention", &tiny_llm_ext::decode_attention, "query"_a, "key"_a, "value"_a, "mask"_a, "scale"_a, - "is_causal"_a, "has_mask"_a, "num_heads"_a, "num_kv_heads"_a, "stream"_a = nb::none()); + m.def("_dense_attention_prefill_mma", &tiny_llm_ext::_dense_attention_prefill_mma, "query"_a, "key"_a, "value"_a, + "mask"_a, "scale"_a, "is_causal"_a, "has_mask"_a, "num_heads"_a, "num_kv_heads"_a, "stream"_a = nb::none()); // Week 3, Day 3. m.def("paged_cache_update", &tiny_llm_ext::paged_cache_update, "pages"_a, "values"_a, "page_id"_a, "start"_a, diff --git a/src/extensions/src/tiny_llm_ext.h b/src/extensions/src/tiny_llm_ext.h index 8580b3cc..e51bc8bc 100644 --- a/src/extensions/src/tiny_llm_ext.h +++ b/src/extensions/src/tiny_llm_ext.h @@ -100,14 +100,16 @@ class Week2SwiGLU : public mx::Primitive { const char *name() const override { return "Week2SwiGLU"; } }; -// Week 2, Day 6: implement online-softmax decode attention. -mx::array decode_attention(const mx::array &q, const mx::array &k, const mx::array &v, const mx::array &mask, - float scale, bool is_causal, bool has_mask, int num_heads, int num_kv_heads, - mx::StreamOrDevice s = {}); - -class Week2DecodeAttention : public mx::Primitive { +// Week 2, Day 5: implement BQ32/BK16 tiled prefill with online max/sum. +// The leading underscore keeps the raw native shape contract private; learners +// call the checked Python wrapper in tiny_llm.week2_kernels. +mx::array _dense_attention_prefill_mma(const mx::array &q, const mx::array &k, const mx::array &v, + const mx::array &mask, float scale, bool is_causal, bool has_mask, int num_heads, + int num_kv_heads, mx::StreamOrDevice s = {}); + +class Week2DensePrefillMMA : public mx::Primitive { public: - Week2DecodeAttention(mx::Stream stream, float scale, bool is_causal, bool has_mask, int num_heads, int num_kv_heads) + Week2DensePrefillMMA(mx::Stream stream, float scale, bool is_causal, bool has_mask, int num_heads, int num_kv_heads) : mx::Primitive(stream), scale_(scale), is_causal_(is_causal), @@ -118,9 +120,9 @@ class Week2DecodeAttention : public mx::Primitive { void eval_gpu(const std::vector &inputs, std::vector &outputs) override; std::pair, std::vector> vmap(const std::vector &, const std::vector &) override { - throw std::runtime_error("Week2DecodeAttention has no vmap implementation."); + throw std::runtime_error("Week2DensePrefillMMA has no vmap implementation."); } - const char *name() const override { return "Week2DecodeAttention"; } + const char *name() const override { return "Week2DensePrefillMMA"; } private: float scale_; diff --git a/src/extensions/src/week2_kernels.cpp b/src/extensions/src/week2_kernels.cpp index f14aebb2..d03bc5cd 100644 --- a/src/extensions/src/week2_kernels.cpp +++ b/src/extensions/src/week2_kernels.cpp @@ -50,18 +50,17 @@ void Week2SwiGLU::eval_gpu(const std::vector &, std::vector &, std::vector &) { - checkpoint_todo("Week2DecodeAttention::eval_cpu", "Week 2, Day 6"); +void Week2DensePrefillMMA::eval_cpu(const std::vector &, std::vector &) { + checkpoint_todo("Week2DensePrefillMMA::eval_cpu", "Week 2, Day 5"); } -void Week2DecodeAttention::eval_gpu(const std::vector &, std::vector &) { - checkpoint_todo("Week2DecodeAttention::eval_gpu", "Week 2, Day 6"); +void Week2DensePrefillMMA::eval_gpu(const std::vector &, std::vector &) { + checkpoint_todo("Week2DensePrefillMMA::eval_gpu", "Week 2, Day 5"); } } // namespace tiny_llm_ext diff --git a/src/extensions/src/week2_kernels.metal b/src/extensions/src/week2_kernels.metal index 0539ddb4..88f6835e 100644 --- a/src/extensions/src/week2_kernels.metal +++ b/src/extensions/src/week2_kernels.metal @@ -6,8 +6,8 @@ using namespace metal; // week2_rms_norm_register_cached (keep a readable fallback for wide rows) // week2_rope // week2_swiglu -// Week 2, Day 6: -// week2_decode_attention +// Week 2, Day 5: +// week2_dense_prefill_mma_bf16_d128 // // Add each [[kernel]] function when its task asks for it. The C++ starter // wrapper and binding already exist, but remain fail-closed until then. diff --git a/src/extensions/test.py b/src/extensions/test.py index e8d98a8a..137f07ec 100644 --- a/src/extensions/test.py +++ b/src/extensions/test.py @@ -9,7 +9,6 @@ "rms_norm": "Week 2, Day 4", "rope": "Week 2, Day 4", "swiglu": "Week 2, Day 4", - "decode_attention": "Week 2, Day 6", "paged_cache_update": "Week 3, Day 3", "quantized_embedding": "Week 3, Day 4", "paged_attention": "Week 3, Day 4", diff --git a/src/extensions_ref/bindings.cpp b/src/extensions_ref/bindings.cpp index 5693c837..267b8e1d 100644 --- a/src/extensions_ref/bindings.cpp +++ b/src/extensions_ref/bindings.cpp @@ -39,7 +39,9 @@ NB_MODULE(_ext, m) { m.def("swiglu", &tiny_llm_ext_ref::swiglu, "gate"_a, "up"_a, "stream"_a = nb::none()); m.def("decode_attention", &tiny_llm_ext_ref::decode_attention, "query"_a, "key"_a, "value"_a, "mask"_a, "scale"_a, "is_causal"_a, "has_mask"_a, "num_heads"_a, "num_kv_heads"_a, "stream"_a = nb::none()); - + m.def("_dense_attention_prefill_mma", &tiny_llm_ext_ref::_dense_attention_prefill_mma, "query"_a, "key"_a, + "value"_a, "mask"_a, "scale"_a, "is_causal"_a, "has_mask"_a, "num_heads"_a, "num_kv_heads"_a, + "stream"_a = nb::none()); m.def("paged_cache_update", &tiny_llm_ext_ref::paged_cache_update, "pages"_a, "values"_a, "page_id"_a, "start"_a, "stream"_a = nb::none()); m.def("paged_attention", &tiny_llm_ext_ref::paged_attention, "query"_a, "key_pages"_a, "value_pages"_a, diff --git a/src/extensions_ref/src/tiny_llm_ext.h b/src/extensions_ref/src/tiny_llm_ext.h index 7d75b189..02733d71 100644 --- a/src/extensions_ref/src/tiny_llm_ext.h +++ b/src/extensions_ref/src/tiny_llm_ext.h @@ -62,6 +62,9 @@ mx::array swiglu(const mx::array &gate, const mx::array &up, mx::StreamOrDevice mx::array decode_attention(const mx::array &q, const mx::array &k, const mx::array &v, const mx::array &mask, float scale, bool is_causal, bool has_mask, int num_heads, int num_kv_heads, mx::StreamOrDevice s = {}); +mx::array _dense_attention_prefill_mma(const mx::array &q, const mx::array &k, const mx::array &v, + const mx::array &mask, float scale, bool is_causal, bool has_mask, + int num_heads, int num_kv_heads, mx::StreamOrDevice s = {}); class Week2RMSNorm : public mx::Primitive { public: @@ -133,6 +136,32 @@ class Week2DecodeAttention : public mx::Primitive { int num_kv_heads_; }; +class Week2DensePrefillMMA : public mx::Primitive { +public: + Week2DensePrefillMMA(mx::Stream stream, float scale, bool is_causal, bool has_mask, int num_heads, + int num_kv_heads) + : mx::Primitive(stream), + scale_(scale), + is_causal_(is_causal), + has_mask_(has_mask), + num_heads_(num_heads), + num_kv_heads_(num_kv_heads) {} + void eval_cpu(const std::vector &inputs, std::vector &outputs) override; + void eval_gpu(const std::vector &inputs, std::vector &outputs) override; + std::pair, std::vector> vmap(const std::vector &, + const std::vector &) override { + throw std::runtime_error("Week2DensePrefillMMA has no vmap implementation."); + } + const char *name() const override { return "Week2DensePrefillMMA"; } + +private: + float scale_; + bool is_causal_; + bool has_mask_; + int num_heads_; + int num_kv_heads_; +}; + mx::array paged_attention(const mx::array &q, const mx::array &key_pages, const mx::array &value_pages, const mx::array &block_table, const mx::array &context_lens, const float scale, const bool is_causal, const int num_kv_heads, const int num_heads, mx::StreamOrDevice s = {}); diff --git a/src/extensions_ref/src/week2_kernels.cpp b/src/extensions_ref/src/week2_kernels.cpp index 569645ca..56a341ae 100644 --- a/src/extensions_ref/src/week2_kernels.cpp +++ b/src/extensions_ref/src/week2_kernels.cpp @@ -84,6 +84,28 @@ mx::array decode_attention(const mx::array &q, const mx::array &k, const mx::arr {q, k, v, mask}); } +mx::array _dense_attention_prefill_mma(const mx::array &q, const mx::array &k, const mx::array &v, + const mx::array &mask, float scale, bool is_causal, bool has_mask, int num_heads, + int num_kv_heads, mx::StreamOrDevice s) { + if (q.dtype() != mx::bfloat16 || k.dtype() != mx::bfloat16 || v.dtype() != mx::bfloat16 || + mask.dtype() != mx::float32) { + throw std::runtime_error("dense_attention_prefill_mma: expected BF16 q/k/v and float32 mask"); + } + if (q.ndim() != 3 || k.ndim() != 3 || v.ndim() != 3 || q.shape()[2] != 128 || k.shape()[2] != 128 || + v.shape()[2] != 128 || k.shape() != v.shape() || q.shape()[1] <= 8 || num_heads % num_kv_heads != 0 || + q.shape()[0] % num_heads != 0 || k.shape()[0] * num_heads != q.shape()[0] * num_kv_heads) { + throw std::runtime_error("dense_attention_prefill_mma: incompatible BF16/D128 attention shapes"); + } + if (has_mask && (mask.ndim() != 3 || mask.shape()[0] != q.shape()[0] || mask.shape()[1] != q.shape()[1] || + mask.shape()[2] != k.shape()[1])) { + throw std::runtime_error("dense_attention_prefill_mma: mask must have shape [B*Hq,L,S]"); + } + return mx::array( + q.shape(), q.dtype(), + std::make_shared(to_stream(s), scale, is_causal, has_mask, num_heads, num_kv_heads), + {q, k, v, mask}); +} + void Week2RMSNorm::eval_cpu(const std::vector &inputs, std::vector &outputs) { throw std::runtime_error("rms_norm: the course extension is GPU-only"); } @@ -100,6 +122,10 @@ void Week2DecodeAttention::eval_cpu(const std::vector &inputs, std::v throw std::runtime_error("decode_attention: the course extension is GPU-only"); } +void Week2DensePrefillMMA::eval_cpu(const std::vector &, std::vector &) { + throw std::runtime_error("dense_attention_prefill_mma: the course extension is GPU-only"); +} + #ifdef _METAL_ void Week2RMSNorm::eval_gpu(const std::vector &inputs, std::vector &outputs) { @@ -216,6 +242,41 @@ void Week2DecodeAttention::eval_gpu(const std::vector &inputs, std::v encoder.dispatch_threadgroups(MTL::Size(q_rows * length, 1, 1), MTL::Size(simdgroups_per_query * 32, 1, 1)); } +void Week2DensePrefillMMA::eval_gpu(const std::vector &inputs, std::vector &outputs) { + const auto &q = inputs[0]; + const auto &k = inputs[1]; + const auto &v = inputs[2]; + const auto &mask = inputs[3]; + auto &out = outputs[0]; + out.set_data(mx::allocator::malloc(out.nbytes())); + auto &d = mx::metal::device(stream().device); + auto kernel = d.get_kernel("week2_dense_prefill_mma_bf16_d128", d.get_library("tiny_llm_ext_ref")); + auto &encoder = mx::metal::get_command_encoder(stream()); + encoder.set_compute_pipeline_state(kernel); + encoder.set_input_array(q, 0); + encoder.set_input_array(k, 1); + encoder.set_input_array(v, 2); + encoder.set_input_array(mask, 3); + encoder.set_output_array(out, 4); + const int q_rows = q.shape()[0]; + const int length = q.shape()[1]; + const int context = k.shape()[1]; + const int is_causal = is_causal_; + const int has_mask = has_mask_; + encoder.set_bytes(q_rows, 5); + encoder.set_bytes(length, 6); + encoder.set_bytes(context, 7); + encoder.set_bytes(scale_, 8); + encoder.set_bytes(is_causal, 9); + encoder.set_bytes(has_mask, 10); + encoder.set_bytes(num_heads_, 11); + encoder.set_bytes(num_kv_heads_, 12); + constexpr int query_block = 32; + constexpr int threads = 128; + encoder.dispatch_threadgroups(MTL::Size((length + query_block - 1) / query_block, q_rows, 1), + MTL::Size(threads, 1, 1)); +} + #else void Week2RMSNorm::eval_gpu(const std::vector &, std::vector &) { @@ -230,7 +291,9 @@ void Week2SwiGLU::eval_gpu(const std::vector &, std::vector &, std::vector &) { throw std::runtime_error("Metal unavailable"); } - +void Week2DensePrefillMMA::eval_gpu(const std::vector &, std::vector &) { + throw std::runtime_error("Metal unavailable"); +} #endif } // namespace tiny_llm_ext_ref diff --git a/src/extensions_ref/src/week2_kernels.metal b/src/extensions_ref/src/week2_kernels.metal index 9ac7f2c7..92915b52 100644 --- a/src/extensions_ref/src/week2_kernels.metal +++ b/src/extensions_ref/src/week2_kernels.metal @@ -1,8 +1,65 @@ +#include #include #include "mlx/backend/metal/kernels/utils.h" +#include "cooperative_matrix.h" using namespace metal; +namespace { + +constant constexpr int DENSE_MMA_SIZE = 8; + +template +inline float dense_row_max( + thread const simdgroup_matrix* matrices) { + float value = -INFINITY; + for (int fragment_idx = 0; fragment_idx < N; ++fragment_idx) { + const auto elements = matrices[fragment_idx].thread_elements(); + value = max(value, max(elements[0], elements[1])); + } + value = max(value, simd_shuffle_xor(value, ushort(1))); + value = max(value, simd_shuffle_xor(value, ushort(8))); + return value; +} + +template +inline float dense_row_sum( + thread const simdgroup_matrix* matrices) { + float value = 0.0f; + for (int fragment_idx = 0; fragment_idx < N; ++fragment_idx) { + const auto elements = matrices[fragment_idx].thread_elements(); + value += elements[0] + elements[1]; + } + value += simd_shuffle_xor(value, ushort(1)); + value += simd_shuffle_xor(value, ushort(8)); + return value; +} + +inline void dense_clear_matrix( + thread simdgroup_matrix& matrix) { + matrix.thread_elements()[0] = 0.0f; + matrix.thread_elements()[1] = 0.0f; +} + +inline void dense_scale_matrix_rows( + thread simdgroup_matrix& matrix, + float scale) { + matrix.thread_elements()[0] *= scale; + matrix.thread_elements()[1] *= scale; +} + +template +inline void dense_matrix_multiply_accumulate( + thread simdgroup_matrix& accumulator, + thread simdgroup_matrix& left, + thread simdgroup_matrix& right) { + simdgroup_matrix result; + simdgroup_multiply_accumulate(result, left, right, accumulator); + accumulator = result; +} + +} // namespace + template [[kernel]] void week2_rms_norm( device const T* x [[buffer(0)]], @@ -284,6 +341,199 @@ template } } +[[kernel, max_total_threads_per_threadgroup(128)]] void week2_dense_prefill_mma_bf16_d128( + device const bfloat* q [[buffer(0)]], + device const bfloat* k [[buffer(1)]], + device const bfloat* v [[buffer(2)]], + device const float* mask [[buffer(3)]], + device bfloat* out [[buffer(4)]], + constant const int& q_rows [[buffer(5)]], + constant const int& length [[buffer(6)]], + constant const int& context [[buffer(7)]], + constant const float& scale [[buffer(8)]], + constant const int& is_causal [[buffer(9)]], + constant const int& has_mask [[buffer(10)]], + constant const int& num_heads [[buffer(11)]], + constant const int& num_kv_heads [[buffer(12)]], + uint2 group_id [[threadgroup_position_in_grid]], + ushort simd_gid [[simdgroup_index_in_threadgroup]], + ushort lane [[thread_index_in_simdgroup]]) { + constexpr int HEAD_DIM = 128; + constexpr int BQ = 32; + constexpr int BK = 16; + constexpr int SIMD_GROUPS = 4; + constexpr int THREADS = SIMD_GROUPS * 32; + constexpr int LDQ = HEAD_DIM + 2; + constexpr int LDK = BK + 2; + constexpr int LDV = HEAD_DIM + 2; + constexpr int KV_STORAGE = HEAD_DIM * LDK; + constexpr int SCORE_FRAGMENTS = BK / DENSE_MMA_SIZE; + constexpr int OUTPUT_FRAGMENTS = HEAD_DIM / DENSE_MMA_SIZE; + constexpr float LOG2_E = 1.44269504089f; + using QBlockLoader = tiny_llm::CooperativeTileLoader< + bfloat, BQ, HEAD_DIM, LDQ, THREADS>; + using KBlockLoader = tiny_llm::CooperativeTileLoader< + bfloat, BK, HEAD_DIM, LDK, THREADS, true>; + using VBlockLoader = tiny_llm::CooperativeTileLoader< + bfloat, BK, HEAD_DIM, LDV, THREADS>; + + const int query_block = group_id.x; + const int query_row_group = group_id.y; + if (query_row_group >= q_rows) return; + const int batch = query_row_group / num_heads; + const int query_head = query_row_group - batch * num_heads; + const int kv_head = query_head / (num_heads / num_kv_heads); + const int kv_row = batch * num_kv_heads + kv_head; + const int thread_idx = simd_gid * 32 + lane; + const ushort2 coordinate = tiny_llm::course_matrix_coordinate(lane); + const int query_row_in_block = simd_gid * DENSE_MMA_SIZE + coordinate.y; + const int query_position = query_block * BQ + query_row_in_block; + const bool query_valid = query_position < length; + const int live_queries = clamp(length - query_block * BQ, 0, BQ); + const float scale_log2 = scale * LOG2_E; + + threadgroup bfloat q_tile[BQ * LDQ]; + threadgroup bfloat kv_tile[KV_STORAGE]; + QBlockLoader::load( + q + query_row_group * length * HEAD_DIM + query_block * BQ * HEAD_DIM, + HEAD_DIM, + q_tile, + thread_idx, + live_queries, + HEAD_DIM); + threadgroup_barrier(mem_flags::mem_threadgroup); + + simdgroup_matrix output[OUTPUT_FRAGMENTS]; + for (int fragment_idx = 0; fragment_idx < OUTPUT_FRAGMENTS; ++fragment_idx) { + dense_clear_matrix(output[fragment_idx]); + } + float running_max = -INFINITY; + float running_sum = 0.0f; + + const int total_tiles = (context + BK - 1) / BK; + int tile_limit = total_tiles; + if (is_causal) { + const int last_query = min((query_block + 1) * BQ, length) - 1; + const int last_key = last_query + (context - length); + tile_limit = clamp((last_key + 1 + BK - 1) / BK, 0, total_tiles); + } + + for (int tile = 0; tile < tile_limit; ++tile) { + const int tile_start = tile * BK; + const int live_keys = clamp(context - tile_start, 0, BK); + KBlockLoader::load( + k + (kv_row * context + tile_start) * HEAD_DIM, + HEAD_DIM, + kv_tile, + thread_idx, + live_keys, + HEAD_DIM); + threadgroup_barrier(mem_flags::mem_threadgroup); + + simdgroup_matrix scores[SCORE_FRAGMENTS]; + for (int fragment_idx = 0; fragment_idx < SCORE_FRAGMENTS; ++fragment_idx) { + dense_clear_matrix(scores[fragment_idx]); + } + for (int dim = 0; dim < HEAD_DIM; dim += DENSE_MMA_SIZE) { + simdgroup_matrix q_fragment; + tiny_llm::course_load_matrix( + q_fragment, + q_tile + simd_gid * DENSE_MMA_SIZE * LDQ + dim, + LDQ, + lane); + for (int key_fragment = 0; key_fragment < SCORE_FRAGMENTS; ++key_fragment) { + simdgroup_matrix k_fragment; + tiny_llm::course_load_matrix( + k_fragment, + kv_tile + dim * LDK + key_fragment * DENSE_MMA_SIZE, + LDK, + lane); + dense_matrix_multiply_accumulate(scores[key_fragment], q_fragment, k_fragment); + } + } + + for (int key_fragment = 0; key_fragment < SCORE_FRAGMENTS; ++key_fragment) { + thread auto& values = scores[key_fragment].thread_elements(); + for (int element = 0; element < 2; ++element) { + const int tile_key = key_fragment * DENSE_MMA_SIZE + coordinate.x + element; + const int key_position = tile_start + tile_key; + bool valid = query_valid && key_position < context; + if (is_causal) { + valid = valid && key_position <= query_position + (context - length); + } + float score = valid ? values[element] * scale_log2 : -INFINITY; + if (valid && has_mask) { + score += mask[(query_row_group * length + query_position) * context + key_position] * LOG2_E; + } + values[element] = score; + } + } + + const float tile_max = dense_row_max(scores); + const float new_max = max(running_max, tile_max); + const bool finite_row = query_valid && new_max != -INFINITY; + const float previous_scale = running_max == -INFINITY || !finite_row + ? 0.0f + : fast::exp2(running_max - new_max); + for (int fragment_idx = 0; fragment_idx < SCORE_FRAGMENTS; ++fragment_idx) { + thread auto& values = scores[fragment_idx].thread_elements(); + for (int element = 0; element < 2; ++element) { + values[element] = values[element] == -INFINITY || !finite_row + ? 0.0f + : fast::exp2(values[element] - new_max); + } + } + const float tile_sum = dense_row_sum(scores); + running_max = new_max; + running_sum = previous_scale * running_sum + tile_sum; + for (int fragment_idx = 0; fragment_idx < OUTPUT_FRAGMENTS; ++fragment_idx) { + dense_scale_matrix_rows(output[fragment_idx], previous_scale); + } + + simdgroup_matrix probabilities[SCORE_FRAGMENTS]; + for (int fragment_idx = 0; fragment_idx < SCORE_FRAGMENTS; ++fragment_idx) { + const auto values = scores[fragment_idx].thread_elements(); + probabilities[fragment_idx].thread_elements()[0] = bfloat(values[0]); + probabilities[fragment_idx].thread_elements()[1] = bfloat(values[1]); + } + + threadgroup_barrier(mem_flags::mem_threadgroup); + VBlockLoader::load( + v + (kv_row * context + tile_start) * HEAD_DIM, + HEAD_DIM, + kv_tile, + thread_idx, + live_keys, + HEAD_DIM); + threadgroup_barrier(mem_flags::mem_threadgroup); + for (int output_fragment = 0; output_fragment < OUTPUT_FRAGMENTS; ++output_fragment) { + for (int key_fragment = 0; key_fragment < SCORE_FRAGMENTS; ++key_fragment) { + simdgroup_matrix value_fragment; + tiny_llm::course_load_matrix( + value_fragment, + kv_tile + key_fragment * DENSE_MMA_SIZE * LDV + output_fragment * DENSE_MMA_SIZE, + LDV, + lane); + dense_matrix_multiply_accumulate( + output[output_fragment], probabilities[key_fragment], value_fragment); + } + } + threadgroup_barrier(mem_flags::mem_threadgroup); + } + + if (query_valid) { + for (int fragment_idx = 0; fragment_idx < OUTPUT_FRAGMENTS; ++fragment_idx) { + const auto values = output[fragment_idx].thread_elements(); + for (int element = 0; element < 2; ++element) { + const int dim = fragment_idx * DENSE_MMA_SIZE + coordinate.x + element; + out[(query_row_group * length + query_position) * HEAD_DIM + dim] = running_sum == 0.0f + ? bfloat(0.0f) + : bfloat(values[element] / running_sum); + } + } + } +} + instantiate_kernel("week2_rms_norm_f32", week2_rms_norm, float); instantiate_kernel("week2_rms_norm_f16", week2_rms_norm, half); instantiate_kernel("week2_rms_norm_bf16", week2_rms_norm, bfloat16_t); diff --git a/src/tiny_llm/qwen3_week2.py b/src/tiny_llm/qwen3_week2.py index 33579a68..238c88e8 100644 --- a/src/tiny_llm/qwen3_week2.py +++ b/src/tiny_llm/qwen3_week2.py @@ -23,8 +23,7 @@ class Week2CheckpointFeatures: fast_rope: bool = False fast_swiglu: bool = False simdgroup_matmul: bool = False - decode_attention: bool = False - split_k_matmul: bool = False + tiled_prefill_attention: bool = False WEEK2_CHECKPOINT_FEATURES = MappingProxyType( @@ -60,13 +59,28 @@ class Week2CheckpointFeatures: fast_rope=True, fast_swiglu=True, ), + "tiled-prefill": Week2CheckpointFeatures( + bounded_kv_capacity=True, + quantized_weights=True, + simdgroup_matmul=True, + fast_rms_norm=True, + fast_rope=True, + fast_swiglu=True, + tiled_prefill_attention=True, + ), + "selected": Week2CheckpointFeatures( + bounded_kv_capacity=True, + quantized_weights=True, + simdgroup_matmul=True, + fast_rms_norm=True, + fast_rope=True, + fast_swiglu=True, + tiled_prefill_attention=True, + ), } ) WEEK2_CHECKPOINTS = tuple(WEEK2_CHECKPOINT_FEATURES) -DECODE_ATTENTION_MAX_CONTEXT = 256 -DECODE_ATTENTION_MAX_QUERY = 2 - class Qwen3MultiHeadAttention: def __init__( @@ -86,7 +100,7 @@ def __init__( rms_norm_eps: float = 1e-5, use_fast_rms_norm: bool = True, use_fast_rope: bool = True, - use_decode_attention: bool = True, + use_tiled_prefill_attention: bool = False, ): pass @@ -141,7 +155,7 @@ def __init__( use_fast_rms_norm: bool = True, use_fast_rope: bool = True, use_fast_swiglu: bool = True, - use_decode_attention: bool = True, + use_tiled_prefill_attention: bool = False, ): pass @@ -159,9 +173,11 @@ class Qwen3ModelWeek2: def __init__( self, mlx_model: Any, - checkpoint: str = "swiglu", + checkpoint: str = "selected", use_mlx_quantized_linear: bool = False, use_bounded_kv_capacity: bool | None = None, + use_register_cached_rms_norm: bool | None = None, + use_tiled_prefill_attention: bool | None = None, ): self.num_hidden_layers = mlx_model.args.num_hidden_layers pass diff --git a/src/tiny_llm/week2_kernels.py b/src/tiny_llm/week2_kernels.py index 1aac3bbd..23f1c5d3 100644 --- a/src/tiny_llm/week2_kernels.py +++ b/src/tiny_llm/week2_kernels.py @@ -1,9 +1,20 @@ import mlx.core as mx +from extensions.tiny_llm_ext import _ext as tiny_llm_ext_private + + +_NO_ATTENTION_MASK = mx.zeros((1,), dtype=mx.float32) +DENSE_PREFILL_MIN_QUERY = 9 class FastRMSNorm: def __init__(self, dim: int, weight: mx.array, eps: float = 1e-5): - pass + self.dim = dim + self.weight = weight + self.eps = eps + self.dispatch_counts = { + "register_cached": 0, + "fixed_width_fallback": 0, + } def __call__(self, x: mx.array) -> mx.array: pass @@ -44,4 +55,81 @@ def decode_attention_custom( scale: float, mask: mx.array | str | None = None, ) -> mx.array: - pass + # The bounded decode experiment is not learner-owned. Keep this compatibility + # name on the readable path for Week 3 callers. + return scaled_dot_product_attention(query, key, value, scale, mask) + + +def _prepare_dense_attention_inputs( + query: mx.array, + key: mx.array, + value: mx.array, + mask: mx.array | str | None, +) -> tuple[mx.array, mx.array, mx.array, mx.array, bool, bool]: + """Validate and flatten grouped-query inputs for the Day 5 native seam.""" + if query.ndim != 4 or key.ndim != 4 or value.ndim != 4: + raise ValueError("dense attention expects [B,H,L,D] query, key, and value") + batch_size, num_heads, query_length, head_dim = query.shape + key_batch_size, num_kv_heads, context_length, key_head_dim = key.shape + if batch_size != key_batch_size or key.shape != value.shape: + raise ValueError("query, key, and value batch dimensions must match") + if head_dim != key_head_dim or num_heads % num_kv_heads != 0: + raise ValueError("incompatible grouped-query attention shapes") + if isinstance(mask, str) and mask != "causal": + raise ValueError(f"unsupported attention mask: {mask}") + + query = mx.contiguous(query.reshape(batch_size * num_heads, query_length, head_dim)) + key = mx.contiguous( + key.reshape(batch_size * num_kv_heads, context_length, head_dim) + ) + value = mx.contiguous( + value.reshape(batch_size * num_kv_heads, context_length, head_dim) + ) + is_causal = isinstance(mask, str) and mask == "causal" + has_mask = isinstance(mask, mx.array) + if has_mask: + mask = mx.broadcast_to( + mask, (batch_size, num_heads, query_length, context_length) + ) + mask = mx.contiguous( + mask.astype(mx.float32).reshape( + batch_size * num_heads, query_length, context_length + ) + ) + else: + mask = _NO_ATTENTION_MASK + return query, key, value, mask, is_causal, has_mask + + +def dense_prefill_attention_mma( + query: mx.array, + key: mx.array, + value: mx.array, + scale: float, + mask: mx.array | str | None = None, +) -> mx.array: + """Run the learner-owned BQ32/BK16 online-softmax prefill kernel.""" + if query.ndim != 4 or key.ndim != 4 or value.ndim != 4: + raise ValueError("dense attention expects [B,H,L,D] query, key, and value") + expected_shape = query.shape + _, num_heads, query_length, head_dim = query.shape + num_kv_heads = key.shape[1] + if query.dtype != mx.bfloat16 or head_dim != 128: + raise ValueError("tiled dense prefill requires BF16 query/key/value with D=128") + if query_length < DENSE_PREFILL_MIN_QUERY: + raise ValueError(f"tiled dense prefill requires L >= {DENSE_PREFILL_MIN_QUERY}") + query, key, value, mask, is_causal, has_mask = _prepare_dense_attention_inputs( + query, key, value, mask + ) + result = tiny_llm_ext_private._dense_attention_prefill_mma( + query, + key, + value, + mask, + scale, + is_causal, + has_mask, + num_heads, + num_kv_heads, + ) + return result.reshape(expected_shape) diff --git a/src/tiny_llm_ref/qwen3_week2.py b/src/tiny_llm_ref/qwen3_week2.py index f98355b8..ead11838 100644 --- a/src/tiny_llm_ref/qwen3_week2.py +++ b/src/tiny_llm_ref/qwen3_week2.py @@ -12,9 +12,10 @@ from .positional_encoding import RoPE from .quantize import QuantizedWeights, dequantize_linear, quantized_linear from .week2_kernels import ( + DENSE_PREFILL_MIN_QUERY, FastRMSNorm, FastRoPE, - decode_attention_custom, + dense_prefill_attention_mma, swiglu, ) @@ -27,8 +28,7 @@ class Week2CheckpointFeatures: fast_rope: bool = False fast_swiglu: bool = False simdgroup_matmul: bool = False - decode_attention: bool = False - split_k_matmul: bool = False + tiled_prefill_attention: bool = False WEEK2_CHECKPOINT_FEATURES = MappingProxyType( @@ -64,13 +64,28 @@ class Week2CheckpointFeatures: fast_rope=True, fast_swiglu=True, ), + "tiled-prefill": Week2CheckpointFeatures( + bounded_kv_capacity=True, + quantized_weights=True, + simdgroup_matmul=True, + fast_rms_norm=True, + fast_rope=True, + fast_swiglu=True, + tiled_prefill_attention=True, + ), + "selected": Week2CheckpointFeatures( + bounded_kv_capacity=True, + quantized_weights=True, + simdgroup_matmul=True, + fast_rms_norm=True, + fast_rope=True, + fast_swiglu=True, + tiled_prefill_attention=True, + ), } ) WEEK2_CHECKPOINTS = tuple(WEEK2_CHECKPOINT_FEATURES) -DECODE_ATTENTION_MAX_CONTEXT = 256 -DECODE_ATTENTION_MAX_QUERY = 2 - def _linear(x: mx.array, weight: mx.array | QuantizedWeights) -> mx.array: if isinstance(weight, QuantizedWeights): @@ -109,7 +124,7 @@ def __init__( rms_norm_eps: float = 1e-5, use_fast_rms_norm: bool = True, use_fast_rope: bool = True, - use_decode_attention: bool = True, + use_tiled_prefill_attention: bool = False, ): self.hidden_size = hidden_size self.num_heads = num_heads @@ -127,7 +142,12 @@ def __init__( self.wv = wv self.wo = wo self.use_fast_rope = use_fast_rope - self.use_decode_attention = use_decode_attention + self.use_tiled_prefill_attention = use_tiled_prefill_attention + self.attention_dispatch_counts = { + "tiled_prefill": 0, + "tiled_prefill_fallback": 0, + "readable": 0, + } rope_cls = FastRoPE if use_fast_rope else RoPE norm_cls = FastRMSNorm if use_fast_rms_norm else RMSNorm self.rope = rope_cls(self.head_dim, max_seq_len, theta) @@ -162,13 +182,17 @@ def __call__( projection_k, projection_v, _, mask = cache.update_and_fetch( projection_k, projection_v, mask_length=L, mask=mask ) - if ( - self.use_decode_attention - and L <= DECODE_ATTENTION_MAX_QUERY - and projection_k.shape[-2] <= DECODE_ATTENTION_MAX_CONTEXT - and not isinstance(mask, mx.array) - ): - x = decode_attention_custom( + can_tiled_prefill = ( + L >= DENSE_PREFILL_MIN_QUERY + and projection_q.dtype == mx.bfloat16 + and self.head_dim == 128 + ) + if self.use_tiled_prefill_attention and not can_tiled_prefill: + self.attention_dispatch_counts["tiled_prefill_fallback"] += 1 + + if self.use_tiled_prefill_attention and can_tiled_prefill: + self.attention_dispatch_counts["tiled_prefill"] += 1 + x = dense_prefill_attention_mma( projection_q, projection_k, projection_v, @@ -176,6 +200,7 @@ def __call__( mask=mask, ) else: + self.attention_dispatch_counts["readable"] += 1 x = scaled_dot_product_attention_grouped( projection_q.astype(mx.float32), projection_k.astype(mx.float32), @@ -236,7 +261,7 @@ def __init__( use_fast_rms_norm: bool = True, use_fast_rope: bool = True, use_fast_swiglu: bool = True, - use_decode_attention: bool = True, + use_tiled_prefill_attention: bool = False, ): self.num_attention_heads = num_attention_heads self.hidden_size = hidden_size @@ -271,7 +296,7 @@ def __init__( rms_norm_eps=rms_norm_eps, use_fast_rms_norm=use_fast_rms_norm, use_fast_rope=use_fast_rope, - use_decode_attention=use_decode_attention, + use_tiled_prefill_attention=use_tiled_prefill_attention, ) def __call__( @@ -292,32 +317,49 @@ class Qwen3ModelWeek2: def __init__( self, mlx_model: Any, - checkpoint: str = "swiglu", + checkpoint: str = "selected", use_mlx_quantized_linear: bool = False, use_bounded_kv_capacity: bool | None = None, + use_register_cached_rms_norm: bool | None = None, + use_tiled_prefill_attention: 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 + for name, value in ( + ("use_bounded_kv_capacity", use_bounded_kv_capacity), + ("use_register_cached_rms_norm", use_register_cached_rms_norm), + ("use_tiled_prefill_attention", use_tiled_prefill_attention), ): - raise ValueError("use_bounded_kv_capacity must be a bool or None") + if value is not None and not isinstance(value, bool): + raise ValueError(f"{name} 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 + if use_register_cached_rms_norm is None: + use_register_cached_rms_norm = features.fast_rms_norm + if use_tiled_prefill_attention is None: + use_tiled_prefill_attention = features.tiled_prefill_attention use_quantized_weights = features.quantized_weights - use_fast_rms_norm = features.fast_rms_norm + use_fast_rms_norm = use_register_cached_rms_norm use_fast_rope = features.fast_rope use_fast_swiglu = features.fast_swiglu - use_decode_attention = features.decode_attention use_simdgroup_matmul = features.simdgroup_matmul - use_split_k_matmul = features.split_k_matmul + use_split_k_matmul = False self.num_hidden_layers = mlx_model.args.num_hidden_layers + self.use_bounded_kv_capacity = use_bounded_kv_capacity + self.use_register_cached_rms_norm = use_register_cached_rms_norm + self.use_tiled_prefill_attention = use_tiled_prefill_attention + self.mechanism_controls = MappingProxyType( + { + "capacity_cache": use_bounded_kv_capacity, + "register_cached_rmsnorm": use_register_cached_rms_norm, + "tiled_prefill": use_tiled_prefill_attention, + } + ) self.use_fast_rope = use_fast_rope self.hidden_size = mlx_model.args.hidden_size self.vocab_size = mlx_model.args.vocab_size @@ -383,7 +425,7 @@ def model_weight(layer: Any) -> mx.array | QuantizedWeights: use_fast_rms_norm=use_fast_rms_norm, use_fast_rope=use_fast_rope, use_fast_swiglu=use_fast_swiglu, - use_decode_attention=use_decode_attention, + use_tiled_prefill_attention=use_tiled_prefill_attention, ) self.layers_inner.append(layer) norm_cls = FastRMSNorm if use_fast_rms_norm else RMSNorm diff --git a/src/tiny_llm_ref/week2_kernels.py b/src/tiny_llm_ref/week2_kernels.py index 661b3c98..bb32169f 100644 --- a/src/tiny_llm_ref/week2_kernels.py +++ b/src/tiny_llm_ref/week2_kernels.py @@ -1,19 +1,30 @@ import mlx.core as mx from extensions_ref import tiny_llm_ext_ref +from extensions_ref.tiny_llm_ext_ref import _ext as tiny_llm_ext_ref_private from .basics import softmax _NO_ATTENTION_MASK = mx.zeros((1,), dtype=mx.float32) +DENSE_PREFILL_MIN_QUERY = 9 + class FastRMSNorm: def __init__(self, dim: int, weight: mx.array, eps: float = 1e-5): self.dim = dim self.weight = weight self.eps = eps + self.dispatch_counts = { + "register_cached": 0, + "fixed_width_fallback": 0, + } def __call__(self, x: mx.array) -> mx.array: + if x.shape[-1] <= 4096: + self.dispatch_counts["register_cached"] += 1 + else: + self.dispatch_counts["fixed_width_fallback"] += 1 return tiny_llm_ext_ref.rms_norm( mx.contiguous(x), mx.contiguous(self.weight.astype(x.dtype)), self.eps ) @@ -145,3 +156,76 @@ def decode_attention_custom( num_kv_heads, ) return result.reshape(batch_size, num_heads, query_length, head_dim) + + +def _prepare_dense_attention_inputs( + query: mx.array, + key: mx.array, + value: mx.array, + mask: mx.array | str | None, +) -> tuple[mx.array, mx.array, mx.array, mx.array, bool, bool]: + if query.ndim != 4 or key.ndim != 4 or value.ndim != 4: + raise ValueError("dense attention expects [B,H,L,D] query, key, and value") + batch_size, num_heads, query_length, head_dim = query.shape + key_batch_size, num_kv_heads, context_length, key_head_dim = key.shape + if batch_size != key_batch_size or key.shape != value.shape: + raise ValueError("query, key, and value batch dimensions must match") + if head_dim != key_head_dim or num_heads % num_kv_heads != 0: + raise ValueError("incompatible grouped-query attention shapes") + if isinstance(mask, str) and mask != "causal": + raise ValueError(f"unsupported attention mask: {mask}") + + query = mx.contiguous(query.reshape(batch_size * num_heads, query_length, head_dim)) + key = mx.contiguous( + key.reshape(batch_size * num_kv_heads, context_length, head_dim) + ) + value = mx.contiguous( + value.reshape(batch_size * num_kv_heads, context_length, head_dim) + ) + is_causal = isinstance(mask, str) and mask == "causal" + has_mask = isinstance(mask, mx.array) + if has_mask: + mask = mx.broadcast_to( + mask, (batch_size, num_heads, query_length, context_length) + ) + mask = mx.contiguous( + mask.astype(mx.float32).reshape( + batch_size * num_heads, query_length, context_length + ) + ) + else: + mask = _NO_ATTENTION_MASK + return query, key, value, mask, is_causal, has_mask + + +def dense_prefill_attention_mma( + query: mx.array, + key: mx.array, + value: mx.array, + scale: float, + mask: mx.array | str | None = None, +) -> mx.array: + if query.ndim != 4 or key.ndim != 4 or value.ndim != 4: + raise ValueError("dense attention expects [B,H,L,D] query, key, and value") + expected_shape = query.shape + batch_size, num_heads, query_length, head_dim = query.shape + num_kv_heads = key.shape[1] + if query.dtype != mx.bfloat16 or head_dim != 128: + raise ValueError("tiled dense prefill requires BF16 query/key/value with D=128") + if query_length < DENSE_PREFILL_MIN_QUERY: + raise ValueError(f"tiled dense prefill requires L >= {DENSE_PREFILL_MIN_QUERY}") + query, key, value, mask, is_causal, has_mask = _prepare_dense_attention_inputs( + query, key, value, mask + ) + result = tiny_llm_ext_ref_private._dense_attention_prefill_mma( + query, + key, + value, + mask, + scale, + is_causal, + has_mask, + num_heads, + num_kv_heads, + ) + return result.reshape(expected_shape) diff --git a/tests/utils.py b/tests/utils.py index 71b8f97f..fe15fb7f 100644 --- a/tests/utils.py +++ b/tests/utils.py @@ -9,8 +9,11 @@ PRECISION_IDS = ["f32", "f16"] -def tiny_qwen3_mlx_model(num_hidden_layers: int = 1) -> SimpleNamespace: +def tiny_qwen3_mlx_model( + num_hidden_layers: int = 1, *, head_dim: int = 64 +) -> SimpleNamespace: """Build a small MLX-shaped Qwen3 model for integration tests.""" + hidden_size = 2 * head_dim def quantized_layer(out_dim: int, in_dim: int) -> SimpleNamespace: weight = mx.random.normal((out_dim, in_dim)).astype(mx.bfloat16) @@ -25,12 +28,12 @@ def quantized_layer(out_dim: int, in_dim: int) -> SimpleNamespace: args = SimpleNamespace( num_hidden_layers=num_hidden_layers, - hidden_size=128, + hidden_size=hidden_size, vocab_size=128, num_attention_heads=2, num_key_value_heads=1, - head_dim=64, - intermediate_size=128, + head_dim=head_dim, + intermediate_size=hidden_size, rms_norm_eps=1e-5, max_position_embeddings=256, rope_theta=10000, @@ -41,30 +44,32 @@ def quantized_layer(out_dim: int, in_dim: int) -> SimpleNamespace: layers.append( SimpleNamespace( self_attn=SimpleNamespace( - q_proj=quantized_layer(128, 128), - k_proj=quantized_layer(64, 128), - v_proj=quantized_layer(64, 128), - o_proj=quantized_layer(128, 128), - q_norm=SimpleNamespace(weight=mx.ones((64,), mx.bfloat16)), - k_norm=SimpleNamespace(weight=mx.ones((64,), mx.bfloat16)), + q_proj=quantized_layer(hidden_size, hidden_size), + k_proj=quantized_layer(head_dim, hidden_size), + v_proj=quantized_layer(head_dim, hidden_size), + o_proj=quantized_layer(hidden_size, hidden_size), + q_norm=SimpleNamespace(weight=mx.ones((head_dim,), mx.bfloat16)), + k_norm=SimpleNamespace(weight=mx.ones((head_dim,), mx.bfloat16)), ), mlp=SimpleNamespace( - gate_proj=quantized_layer(128, 128), - up_proj=quantized_layer(128, 128), - down_proj=quantized_layer(128, 128), + gate_proj=quantized_layer(hidden_size, hidden_size), + up_proj=quantized_layer(hidden_size, hidden_size), + down_proj=quantized_layer(hidden_size, hidden_size), + ), + input_layernorm=SimpleNamespace( + weight=mx.ones((hidden_size,), mx.bfloat16) ), - input_layernorm=SimpleNamespace(weight=mx.ones((128,), mx.bfloat16)), post_attention_layernorm=SimpleNamespace( - weight=mx.ones((128,), mx.bfloat16) + weight=mx.ones((hidden_size,), mx.bfloat16) ), ) ) return SimpleNamespace( args=args, model=SimpleNamespace( - embed_tokens=quantized_layer(128, 128), + embed_tokens=quantized_layer(128, hidden_size), layers=layers, - norm=SimpleNamespace(weight=mx.ones((128,), mx.bfloat16)), + norm=SimpleNamespace(weight=mx.ones((hidden_size,), mx.bfloat16)), ), ) diff --git a/tests_refsol/test_dense_attention_prototypes.py b/tests_refsol/test_dense_attention_prototypes.py new file mode 100644 index 00000000..16482843 --- /dev/null +++ b/tests_refsol/test_dense_attention_prototypes.py @@ -0,0 +1,108 @@ +"""Focused model-free controls for tiled dense-prefill attention.""" + +from math import prod + +import mlx.core as mx +import pytest + +from tiny_llm_ref.week2_kernels import ( + dense_prefill_attention_mma, + scaled_dot_product_attention, +) +from .utils import assert_allclose + + +def _fixture(shape: tuple[int, ...], phase: float, dtype=mx.bfloat16) -> mx.array: + values = mx.sin(mx.arange(prod(shape), dtype=mx.float32) * 0.017 + phase) + return values.reshape(shape).astype(dtype) + + +def _explicit_mask(length: int, context: int) -> mx.array: + values = mx.where( + mx.arange(context) % 5 == 0, + mx.array(-1.25, dtype=mx.float32), + mx.array(0.0, dtype=mx.float32), + ) + return mx.broadcast_to(values.reshape(1, 1, 1, context), (1, 1, length, context)) + + +@pytest.mark.parametrize( + ("length", "context", "gqa_ratio", "mask_kind"), + ( + (9, 9, 1, "none"), + (15, 16, 2, "causal"), + (16, 31, 4, "explicit"), + (31, 32, 1, "causal"), + (32, 33, 2, "explicit"), + (33, 47, 4, "none"), + ), +) +def test_tiled_prefill_matches_readable_boundary_sweep( + length: int, + context: int, + gqa_ratio: int, + mask_kind: str, +): + query_heads = 4 + kv_heads = query_heads // gqa_ratio + query = _fixture((1, query_heads, length, 128), 0.1) + key = _fixture((1, kv_heads, context, 128), 0.7) + value = _fixture(key.shape, 1.3) + mask = { + "none": None, + "causal": "causal", + "explicit": _explicit_mask(length, context), + }[mask_kind] + + actual = dense_prefill_attention_mma(query, key, value, 128**-0.5, mask) + expected = scaled_dot_product_attention( + query.astype(mx.float32), + key.astype(mx.float32), + value.astype(mx.float32), + 128**-0.5, + mask, + ).astype(mx.bfloat16) + + assert actual.shape == query.shape + assert actual.dtype == mx.bfloat16 + assert_allclose( + actual, + expected, + mx.bfloat16, + atol=2e-2, + rtol=2e-2, + message=f"L={length}, S={context}, GQA={gqa_ratio}, mask={mask_kind}", + ) + + +def test_tiled_prefill_fully_masked_rows_are_finite_zero_and_contract_is_narrow(): + query = _fixture((1, 4, 9, 128), 0.1) + key = _fixture((1, 2, 17, 128), 0.7) + value = _fixture(key.shape, 1.3) + mask = mx.full((1, 1, 9, 17), -mx.inf, dtype=mx.float32) + + actual = dense_prefill_attention_mma(query, key, value, 128**-0.5, mask) + mx.eval(actual) + assert bool(mx.all(mx.isfinite(actual)).item()) + assert float(mx.max(mx.abs(actual)).item()) == 0.0 + + with pytest.raises(ValueError, match="L >= 9"): + dense_prefill_attention_mma(query[:, :, :8], key, value, 128**-0.5) + with pytest.raises(ValueError, match="BF16"): + dense_prefill_attention_mma( + query.astype(mx.float32), + key.astype(mx.float32), + value.astype(mx.float32), + 128**-0.5, + ) + + +@pytest.mark.parametrize("context", (2048, 8192)) +def test_tiled_prefill_long_context_shape_smoke(context: int): + query = _fixture((1, 1, context, 128), 0.1) + key = _fixture((1, 1, context, 128), 0.7) + value = _fixture(key.shape, 1.3) + actual = dense_prefill_attention_mma(query, key, value, 128**-0.5) + mx.eval(actual) + assert actual.shape == query.shape + assert bool(mx.all(mx.isfinite(actual)).item()) diff --git a/tests_refsol/test_extension_interface_sync.py b/tests_refsol/test_extension_interface_sync.py index a09d3df2..21f0c289 100644 --- a/tests_refsol/test_extension_interface_sync.py +++ b/tests_refsol/test_extension_interface_sync.py @@ -14,7 +14,7 @@ "rms_norm": ("Week 2, Day 4", "week2_kernels.cpp"), "rope": ("Week 2, Day 4", "week2_kernels.cpp"), "swiglu": ("Week 2, Day 4", "week2_kernels.cpp"), - "decode_attention": ("Week 2, Day 6", "week2_kernels.cpp"), + "_dense_attention_prefill_mma": ("Week 2, Day 5", "week2_kernels.cpp"), "paged_cache_update": ("Week 3, Day 3", "paged_attention.cpp"), "quantized_embedding": ("Week 3, Day 4", "quantized_matmul.cpp"), "paged_attention": ("Week 3, Day 4", "paged_attention.cpp"), @@ -25,7 +25,7 @@ "Week2RMSNorm": ("Week 2, Day 4", "week2_kernels.cpp"), "Week2RoPE": ("Week 2, Day 4", "week2_kernels.cpp"), "Week2SwiGLU": ("Week 2, Day 4", "week2_kernels.cpp"), - "Week2DecodeAttention": ("Week 2, Day 6", "week2_kernels.cpp"), + "Week2DensePrefillMMA": ("Week 2, Day 5", "week2_kernels.cpp"), "PagedCacheUpdate": ("Week 3, Day 3", "paged_attention.cpp"), "QuantizedEmbedding": ("Week 3, Day 4", "quantized_matmul.cpp"), "PagedAttention": ("Week 3, Day 4", "paged_attention.cpp"), @@ -42,7 +42,7 @@ "week2_rms_norm_register_cached": "Week 2, Day 4", "week2_rope": "Week 2, Day 4", "week2_swiglu": "Week 2, Day 4", - "week2_decode_attention": "Week 2, Day 6", + "week2_dense_prefill_mma_bf16_d128": "Week 2, Day 5", }, "paged_attention.metal": { "paged_cache_update_kernel": "Week 3, Day 3", @@ -53,7 +53,7 @@ } DOC_TASK_MARKERS = { - "book/src/week2-03-quantize-model.md": { + "book/src/week2-02-quantize-model.md": { "Task 1": {"QuantizedWeights.from_mlx_layer", "QuantizedEmbedding.__call__"}, "Task 2": {"tiny_llm_ext::quantized_matmul", "QuantizedMatmul::eval_cpu"}, "Task 3": { @@ -67,29 +67,15 @@ "Task 1": { "tiny_llm_ext::rms_norm", "Week2RMSNorm::eval_gpu", - "week2_rms_norm", + "week2_rms_norm_register_cached", }, "Task 2": {"tiny_llm_ext::rope", "Week2RoPE::eval_gpu", "week2_rope"}, "Task 3": {"tiny_llm_ext::swiglu", "Week2SwiGLU::eval_gpu", "week2_swiglu"}, "Task 4": {"Qwen3ModelWeek2.__init__", "Qwen3MLP.__call__"}, }, - "book/src/week2-05-simd-matrix-prefill.md": { - "Task 2": { - "QuantizedMatmul::eval_gpu", - "quantized_matmul_simdgroup_w4a16_g128", - }, - }, - "book/src/week2-06-operator-lab.md": { - "Task 2": { - "tiny_llm_ext::decode_attention", - "Week2DecodeAttention::eval_gpu", - "week2_decode_attention", - "Qwen3MultiHeadAttention.__call__", - "decode_attention_custom", - }, - }, - "book/src/week2-07-split-k-prefill.md": { - "Task 3": {"QuantizedMatmul::eval_gpu"}, + "book/src/week2-03-simd-matrix-prefill.md": { + "Task 1": {"quantized_matmul_simdgroup_w4a16_g128"}, + "Task 2": {"QuantizedMatmul::eval_gpu"}, }, "book/src/week3-03-paged-attention-part1.md": { "Task 1": { @@ -131,7 +117,7 @@ } EXTENSION_TASK_PAIRS = { - "book/src/week2-03-quantize-model.md": { + "book/src/week2-02-quantize-model.md": { "Task 2": { ( "src/extensions/src/quantized_matmul.cpp", @@ -156,7 +142,10 @@ ("src/extensions/src/week2_kernels.cpp", "tiny_llm_ext::rms_norm"), ("src/extensions/src/week2_kernels.cpp", "Week2RMSNorm::eval_cpu"), ("src/extensions/src/week2_kernels.cpp", "Week2RMSNorm::eval_gpu"), - ("src/extensions/src/week2_kernels.metal", "week2_rms_norm"), + ( + "src/extensions/src/week2_kernels.metal", + "week2_rms_norm_register_cached", + ), }, "Task 2": { ("src/extensions/src/week2_kernels.cpp", "tiny_llm_ext::rope"), @@ -171,33 +160,15 @@ ("src/extensions/src/week2_kernels.metal", "week2_swiglu"), }, }, - "book/src/week2-05-simd-matrix-prefill.md": { - "Task 2": { - ("src/extensions/src/quantized_matmul.cpp", "QuantizedMatmul::eval_gpu"), + "book/src/week2-03-simd-matrix-prefill.md": { + "Task 1": { ( "src/extensions/src/quantized_matmul.metal", "quantized_matmul_simdgroup_w4a16_g128", ), }, - }, - "book/src/week2-06-operator-lab.md": { "Task 2": { - ( - "src/extensions/src/week2_kernels.cpp", - "tiny_llm_ext::decode_attention", - ), - ( - "src/extensions/src/week2_kernels.cpp", - "Week2DecodeAttention::eval_cpu", - ), - ( - "src/extensions/src/week2_kernels.cpp", - "Week2DecodeAttention::eval_gpu", - ), - ( - "src/extensions/src/week2_kernels.metal", - "week2_decode_attention", - ), + ("src/extensions/src/quantized_matmul.cpp", "QuantizedMatmul::eval_gpu"), }, }, "book/src/week3-03-paged-attention-part1.md": { @@ -302,7 +273,7 @@ def _normalized_cpp_body(body: str) -> str: def _learner_owned_cpp_definitions(source: str) -> set[str]: source = _strip_comments(source) - wrappers = set(re.findall(r"\bmx::array\s+([a-z][a-z0-9_]*)\s*\(", source)) + wrappers = set(re.findall(r"\bmx::array\s+([a-z_][a-z0-9_]*)\s*\(", source)) evaluators = set( re.findall(r"\bvoid\s+([A-Z][A-Za-z0-9_]*::eval_(?:cpu|gpu))\s*\(", source) ) @@ -333,8 +304,13 @@ def _assert_metal_scaffold_only(source: str) -> None: ) +def _assert_metal_checkpoint_marker(source: str, kernel: str, checkpoint: str) -> None: + assert kernel in source + assert checkpoint in source + + def _header_functions(path: str) -> set[str]: - return set(re.findall(r"mx::array\s+([a-z][a-z0-9_]*)\s*\(", _read(path))) + return set(re.findall(r"mx::array\s+([a-z_][a-z0-9_]*)\s*\(", _read(path))) def _header_declaration(path: str, function: str) -> str: @@ -347,7 +323,7 @@ def _header_declaration(path: str, function: str) -> str: def _binding_functions(path: str) -> set[str]: - return set(re.findall(r'm\.def\("([a-z][a-z0-9_]*)"', _read(path))) + return set(re.findall(r'm\.def\("([a-z_][a-z0-9_]*)"', _read(path))) def _binding_contract( @@ -358,7 +334,7 @@ def _binding_contract( target = re.search(r"&([A-Za-z_][A-Za-z0-9_:]*)", match.group(1)) assert target is not None, f"missing binding target for {function}" arguments = re.findall( - r'"([a-z][a-z0-9_]*)"_a(?:\s*=\s*([A-Za-z0-9_:().+\-]+))?', + r'"([a-z_][a-z0-9_]*)"_a(?:\s*=\s*([A-Za-z0-9_:().+\-]+))?', match.group(1), ) return target.group(1), [(name, default or None) for name, default in arguments] @@ -442,7 +418,9 @@ def _assert_task_pairs(chapter: str, task: str, pairs: set[tuple[str, str]]) -> def test_starter_and_reference_publish_the_same_learner_extension_functions(): expected = set(INTERFACES) - assert _header_functions("src/extensions_ref/src/tiny_llm_ext.h") == expected + assert _header_functions("src/extensions_ref/src/tiny_llm_ext.h") == expected | { + "decode_attention" + } assert _header_functions("src/extensions/src/tiny_llm_ext.h") == expected reference_bindings = _binding_functions("src/extensions_ref/bindings.cpp") - { @@ -452,7 +430,7 @@ def test_starter_and_reference_publish_the_same_learner_extension_functions(): "load_library", "axpby", } - assert reference_bindings == expected + assert reference_bindings == expected | {"decode_attention"} assert starter_bindings == expected for function in expected: @@ -508,8 +486,7 @@ def test_starter_metal_stubs_name_each_learner_owned_kernel_and_checkpoint(): source = _read(f"src/extensions/src/{filename}") assert f"src/{filename}" in cmake for kernel, checkpoint in kernels.items(): - assert kernel in source - assert checkpoint in source + _assert_metal_checkpoint_marker(source, kernel, checkpoint) _assert_metal_scaffold_only(source) @@ -577,7 +554,9 @@ def test_built_starter_extension_fails_closed_for_every_public_operation(monkeyp "rms_norm": lambda: extension.rms_norm(scalar, scalar, 1e-5), "rope": lambda: extension.rope(scalar, scalar, 1, 10_000.0), "swiglu": lambda: extension.swiglu(scalar, scalar), - "decode_attention": lambda: extension.decode_attention( + "_dense_attention_prefill_mma": lambda: importlib.import_module( + "tiny_llm_ext._ext" + )._dense_attention_prefill_mma( scalar, scalar, scalar, scalar, 1.0, False, False, 1, 1 ), "paged_cache_update": lambda: extension.paged_cache_update( @@ -617,20 +596,19 @@ def test_binding_guard_rejects_a_changed_python_default(): ) -def test_task_pair_guard_rejects_a_nonexistent_source_path(): - path = "book/src/week2-04-fused-model-kernels.md" - source = _read(path) - wrong_path = source.replace( - "`Week2RMSNorm::eval_gpu` in `src/extensions/src/week2_kernels.cpp`", - "`Week2RMSNorm::eval_gpu` in `src/extensions/src/wrong.cpp`", +def test_day_3_metal_guard_rejects_a_missing_simd_matmul_kernel(): + source = _read("src/extensions/src/quantized_matmul.metal") + missing_kernel = source.replace( + "quantized_matmul_simdgroup_w4a16_g128", + "quantized_matmul_removed_w4a16_g128", 1, ) - assert wrong_path != source + assert missing_kernel != source with pytest.raises(AssertionError): - _assert_task_pairs( - wrong_path, - "Task 1", - EXTENSION_TASK_PAIRS[path]["Task 1"], + _assert_metal_checkpoint_marker( + missing_kernel, + "quantized_matmul_simdgroup_w4a16_g128", + "Week 2, Day 3", ) diff --git a/tests_refsol/test_week_2_day_1.py b/tests_refsol/test_week_2_day_1.py index 6c4f9e00..8c530f96 100644 --- a/tests_refsol/test_week_2_day_1.py +++ b/tests_refsol/test_week_2_day_1.py @@ -3,10 +3,19 @@ import mlx.core as mx import pytest -from .tiny_llm_base import Embedding, Qwen3ModelWeek2, RMSNorm, RoPE, TinyKvFullCache +from .tiny_llm_base import Qwen3ModelWeek2, TinyKvFullCache from .utils import assert_allclose, tiny_qwen3_mlx_model +def _fixed_fixture(seed: int): + random_state = mx.random.state[:] + try: + mx.random.seed(seed) + return tiny_qwen3_mlx_model() + finally: + mx.random.state[:] = random_state + + def test_task_1_full_cache_appends_chunks(): cache = TinyKvFullCache() key_1 = mx.random.normal((1, 2, 3, 4)).astype(mx.bfloat16) @@ -31,23 +40,22 @@ def test_task_1_full_cache_appends_chunks(): def test_tasks_2_and_3_cached_checkpoint_is_runnable_and_readable(): - model = Qwen3ModelWeek2(tiny_qwen3_mlx_model(), checkpoint="kv-cache") - layer = model.layers_inner[0] - - 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 + fixture = _fixed_fixture(0) + model = Qwen3ModelWeek2(fixture, checkpoint="kv-cache") + cache = model.create_kv_cache() + assert len(cache) == fixture.args.num_hidden_layers - output = model(mx.array([[1, 2]], dtype=mx.int32), 0, model.create_kv_cache()) - assert output.dtype == mx.bfloat16 + prefill = model(mx.array([[1, 2]], dtype=mx.int32), 0, cache) + decoded = model(mx.array([[3]], dtype=mx.int32), 2, cache) + complete = model(mx.array([[1, 2, 3]], dtype=mx.int32), 0, model.create_kv_cache()) + assert prefill.dtype == decoded.dtype == complete.dtype == mx.bfloat16 + assert prefill.shape == (1, 2, fixture.args.vocab_size) + assert decoded.shape == (1, 1, fixture.args.vocab_size) + assert_allclose(decoded, complete[:, -1:, :], mx.bfloat16) def test_task_3_rejects_a_position_that_disagrees_with_the_cache(): - model = Qwen3ModelWeek2(tiny_qwen3_mlx_model(), checkpoint="kv-cache") + model = Qwen3ModelWeek2(_fixed_fixture(0), checkpoint="kv-cache") with pytest.raises(ValueError): model(mx.array([[1]], dtype=mx.int32), 1, model.create_kv_cache()) @@ -77,17 +85,14 @@ def test_capacity_cache_exposes_only_the_logical_prefix(): 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 - ) + + # Unused request capacity must never appear in the returned K/V prefix. + assert cached_key.shape[2] == 3 < cache.capacity -def test_capacity_rewind_reuses_storage_and_overflow_is_transactional(): +def test_capacity_rewind_and_overflow_preserve_the_logical_prefix(): cache = TinyKvFullCache(capacity=3) key, value = _chunk(0, 2) cache.update_and_fetch(key, value) @@ -101,7 +106,6 @@ def test_capacity_rewind_reuses_storage_and_overflow_is_transactional(): 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), @@ -113,43 +117,51 @@ def test_capacity_rewind_reuses_storage_and_overflow_is_transactional(): 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, + # A rejected append must leave the retained prefix available for rewind. + cache.rewind(1) + final_key, final_value = _chunk(40, 1) + cached_key, cached_value, offset, _ = cache.update_and_fetch(final_key, final_value) + assert offset == 3 + assert_allclose( + cached_key, + mx.concat([key[:, :, :1], replacement_key[:, :, :1], final_key], axis=2), + mx.float32, + ) + assert_allclose( + cached_value, + mx.concat([value[:, :, :1], replacement_value[:, :, :1], final_value], axis=2), + mx.float32, ) + 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 + restarted_key, restarted_value = _chunk(60, 1) + cached_key, cached_value, offset, _ = cache.update_and_fetch( + restarted_key, restarted_value + ) + assert offset == 1 + assert_allclose(cached_key, restarted_key, mx.float32) + assert_allclose(cached_value, restarted_value, mx.float32) def test_capacity_checkpoint_runs_the_week2_engine(): - model = Qwen3ModelWeek2(tiny_qwen3_mlx_model(), checkpoint="capacity-cache") - assert model.use_bounded_kv_capacity - + fixture = _fixed_fixture(0) + model = Qwen3ModelWeek2(fixture, checkpoint="capacity-cache") + readable_model = Qwen3ModelWeek2(fixture, checkpoint="kv-cache") bounded = model.create_kv_cache(capacity=3) - readable = Qwen3ModelWeek2( - tiny_qwen3_mlx_model(), checkpoint="kv-cache" - ).create_kv_cache() + readable = readable_model.create_kv_cache() inputs = mx.array([[1, 2, 3]], dtype=mx.int32) actual = model(inputs, 0, bounded) - expected = model(inputs, 0, readable) + expected = readable_model(inputs, 0, readable) + assert actual.dtype == expected.dtype == mx.bfloat16 + assert actual.shape == expected.shape == (1, 3, fixture.args.vocab_size) 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) + assert all(cache.offset == 3 for cache in bounded) + with pytest.raises(ValueError, match="capacity"): + model(mx.array([[4]], dtype=mx.int32), 3, bounded) diff --git a/tests_refsol/test_week_2_day_2.py b/tests_refsol/test_week_2_day_2.py index c6b83415..41c16180 100644 --- a/tests_refsol/test_week_2_day_2.py +++ b/tests_refsol/test_week_2_day_2.py @@ -1,8 +1,5 @@ """Week 2 Day 2 quantized-matvec tests.""" -import importlib -import inspect - import mlx.core as mx import pytest @@ -16,16 +13,21 @@ Qwen3ModelWeek2, QuantizedEmbedding, QuantizedWeights, - RMSNorm, - RoPE, + dequantize_weights, 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 _fixed_fixture(seed: int): + random_state = mx.random.state[:] + try: + mx.random.seed(seed) + return tiny_qwen3_mlx_model() + finally: + mx.random.state[:] = random_state def test_task_1_quantized_embedding_dequantizes_selected_rows(): @@ -60,31 +62,40 @@ def test_task_1_quantized_embedding_accepts_sampled_uint32_tokens(): 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__) +def test_task_1_dequantize_weights_matches_packed_weight_values(): + weight = mx.sin(mx.arange(5 * 256, dtype=mx.float32) * 0.07).reshape(5, 256) + packed, scales, biases = mx.quantize( + weight.astype(mx.bfloat16), group_size=128, bits=4 ) - 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") - assert hasattr(model, "layers_inner"), ( - "implement the Qwen3ModelWeek2 quantized-matvec learner seam" + actual = dequantize_weights(packed, scales, biases, 128, 4) + expected = mx.dequantize(packed, scales, biases, group_size=128, bits=4) + assert actual is not None, "implement the dequantize_weights learner seam" + assert actual.shape == (5, 256) + assert_allclose(actual, expected, mx.bfloat16, atol=2e-2, rtol=2e-2) + + +@pytest.mark.parametrize("seed", (0, 2)) +def test_task_4_quantized_model_matches_packed_mlx_control(seed: int): + fixture = _fixed_fixture(seed) + model = Qwen3ModelWeek2(fixture, checkpoint="quantized-matvec") + control = Qwen3ModelWeek2( + fixture, checkpoint="quantized-matvec", use_mlx_quantized_linear=True ) - layer = model.layers_inner[0] + tokens = mx.array([[1, 2, 3]], dtype=mx.int32) + with mx.stream(mx.gpu): + actual_cache = model.create_kv_cache(capacity=4) + expected_cache = control.create_kv_cache(capacity=4) + actual = model(tokens, 0, actual_cache) + expected = control(tokens, 0, expected_cache) + decoded = model(mx.array([[4]], dtype=mx.int32), 3, actual_cache) + decoded_control = control(mx.array([[4]], dtype=mx.int32), 3, expected_cache) + mx.eval(actual, expected, decoded, decoded_control) - 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 model.use_bounded_kv_capacity - assert not layer.self_attn.use_decode_attention - assert not layer.mlp.use_fast_swiglu + assert actual.dtype == expected.dtype == mx.bfloat16 + assert actual.shape == expected.shape == (1, 3, fixture.args.vocab_size) + assert_allclose(actual, expected, mx.bfloat16, atol=1.0, rtol=0.05) + assert decoded.shape == decoded_control.shape == (1, 1, fixture.args.vocab_size) + assert_allclose(decoded, decoded_control, mx.bfloat16, atol=1.0, rtol=0.05) def quantized_matmul_helper( @@ -251,4 +262,4 @@ def test_day2_public_selectors_preserve_day2_prefix(): "week2-capacity-cache", "week2-quantized-matvec", ] - assert SECTIONS == ("embedding", "decode-projections", "prefill-projections") + assert SECTIONS[:3] == ("embedding", "decode-projections", "prefill-projections") diff --git a/tests_refsol/test_week_2_day_4.py b/tests_refsol/test_week_2_day_4.py index a0af369b..eb6d4639 100644 --- a/tests_refsol/test_week_2_day_4.py +++ b/tests_refsol/test_week_2_day_4.py @@ -92,7 +92,7 @@ def test_task_4_primitive_checkpoints_compose_public_model(seed): ) -def test_day4_public_selectors_include_only_shipped_checkpoints(): +def test_day4_public_selectors_remain_the_prefix_of_day5(): model_module = importlib.import_module(Qwen3ModelWeek2.__module__) checkpoints = ( "kv-cache", @@ -103,9 +103,9 @@ def test_day4_public_selectors_include_only_shipped_checkpoints(): "rope", "swiglu", ) - assert model_module.WEEK2_CHECKPOINTS == checkpoints - assert KNOWN_CHECKPOINTS == checkpoints - assert DEFAULT_CASES == ( + assert model_module.WEEK2_CHECKPOINTS[:7] == checkpoints + assert KNOWN_CHECKPOINTS[:7] == checkpoints + assert DEFAULT_CASES[:7] == ( "kv-cache:decode:128", "capacity-cache:decode:128", "quantized-matvec:decode:128", @@ -114,7 +114,7 @@ def test_day4_public_selectors_include_only_shipped_checkpoints(): "rope:decode:128", "swiglu:decode:128", ) - assert [variant.key for variant in WEEK2_VARIANTS] == [ + assert [variant.key for variant in WEEK2_VARIANTS][:8] == [ "week1", "week2-kv-cache", "week2-capacity-cache", @@ -123,7 +123,6 @@ def test_day4_public_selectors_include_only_shipped_checkpoints(): "week2-rmsnorm", "week2-rope", "week2-swiglu", - "mlx", ] diff --git a/tests_refsol/test_week_2_day_5.py b/tests_refsol/test_week_2_day_5.py new file mode 100644 index 00000000..dc3fb957 --- /dev/null +++ b/tests_refsol/test_week_2_day_5.py @@ -0,0 +1,100 @@ +"""Week 2 Day 5 public tiled-attention and selected-engine tests.""" + +from math import prod + +import mlx.core as mx +import pytest + +from .tiny_llm_base import Qwen3ModelWeek2, dense_prefill_attention_mma +from .utils import assert_allclose, tiny_qwen3_mlx_model + + +def _fixture(shape: tuple[int, ...], phase: float) -> mx.array: + values = mx.sin(mx.arange(prod(shape), dtype=mx.float32) * 0.017 + phase) + return values.reshape(shape).astype(mx.bfloat16) + + +def _model_fixture(seed: int): + state = mx.random.state[:] + try: + mx.random.seed(seed) + return tiny_qwen3_mlx_model(head_dim=128) + finally: + mx.random.state[:] = state + + +def _run(model, length: int): + tokens = mx.array([list(range(1, length + 1))], dtype=mx.int32) + output = model(tokens, 0, model.create_kv_cache(capacity=length)) + mx.eval(output) + assert output.dtype == mx.bfloat16 + assert output.shape[:2] == (1, length) + return output + + +def test_task_1_tiled_prefill_matches_mlx_causal_gqa(): + query = _fixture((1, 4, 33, 128), 0.1) + key = _fixture((1, 2, 47, 128), 0.7) + value = _fixture(key.shape, 1.3) + actual = dense_prefill_attention_mma(query, key, value, 128**-0.5, "causal") + expected = mx.fast.scaled_dot_product_attention( + query.astype(mx.float32), + key.astype(mx.float32), + value.astype(mx.float32), + scale=128**-0.5, + mask="causal", + ).astype(mx.bfloat16) + assert actual.shape == query.shape + assert actual.dtype == mx.bfloat16 + assert_allclose(actual, expected, mx.bfloat16, atol=2e-2, rtol=2e-2) + + +def test_task_2_fully_masked_rows_are_finite_zero(): + query = _fixture((1, 4, 9, 128), 0.1) + key = _fixture((1, 1, 17, 128), 0.7) + value = _fixture(key.shape, 1.3) + mask = mx.full((1, 1, 9, 17), -mx.inf, dtype=mx.float32) + actual = dense_prefill_attention_mma(query, key, value, 128**-0.5, mask) + mx.eval(actual) + assert bool(mx.all(mx.isfinite(actual)).item()) + assert float(mx.max(mx.abs(actual)).item()) == 0.0 + + +@pytest.mark.parametrize("length", (3, 10)) +def test_task_3_tiled_checkpoint_matches_readable_model(length: int): + fixture = _model_fixture(4) + tiled = Qwen3ModelWeek2(fixture, checkpoint="tiled-prefill") + readable = Qwen3ModelWeek2(fixture, checkpoint="swiglu") + actual = _run(tiled, length) + expected = _run(readable, length) + assert_allclose(actual, expected, mx.bfloat16, atol=0.75, rtol=0.05) + + +@pytest.mark.parametrize( + "disabled", + ( + {"use_bounded_kv_capacity": False}, + {"use_register_cached_rms_norm": False}, + {"use_tiled_prefill_attention": False}, + ), +) +def test_selected_controls_can_be_disabled_independently(disabled): + fixture = _model_fixture(0) + selected = Qwen3ModelWeek2(fixture, checkpoint="selected") + switched = Qwen3ModelWeek2(fixture, checkpoint="selected", **disabled) + actual = _run(selected, 10) + expected = _run(switched, 10) + assert_allclose(actual, expected, mx.bfloat16, atol=0.75, rtol=0.05) + + +def test_selected_is_default_and_matches_explicit_checkpoint(): + fixture = _model_fixture(0) + default = Qwen3ModelWeek2(fixture) + explicit = Qwen3ModelWeek2(fixture, checkpoint="selected") + assert_allclose(_run(default, 10), _run(explicit, 10), mx.bfloat16) + + +@pytest.mark.parametrize("checkpoint", ("decode-attention", "split-k")) +def test_retired_experiments_are_not_week2_checkpoints(checkpoint): + with pytest.raises(ValueError, match="unknown Week 2 checkpoint"): + Qwen3ModelWeek2(_model_fixture(0), checkpoint=checkpoint)