Skip to content

feat(mtp): per-row cache_len in the fused step — the pipeline blocker #100 named - #102

Merged
heydryft merged 2 commits into
masterfrom
feat/mtp-step-per-seq-cachelen
Aug 17, 2026
Merged

heydryft merged 2 commits into
masterfrom
feat/mtp-step-per-seq-cachelen

Conversation

@heydryft

Copy link
Copy Markdown
Contributor

Stacked on #100 (feat/v4-ragged-mask-wire), which is stacked on #95, which is stacked on #92. Review those first.

The blocker, in one sentence

MtpSpeculativePipeline::step keyed the entire fused step on ONE scalar cache_len — the batched buffer's width, which left-alignment makes max_j L_j and not any row's length — so all four of the quantities a per-sequence step needs were computed against the wrong number.

#92 named two blockers; #95 removed the first, #100 removed the second and named this one. This removes it. cache_supports_per_sequence_advance, model_masks_ragged_batches and this step's own arithmetic all pass for a V4 target now.

🔑 The per-row vector already existed — clone_in_cache was throwing it away

front_align_batch returns lead_pad[i], the dead prefix it left ahead of row i. MtpSpeculativePipeline::clone_in_cache bound it to _lead_pad and dropped it. Once front_pad_single has run, the slot's own current_seq_len is the batch width for every row — so that discarded vector is the only surviving record of who is actually where, and discarding it is what forced every downstream quantity onto one shared scalar.

resolve_row_cache_lens hands it straight back: cache_lens[i] = cache_len - lead_pad[i]. Nothing is plumbed, no HashMap<seq_id, …> appears, no signature changes through the engine. Same shape as #100's discovery that seqlen_offsets[i] was already the per-row q0.

The four defects, and how each closes

# defect fix
1 uncached[i] = toks_i.len() - cache_len saturates to 0 for every row but the longest ⇒ window_ok fails on every ragged step plan_step_window takes each row's own length; every tail is 1, which is the B=1 invariant at every batch size
2 toks[cache_len..] is a panic site once cache_len > toks_i.len() toks[cache_lens[i]..], and ok now implies the slice is in range (u_i >= 1 ⟺ cache_lens[i] < tok_lens[i])
3 prefill_window = Some((w, cache_len)) collapses seqlen_offsets to one scalar (last_n_context_len.1), so the row_q0 #100 built is fed uniform values per-row route through inputs_processor — see below
4 set_len(cache_len + c_i) counts the dead prefix inside the length, so it accumulates linearly drop_dead_prefix strips it first; the recorded length is cache_lens[i] + c_i, the row's own

🔑 Defect 4: the code was wrong, not the doc

front_pad_single's doc claims the dead prefix tracks a sqrt(steps) random walk. That claim is a property of the code, not of the batch, and it only holds if the prefix is taken back out on the way to the per-sequence caches.

It was not. front_pad_single sets current_seq_len = target_len (it has to — that is where the one shared append offset comes from) and clone_out_cache copies that length verbatim into every row. Recording it as the sequence's own makes the next front_align_batch pad relative to already-padded lengths, so each step adds another max_j c_j - c_i columns.

Worse than "it grows": the rows re-converge onto one length. Every row's recorded length ends up advancing at the fastest row's rate — the cohort barrier this whole workstream exists to remove, reappearing as padding. the_dead_prefix_does_not_accumulate_across_steps runs the loop at 6 and 12 steps and pins the inflation at 10 and 22 columns; with the strip it is 0 at both.

So the fix is drop_dead_prefix, the inverse of front_pad_kv_cache — a ~25-line symmetric primitive that restores the invariant everything else already assumes (front_pad_single's own narrow(dim, 0, live) requires the live run to start at column 0). No new field on SingleCache, and lead == 0 — every B=1 request, every uniform batch — touches no tensor at all.

🔑 Defect 3: chunk_offset_toks was the flattening, one layer down

get_prompt_input did let offset = input_seqs[0].token_offset(); — a per-sequence value collapsed to row 0 — and make_prompt_chunk then used it for seqlen_offsets, position_ids, seqlens_k and the paged slot mappings. Every one of those wants per-row.

So the trait signature is untouched (last_n_context_len: Option<(usize, usize)> stays; ~15 vision processors are not churned). make_prompt_chunk gains row_offsets: Option<&[usize]> which shadows chunk_offset_toks inside the per-sequence loop, and the value reaches it on Sequence next to prefill_prompt_toks — the field whose doc already says "only meant for internal speculative decoding usage". reset_prefill_toks clears both, so an offset can never outlive the window it describes.

Why "flag off is byte-identical" is structural

Two dispatch functions, both returning None for a uniform batch, and None runs the pre-change code verbatim:

  • resolve_row_cache_lens — None unless some row carries a lead and the mode is authorized;
  • resolve_row_offsets — None unless some sequence carries its own offset.

No fourth flag. The gate is ARC_MTP_PER_SEQ_KV, coordinated with #92 and #95 exactly as #100 did it: front_align_batch is the only producer of a ragged dense batch and it runs only under KvAdvance::PerSequence.

The draft grouping key moved from u to (u, chain_start). On a batch that shares a width the chain start is a function of u alone, so the BTreeMap partitions and iterates identically — but keying on u alone under raggedness would not be slow, it would be silent: a row whose seed is not at the group's chain start is skipped, so it would stop drafting with no signal but a falling tok_per_step.

⚠️ window_ok's deferral is no longer reachable for a padded batch — it refuses instead. #100 was right that the deferral lands in the target's own decode, which builds no mask at t_q == 1 (layers_masker.rs:285), so handing it a left-aligned batch would attend the dead prefix as real keys. Deferring is only safe for a batch that is not padded.

📊 Measured — CPU unit tests and plan arithmetic, not hardware (D14)

cargo test -p mistralrs-core --lib → 423 passed, 0 failed (was 407 on #100's tip) · --test synthetic_load_smoke → 13 passed · cargo check --workspace --tests green · scoped clippy lane green · mistralrs-core clippy introduces zero new diagnostics (diffed before/after on the touched files) · zero rustfmt drift, checked like-for-like.

Mutation runs — 17 mutations, all 17 caught

resolve_row_cache_lens: always take the per-row path ............ 1 FAILED
resolve_row_cache_lens: ignore the opt-in flag .................. 1 FAILED
resolve_row_cache_lens: never refuse a lead past the batch width  1 FAILED
resolve_row_cache_lens: never refuse "no row at the end column" .. 1 FAILED
plan_step_window: use the batch max for every row (the old scalar) 1 FAILED
plan_step_window: drop the `c + u >= 2` half .................... 1 FAILED
plan_step_window: drop the `(1..=w)` half ....................... 3 FAILED
draft_group_key: key on the uncached tail alone ................. 1 FAILED
per_sequence_refusal: drop the remaining capability check ....... 1 FAILED
drop_dead_prefix: make it a no-op ............................... 3 FAILED
drop_front_single: keep the length, only move the data .......... 2 FAILED
drop_front_single: narrow from 0 instead of from `lead` ......... 2 FAILED
drop_front_single: clamp instead of refusing an over-wide strip .. 1 FAILED
resolve_row_offsets: always build the vector .................... 1 FAILED
resolve_row_offsets: fall back to 0 instead of the shared offset . 1 FAILED
make_prompt_chunk: ignore the per-row offsets ................... 1 FAILED
make_prompt_chunk: never refuse a wrong-width offsets vector .... 1 FAILED

⚠️ One survived first drafting, and the fixture was wrong, not the mutation. the_dead_prefix_does_not_accumulate_across_steps computed the post-strip live length itself (target - lead) instead of reading it back from the slot — so it asserted my arithmetic against my arithmetic and passed with drop_dead_prefix stubbed to a no-op. It now reads slot.current_seq_len() back and asserts the live run's content is at column 0, which is what makes the no-op, the length-preserving and the narrow-from-0 mutations all fail. This is the same trap in a fourth disguise (#95's shared comp column, #98's non-binding margin, #100's invisible uniform dispatch): the fixture could not reach the condition the code was written for.

🔴 The honest gap — PerSequence is still refused, and the refusal moved AGAIN

This does not move the B=128 number. With the target's time base now per-row end to end, a fourth blocker is visible that none of #92, #95 or #100 could see: the MTP draft chain is a second, independent time base, and it is still keyed on one scalar. Four sites:

  1. MtpHiddenCapture::store(seqlen_offsets.first(), ..) (deepseek4.rs:3985) records row 0's absolute position for the whole capture, so extend_draft_kv_row's off describes only the leading row;
  2. propose_chain_batched(.., start_pos, ..) takes ONE start_pos and hands start_pos + i to MtpBlock::forward_step, whose pos is both the RoPE position and the draft-KV slot;
  3. batch_draft_caches refuses rows whose current_seq_len differs — under per-sequence advance that is all of them;
  4. so rows can only be grouped when they share an absolute chain start. With u_i ≡ 1 and every row at its own length that is one group per sequence: B single-row MTP-block forwards per chain step instead of <= depth + 1 batched ones.

the_draft_group_partition_collapses_when_the_rows_carry_their_own_lengths counts it: 8 rows on a shared width → 4 groups; 8 rows at their own lengths → 8.

Correct and B× slower is not a mode worth granting, and granting it would hide the real fix — a per-row pos in MtpBlock::forward_step (a positions vector for RoPE plus a per-row draft-KV slot) and a per-row capture offset. That is a model change in deepseek4, which is why it is not here. It is stated as draft_chain_carries_per_sequence_positions() so the next change flips exactly one thing.

I am not claiming this is the last blocker. Four agents have now made that claim and been wrong; the useful output of each wave has been naming precisely where the refusal moved.

🔴 The GPU ask

Nothing changes a default, so this is a no-regression check on the box already serving V4. I did not run it (D15 — I never call runcrate).

mistralrs bench -m <v4> -b 1 -b 128 2>&1 | tee /tmp/a.txt
ARC_V4_XS_PER_SEQ=1 ARC_MTP_PER_SEQ_KV=1 mistralrs bench -m <v4> -b 1 -b 128 2>&1 | tee /tmp/b.txt

The one number: decode tok/s in B as a fraction of A, at B=1 and B=128; the claim is B/A == 1.00 at both, because the refusal means B still takes the cohort path. And grep "cannot honour it" /tmp/b.txt must now cite the draft chain and must contain none of XsRolling, the ragged mask, cache_len or window_ok — if it names any of them, this wave did not land.

Noticed, not shipped

  • The strip and the next step's pad are two copies where one would do. drop_dead_prefix narrows-and-contiguates, then the next front_align_batch allocates and shifts again. Giving SingleCache a live_start would fold both into one — but it touches current_data/append/k()/v() on the hot path, so it is its own change. B=1 pays neither copy (both are no-ops at lead == 0).
  • drop_front_single shrinks capacity_seq_len by the lead, so a stripped slot hits CACHE_GROW_SIZE growth marginally sooner than an unstripped one. Bounded and correct, but a capacity that only ever ratchets down is worth a look.
  • kv_advance's doc still links [Self::target_masks_ragged_batches], which does not exist — the function is the free model_masks_ragged_batches. Fixed on the one line I touched; two other references remain.

#100 named

`step_full` keyed the whole fused step on one scalar `cache_len` (the BATCHED
buffer's width, which left-alignment makes `max_j L_j`). All four of PR #100's
defects close here:

* `uncached[i]` no longer saturates to 0 — `plan_step_window` takes each row's
  own length, so every ragged tail stays inside the window;
* `toks[cache_lens[i]..]` is in range by the same invariant, closing the panic
  site the saturation was hiding;
* `seqlen_offsets` routes per row through `inputs_processor` — which is what
  feeds `dsv4_attention`'s `row_q0`;
* the per-sequence commit strips the dead prefix (`drop_dead_prefix`, the
  inverse of `front_pad_kv_cache`) so a row's recorded length is its own.

Flag-off byte-identity is structural: `resolve_row_cache_lens` and
`resolve_row_offsets` return `None` for a uniform batch, and `None` runs the
pre-change scalar code verbatim.

`KvAdvance::PerSequence` is still refused. The refusal moved to the MTP DRAFT
chain, which keys a whole group on one absolute position.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
@github-actions

github-actions Bot commented Aug 17, 2026 •

Copy link
Copy Markdown
Code Metrics Report
━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
 Language              Files        Lines         Code     Comments       Blanks
━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
 C Header                  5          305          210           52           43
 CSS                       2         1181         1036           34          111
 CUDA                     72        24328        17592         4018         2718
 Dockerfile                1           39           22            8            9
 JavaScript               16         3546         2676          482          388
 Jinja2                    7          694          656            5           33
 JSON                     74         4600         4597            0            3
 Makefile                  1            6            5            0            1
 Metal Shading Lan|       33        12224         9431         1142         1651
 PowerShell                1          300          227           30           43
 Python                  143        14830        12217          797         1816
 Shell                    23         5756         3964         1412          380
 Plain Text                4         3801            0         2479         1322
 TOML                     33         1485         1292           43          150
 YAML                      3           25           23            2            0
─────────────────────────────────────────────────────────────────────────────────
 HTML                      4         2687         2604           43           40
 |- CSS                    2          543          479           37           27
 |- JavaScript             1         1233         1215           12            6
 (Total)                             4463         4298           92           73
─────────────────────────────────────────────────────────────────────────────────
 Jupyter Notebooks         4          122           83           23           16
 |- Markdown               1           60           30           22            8
 |- Python                 1          122          113            1            8
 (Total)                              304          226           46           32
─────────────────────────────────────────────────────────────────────────────────
 Markdown                190        37932            0        29111         8821
 |- BASH                  71         1623         1192          315          116
 |- C                      2           12           12            0            0
 |- CUDA                   2           84           56           16           12
 |- JSON                  18          708          708            0            0
 |- PowerShell             1            1            1            0            0
 |- Python                23         1008          787          113          108
 |- Rust                  65         2048         1713           77          258
 |- TOML                   6          207          164            0           43
 |- YAML                   4           38           33            5            0
 (Total)                            43661         4666        29637         9358
─────────────────────────────────────────────────────────────────────────────────
 Rust                    662       309912       268369        14069        27474
 |- Markdown             477        23482          471        20234         2777
 (Total)                           333394       268840        34303        30251
━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
 Total                  1278       454942       331978        74582        48382
━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━

heydryft added a commit that referenced this pull request Aug 17, 2026
…ort it

PR #93 landed `XsRollingCache::trim_tail_to(new_base: usize)` and
`reconcile_xs_bases` on master while this branch was making `base`/`tokens`
per-row `Vec<usize>`. The 7 resulting errors are semantic, not textual:
`base` is the same identifier naming two different quantities.

  - On master `tail` is `tokens - base` wide, so `base` IS the physical left
    edge of the buffer. Divergent `base` at equal `tokens` gives physical
    widths 4 vs 132 and the `slice_set` mismatch #93 exists to fix, reachable
    on plain decode through the prefix cacher.
  - Here `tail` is `[B, W, hidden]` and end-anchored, with
    `base[i] >= tokens[i] - W`. `base` is a logical resume point decoupled
    from the buffer; `BatchSrc::of` sets `v_slack_dim: Some(1)` with
    `v_slack_at_front: true`, so the widths already agree.

So the reconciliation is required on one path and a regression on the other.
Trimming every row to `base_max` raises the shortest-reach row's rollback
floor to the batch's -- a per-row quantity replaced by a batch-wide scalar,
the same shape as the cohort min-rollback #92 removed and the padding traps
#102/#103/#104 each caught. This is the last layer still holding one.

  - `reconcile_xs_bases` takes `xs_per_seq` and returns early when set. Flag
    off it runs verbatim; #93's two mutation-proven tests pass unchanged
    (they run flag-off), with only a private-field read swapped for the
    `resumable_from()` accessor, which is `base[0]` for a single-row cache.
  - `trim_tail_to` REFUSES a multi-row cache rather than reading `base[0]`.
    A scalar `new_base` cannot describe a trim of rows at different resume
    points, and proceeding returns a right-shaped tensor -- exactly what the
    caller checks -- so it names the reason instead (D18).

`per_row_xs_bases_survive_a_batch_round_trip_and_are_not_flattened` is the
test that discriminates: flag on, two rows at divergent `base` and equal
`tokens` must batch AND each keep its own `base` through `split_row`. Both
mutations were run rather than assumed -- trimming the tensors is caught by
the width assertion, but reconciling only the logical `base` passes every
shape check and fails solely on the per-row read.

8/8 mutations caught. One survivor on the first pass (`trim_tail_to`
narrowing from the wrong end) was real: no test read tail CONTENT, because
every fixture fed zeros. `trimming_the_retained_window_drops_the_oldest_rows_not_the_newest`
feeds a ramp and closes it.

430 lib tests pass, 14 synthetic, workspace check clean, scoped clippy exits
0, zero new clippy diagnostics in the touched files, and rustfmt drift is 6
before and 6 after -- no upstream reformat churn.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
@heydryft
heydryft changed the base branch from feat/v4-ragged-mask-wire to master August 17, 2026 22:54
@heydryft
heydryft merged commit 0183947 into master Aug 17, 2026
17 checks passed
@heydryft
heydryft deleted the feat/mtp-step-per-seq-cachelen branch August 17, 2026 23:42
@heydryft

Copy link
Copy Markdown
Contributor Author

Admin-merged without review (no second reviewer). Bypass covered the review requirement only — the required CI complete context was genuinely green, 0 failures, 0 cancelled, after absorbing master.

Merged without --delete-branch; #103 was retargeted to master first, then feat/mtp-step-per-seq-cachelen was deleted (D20). That ordering exists because doing it the other way earlier tonight auto-closed a downstream PR — GitHub closes a PR whose base ref disappears rather than retargeting it, and then refuses to reopen it while the ref is missing.

heydryft added a commit that referenced this pull request Aug 18, 2026
Master moved 10+ commits since 120f348, including #100 (v4 ragged mask wire)
and #102 (per-row cache_len in the fused MTP step) — adjacent territory. Merged
now rather than letting the divergence be discovered later.

Clean auto-merge; the only overlapping file is `kv_cache/mod.rs`. 466/466 green
on the merged tree (441 mine + master's new tests).

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
heydryft added a commit that referenced this pull request Aug 18, 2026
feat(mtp): per-row positions in the MTP draft chain — the blocker #102 named
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant