Skip to content

spec: MTP no longer bridges the previous h-row into a fresh sequence (engine #290) - #83

Merged
bong-water-water-bong merged 1 commit into
1bit/hrx-vulkan-patchedfrom
1bit/mtp-pending-h-continuity
Oct 5, 2026
Merged

bong-water-water-bong merged 1 commit into
1bit/hrx-vulkan-patchedfrom
1bit/mtp-pending-h-continuity

Conversation

@bong-water-water-bong

Copy link
Copy Markdown

A fix for 1bit-MONSTER/engine#290.

Root cause

common_speculative_impl_draft_mtp keeps pending_h — the target hidden row
that precedes the first token of the next batch — and writes it into embl row
i_batch_beg[seq_id] of every draft batch it processes (the cross-ubatch
bridge). Row 0 of that batch has no other writer: process() fills rows
1..n_tokens-1 from the target's shifted nextn embeddings, so the bridge is what
produces row 0.

The bridge is unconditional. i_batch_beg[seq_id] is the index of the first
token of that sequence in this batch, so a fresh request at pos 0 also takes
the bridge — and pending_h at that moment still holds the previous request's
last hidden row. It is only ever initialised to zeros in the constructor.

The MTP head is therefore seeded at pos 0 with another request's state. That is
exactly what the issue measures:

The MTP draft head's first prediction already differs between requests: draft
candidate 0 at pos 0 has p = 0.862 vs 0.847.

— before any draft has been accepted, so it cannot be a verify-shape effect of
the checkpoint split.

eagle3 has the same deferred-boundary concept and guards it: it stores
pending_pos_last alongside the row and only uses it when
pending_pos_last + 1 == pos[beg]. draft-mtp carries no position, so it
cannot tell a continuation from a fresh sequence. (The issue's "places to look"
points at pending_g_last/verify_g, which are eagle3's; the comment above them
even says "MTP doesn't have this issue" — it is the pending_h block in
common_speculative_impl_draft_mtp that does.)

Change

Carry the position with the row:

  • pending_pos[seq_id] — the position pending_h belongs to (-1 = none),
    set wherever pending_h is written: the end of process() and accept().
  • The bridge is used only when pending_pos[seq_id] + 1 == batch_in.pos[i_batch_beg[seq_id]];
    otherwise row 0 is zeroed — the state a brand new sequence starts from.
  • verify_pos_first[seq_id] records the position of verify_h[seq_id][0], so
    accept() can derive the position of the row it selects.

Every legitimate bridge is preserved:

case pos[beg] pending_pos bridge?
fresh request 0 previous request's last pos no (zeros)
prefill chunk 2 (5+5 split) 5 4 yes
prefill → first verify batch n_past n_past-1 yes
continued cached prompt n_past n_past-1 yes

Also implement get_state/set_state for the (pending_pos, pending_h) pair,
mirroring eagle3. create_checkpoint() already stashes speculative state
(data_spec, documented as "e.g. eagle3's deferred-boundary g_embd row") and
draft-mtp returned the base-class false, so a restored checkpoint left the
bridge stale for the restored position.

Verification

  • Builds clean: Release, cpu-only llama-common, which compiles
    common/speculative.cpp and links.
  • Not verified on hardware. The repro needs Qwen3.8-27B UD-Q4_K_XL with
    mtp-Qwen3.8-27B-Q4_0.gguf on Strix Halo. Before merge:
SPECS="nockpt:bin-fix3::--ctx-checkpoints,0,default:bin-fix3:" NREQ=6  # t8.sh

Expected: the default (checkpoints on) gives 6/6 identical top-5 logprobs, same
as --ctx-checkpoints 0, and the draft's first prediction at pos 0 no longer
depends on the previous request. If it still alternates, the next probe is to
dump pending_h/pending_pos per request at -lv 5.

…ml-org#290)

common_speculative_impl_draft_mtp keeps pending_h, the target hidden row that
precedes the first token of the next batch, and writes it into embl row
i_batch_beg[seq_id] of every draft batch it processes. Row 0 of that batch has
no other writer: process() fills rows 1..n_tokens-1 from the target's shifted
nextn embeddings, so the bridge is what produces row 0.

The bridge is unconditional. i_batch_beg[seq_id] is the first token of the
sequence in *this* batch, so a fresh request at pos 0 also takes the bridge,
and pending_h at that moment still holds the previous request's last hidden row
- it is only initialised to zeros in the constructor. The MTP head is therefore
seeded at pos 0 with another request's state, which is exactly what engine ggml-org#290
reports: "draft candidate 0 at pos 0 has p = 0.862 vs 0.847", before any draft
has been accepted.

eagle3 has the same deferred-boundary concept and guards it: it stores
pending_pos_last with the row and only uses it when
"pending_pos_last + 1 == pos[beg]" (common/speculative.cpp). draft-mtp carries
no position, so it cannot tell a continuation from a fresh sequence.

Carry the position with the row:

* pending_pos[seq_id] tracks the position pending_h belongs to (-1 = none).
  Set wherever pending_h is written: the end of process() and accept().
* The bridge is used only when pending_pos[seq_id] + 1 == pos of this batch's
  first row for that sequence; otherwise row 0 is zeroed, which is the state a
  brand new sequence starts from.
* verify_pos_first[seq_id] records the position of verify_h[seq_id][0] so
  accept() can derive the position of the row it selects.

This keeps every legitimate bridge: within a prefill split into ubatches
(chunk 2 at pos 5 follows pending_pos 4) and from prefill into the first verify
batch, while a new request at pos 0 no longer inherits the previous one.

Also implement get_state/set_state for the pair, mirroring eagle3. Checkpoints
already stash the speculative state (tools/server/server-context.cpp,
"stash the draft's speculative state with the checkpoint", data_spec) and the
field is even documented as "e.g. eagle3's deferred-boundary g_embd row"; mtp
returned the base-class false, so a restored checkpoint left pending_h/pending_pos
stale for the restored position.

Verified: builds clean (Release, cpu-only llama-common, which compiles
common/speculative.cpp and links). Not verified on hardware - the ggml-org#290 repro
needs Qwen3.8-27B with the MTP head on Strix Halo.
@bong-water-water-bong

Copy link
Copy Markdown
Author

Review (orchestrator, PR-review duty while coder-llm is down)

Read against f5b7f4a. The root cause is right: in draft_mtp, the bridge set_h(i_batch_beg, pending_h) was unconditional, and pending_h is only zeroed in the constructor. A fresh request at pos 0 is therefore seeded with the previous request's last row. That matches the pos-0 draft-probability split in ggml-org#290.

Checked:

  • pending_pos is set at both writers: the end of process() uses pos[i_batch_end], and accept() uses verify_pos_first + i_h. The guard pending_pos + 1 == pos[beg] keeps every legitimate bridge in the table.
  • get_state/set_state mirror eagle3's layout. common_speculative_set_state broadcasts to every impl, but the size check (n_embd vs eagle3's n_embd_dec) keeps them apart.
  • When a checkpoint has no stash (empty data_spec), set_state leaves the current pending_pos in place. After a restore, that position is ahead of pos_next, so the guard zeroes row 0 instead of bridging a stale row. That's safe.

Unlike eagle3, get_state is not gated on recurrent/hybrid targets, so it stashes one n_embd row (about 20 KiB) per checkpoint on any target. That's harmless, but say so if you mean it.

Verdict: approve, pending the t8.sh run in the description: default checkpoints give 6/6 identical top-5 logprobs, and the pos-0 draft no longer depends on the previous request. Queued for the post-release merge window. strixhalo is reserved from Sat 10:00 to Sun 10:00 ADT.

@bong-water-water-bong

bong-water-water-bong commented Oct 3, 2026 •

Copy link
Copy Markdown
Author

Verified on strixhalo (gfx1151, balanced power mode 85 W).

Setup:

Results (md5 of each request's top-5 logprob list):

config run 1 run 2 (fresh server)
base, checkpoints on (default) alternates 81f2f08e / 594fbe14; r1 vs r2 first diff at step 15, max |dlogprob| 0.0124, ids change at step 53 same alternation
this PR, checkpoints on 6/6 identical (81f2f08e) 6/6 identical (81f2f08e)
this PR, --ctx-checkpoints 0 6/6 identical (54116382)
base, --ctx-checkpoints 0 6/6 identical (54116382)

With checkpoints on, the PR gives the value of the base's first request, which starts from a zeroed pending_h. The odd/even alternation is gone.

Checkpoints on and --ctx-checkpoints 0 give two different, stable results. With checkpoints on, the 10-token prompt is processed as 5 + 5, so the batch shapes differ. That holds on both binaries, so this PR does not cause it.

MTP decode speed is unchanged. Per-request tok/s, median of 6 (r1 includes warm-up):

config median tok/s drafts accepted
base, checkpoints on 14.08 30/32 and 31/33
this PR, checkpoints on 14.06 (run 2: 14.07) 30/32
this PR, --ctx-checkpoints 0 14.47
base, --ctx-checkpoints 0 14.47

@bong-water-water-bong
bong-water-water-bong merged commit 86eee89 into 1bit/hrx-vulkan-patched Oct 5, 2026
9 of 24 checks passed
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