Repository navigation
spec: MTP no longer bridges the previous h-row into a fresh sequence (engine #290) - #83
Conversation
…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.
|
Review (orchestrator, PR-review duty while coder-llm is down) Read against f5b7f4a. The root cause is right: in Checked:
Unlike eagle3, 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. |
|
Verified on strixhalo (gfx1151, balanced power mode 85 W). Setup:
Results (md5 of each request's top-5 logprob list):
With checkpoints on, the PR gives the value of the base's first request, which starts from a zeroed Checkpoints on and MTP decode speed is unchanged. Per-request tok/s, median of 6 (r1 includes warm-up):
|
86eee89
into
1bit/hrx-vulkan-patched
A fix for 1bit-MONSTER/engine#290.
Root cause
common_speculative_impl_draft_mtpkeepspending_h— the target hidden rowthat 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-ubatchbridge). Row 0 of that batch has no other writer:
process()fills rows1..n_tokens-1from the target's shifted nextn embeddings, so the bridge is whatproduces row 0.
The bridge is unconditional.
i_batch_beg[seq_id]is the index of the firsttoken of that sequence in this batch, so a fresh request at pos 0 also takes
the bridge — and
pending_hat that moment still holds the previous request'slast 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:
— 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_lastalongside the row and only uses it whenpending_pos_last + 1 == pos[beg].draft-mtpcarries no position, so itcannot 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 themeven says "MTP doesn't have this issue" — it is the
pending_hblock incommon_speculative_impl_draft_mtpthat does.)Change
Carry the position with the row:
pending_pos[seq_id]— the positionpending_hbelongs to (-1= none),set wherever
pending_his written: the end ofprocess()andaccept().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 ofverify_h[seq_id][0], soaccept()can derive the position of the row it selects.Every legitimate bridge is preserved:
pos[beg]pending_posAlso implement
get_state/set_statefor 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") anddraft-mtpreturned the base-classfalse, so a restored checkpoint left thebridge stale for the restored position.
Verification
llama-common, which compilescommon/speculative.cppand links.mtp-Qwen3.8-27B-Q4_0.ggufon Strix Halo. Before merge: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 longerdepends on the previous request. If it still alternates, the next probe is to
dump
pending_h/pending_posper request at-lv 5.