Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 4 additions & 3 deletions .github/workflows/macos.yml
Original file line number Diff line number Diff line change
Expand Up @@ -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
9 changes: 7 additions & 2 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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 | ✅ | ✅ | ✅ | 🚧 |
Expand All @@ -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.
Expand Down
4 changes: 3 additions & 1 deletion benches/bench.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
)
Expand Down Expand Up @@ -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")
)
)
Expand Down
14 changes: 14 additions & 0 deletions benches/bench_course_progression.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 = {
Expand Down
4 changes: 3 additions & 1 deletion benches/bench_week2_operators.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,8 @@
"embedding",
"decode-projections",
"prefill-projections",
"model-kernels",
"attention",
)


Expand Down Expand Up @@ -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",
Expand Down
24 changes: 22 additions & 2 deletions benches/profile_week2_kernels.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -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:
Expand All @@ -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,
)


Expand Down Expand Up @@ -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,
Expand Down
Loading