|
| 1 | +# Technical Report: PR #2205 - One prefill plan and one trim hook for the CLI and the server |
| 2 | + |
| 3 | +**Date**: 2026-10-07 |
| 4 | + |
| 5 | +**Status**: Implemented and verified on GB10; pending merge. |
| 6 | + |
| 7 | +**Languages**: Rust, Markdown |
| 8 | + |
| 9 | +**Risk Level**: Medium. |
| 10 | +- Every CLI and server prefill now executes one plan, and the server's default prefill chunk changes from 512 to 2048. |
| 11 | +- Greedy output on three real checkpoints is unchanged before and after. |
| 12 | +- The new default's effect on concurrent decode latency is not measured yet; the epic's end-of-run benchmark covers it. |
| 13 | + |
| 14 | +## Executive Summary |
| 15 | + |
| 16 | +Phase 3 of epic #2166 (#2170). |
| 17 | +- **One plan.** `mlxcel_core::prefill_plan::PrefillPlan` now decides, in one place, how a prompt is split into forwards: the adopted prompt-cache prefix, the history-boundary split, chunking, tile padding and the pad trim. Every CLI and server prefill site only executes its pieces. |
| 18 | +- **One trim hook.** `LanguageModel::trim_state(Option<SequenceId>, excess)` replaces the two separate hooks. |
| 19 | +- **One chunk size.** The chunk is a single policy, 2048 per ADR 0007's measurement. |
| 20 | +- **Prompt-cache divergence explained.** The hit vs miss divergence found in Phase 0 is root-caused (the prefill partition) and bounded by an explicit invariant. |
| 21 | + |
| 22 | +## 1. Problem Statement |
| 23 | + |
| 24 | +- **Prefill was decided twice.** The CLI chunked at 2048, the server at 512, and padding and trimming were coded separately on each side. |
| 25 | +- **Two trim hooks.** The model API had `trim_internal_caches` for the CLI and `trim_sequence_state` for the server, so a family had to implement both. #2140 was that failure. |
| 26 | +- **Prompt-cache divergence.** The Phase 0 parity baseline showed a server prompt-cache hit diverging from the uncached run of the same prompt; Qwen3-1.7B greedy diverged at token 0. |
| 27 | + |
| 28 | +## 2. Change Summary |
| 29 | + |
| 30 | +- **`PrefillPlan`:** |
| 31 | + - `new`/`with_prefix(prompt_len, adopted, boundary, chunk, PrefillCaps)` produce ordered `PrefillPiece { range, padded_len }`. |
| 32 | + - The boundary segment is one unpadded piece, and chunks restart at the boundary. |
| 33 | + - The final piece is tile-padded where the model and input allow it. |
| 34 | + - `reproduces` and `split_points` state when a cache hit can reproduce a miss. |
| 35 | +- **CLI:** `prefill_prompt_last_logits`, the `generate_with_stats` copy and both embedding variants run plan pieces through one executor. |
| 36 | +- **Server:** |
| 37 | + - `planned_prefill.rs` runs pieces for full, chunked, batched and handoff prefill. The plan is rebuilt from the sequence each tick, with `prefill_offset` as the cursor. |
| 38 | + - The history-boundary snapshot is inserted after the boundary piece, with no extra forward. |
| 39 | + - Gemma 4's `mtp_prefill_ranges` and the paged block reservation use the plan. |
| 40 | +- **`trim_state`:** seven model-owned families implement it, the VLM wrappers forward it, and `LoadedModel` delegates to it. `rewind_decode_appends` stays a separate exact-rewind contract. |
| 41 | +- **Chunk policy:** `prefill_chunk_len()` reads `MLXCEL_PREFILL_CHUNK`, else 2048. It is now also the server's `--prefill-chunk-size` default, on both server binaries. |
| 42 | + - The memory estimate's activation term follows it (`activation_prefill_tokens()`). |
| 43 | + - The batched-prefill token budget caps each row's share at 512, so the cohort budget stays 4096. |
| 44 | + |
| 45 | +## 3. Technical Decisions |
| 46 | + |
| 47 | +**Root cause of the cache divergence: the partition, not adoption.** Under `MLXCEL_SDPA_DETERMINISTIC=1`, chunking the miss at the cached length made miss and hit identical: |
| 48 | +- Qwen3-1.7B: chunk at 45. |
| 49 | +- Llama-3.2-1B: chunk at 71. |
| 50 | + |
| 51 | +The mechanism: the dense-KV miss is one 52- or 75-row forward (tiled qmm). The hit forwards only a 7- or 4-row suffix, which takes the per-row qmv kernel (`M*B < 8`) and a different attention tiling. The two reduce in a different order and flip a near-tied token. |
| 52 | + |
| 53 | +**Documented bound instead of forcing the split.** A hit reproduces the miss exactly when its adopted prefix ends on a split point of the miss's plan, and the adopted rows were written by those same pieces. Splitting every family at the history boundary would make the harness rows identical, but it was rejected for three reasons: |
| 54 | +- It adds a forward launch to every cold chat prefill. |
| 55 | +- It would make the prompt cache change single-turn output. |
| 56 | +- It does not fix whole-prompt replay anyway. |
| 57 | + |
| 58 | +**One chunk size: 2048.** Phase 0 measured 8 to 29 percent lower TTFT at 8192 tokens with 2048. Review found three places that silently assumed 512, and they were fixed: |
| 59 | +- `mlx_server`'s flag default. |
| 60 | +- A `clamp(1, 0)` panic in the memory estimate. |
| 61 | +- The batched-prefill budget. |
| 62 | + |
| 63 | +## 4. Validation |
| 64 | + |
| 65 | +- **Unit tests:** |
| 66 | + - `prefill_plan` 11, mlxcel-core `prefill` 76, `generate::` 46, `decode_finish` 9. |
| 67 | + - Root modules run one at a time: `scheduler_prompt_cache_plan_tests` 3 (including the chunk-counter and progress-frame cases), `server::batch::scheduler::tests` 94, `scheduler_prompt_cache_tests` 31, `prefill_span_coverage_tests` 3, `finish_step_tests`, `speculative_burst_tests`, `prefill_cohort`, `cli_input`, `memory_estimate`, `gemma4_mtp_target`. |
| 68 | + - `vlm_wrapper_capability_delegation`, including the tile-padding guard audit. |
| 69 | +- **Real checkpoints:** `mlxcel-engine-parity` with `MLXCEL_SDPA_DETERMINISTIC=1` on qwen3-1.7b-4bit, llama-3.2-1b-instruct-4bit and lfm2-350m-8bit. |
| 70 | + - Every pair result was unchanged before, after and after the merge with `main`. |
| 71 | + - CLI and server token streams were identical to the base commit. |
| 72 | +- **Known flakiness:** the full `server::batch::scheduler` filter aborts intermittently in CUDA graph capture. Every module run alone is clean. |
| 73 | + |
| 74 | +## 5. Residual Risks |
| 75 | + |
| 76 | +- **Concurrent decode latency is unmeasured.** Inter-token latency of concurrent decode streams under the 2048 default is left to the end-of-epic benchmark. A longer prefill piece delays decode grants by up to four times. |
| 77 | +- **History segments are not chunked.** A snapshot family's history segment is one unchunked forward regardless of chunk size. This is pre-existing (#1143). |
| 78 | +- **One environment variable now drives both paths.** `MLXCEL_PREFILL_CHUNK` now also sets the server default, and the memory estimate does not see an explicit `--prefill-chunk-size`. |
| 79 | +- **No numeric bound test.** Nothing asserts a numeric bound on hit-vs-miss divergence. The bound is the documented invariant plus the measured table. |
| 80 | + |
| 81 | +## 6. Learning Points |
| 82 | + |
| 83 | +- **Treat a partition as part of the numerics.** Two correct prefills of the same tokens differ in their low bits when they are split differently, because the kernel dispatch depends on row count. |
| 84 | +- **Changing a default needs a sweep for its implicit dependents.** Three of the four review fixes were places that assumed the old 512 without naming it. |
| 85 | + |
| 86 | +## 7. Related |
| 87 | + |
| 88 | +- Epic #2166, issue #2170. |
| 89 | +- ADR 0007, ADR 0005. |
| 90 | +- #2201 (finish step, merged into this branch). |
| 91 | +- #2140, #1143, #2185. |
0 commit comments