Skip to content

Commit 32af709

Browse files
apepojkenclaude
andcommitted
qwen4exp: fix speculative decoding and add native MTP (NextN) support
Speculative decoding on Qwen3.8-Flash-Next previously ran ~6x SLOWER than plain decode. This branch fixes seven issues: the recurrent rollback ring was not enabled for this arch (every spec step serialized the full recurrent state through host RAM), the ring was granted only to MTP/EAGLE3-style types, the qwen4exp conv-state write ignored the ring banks, the PLE n-gram history could not rewind losslessly, the rejection path never restored a draft context that cannot roll back, prompt-cache reuse desynced such drafts, and --spec-draft-model with draft-mtp loaded the main model instead of the given sidecar file. It also implements the model's native MTP head end to end: an --mtp converter export (mtp.* tensors -> an mtp- sidecar GGUF) and the draft-mtp graph. The combiner keeps the hyper-connection residual streams distinct, which is what the head was trained on (draft acceptance 0.87 vs 0.47 when collapsed). Measured on Strix Halo (Ryzen AI Max+ 395, Vulkan/RADV), UD-Q3_K_XL target, greedy: decode 23.6 -> 49.1 t/s (2.08x) at shallow context and 14.0 -> 31.5 t/s (2.24x) at 7.3k tokens, acceptance 0.785 at --spec-draft-n-max 6. Correctness verified with a greedy identity oracle (spec output equals plain decode, modulo replay FP-reduction-order noise). Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
1 parent 6c5afc8 commit 32af709

16 files changed

Lines changed: 636 additions & 77 deletions

‎common/common.cpp‎

Lines changed: 8 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1647,7 +1647,14 @@ void common_memory::init(llama_context * ctx_tgt, llama_context * ctx_dft) {
16471647
void common_memory::seq_rm(llama_seq_id seq_id, llama_pos p0, llama_pos p1) const {
16481648
common_context_seq_rm(ctx_tgt, seq_id, p0, p1);
16491649
if (ctx_dft) {
1650-
common_context_seq_rm(ctx_dft, seq_id, p0, p1);
1650+
// mixed-capability pair: the target may support partial removal (e.g. via the
1651+
// recurrent rollback ring) while a recurrent draft model does not. the draft cache
1652+
// is a pure optimization, so instead of aborting, drop the whole draft sequence and
1653+
// let the driver re-prime it from the prompt on the next draft call.
1654+
auto * mem = llama_get_memory(ctx_dft);
1655+
if (!llama_memory_seq_rm(mem, seq_id, p0, p1)) {
1656+
llama_memory_seq_rm(mem, seq_id, -1, -1);
1657+
}
16511658
}
16521659
}
16531660

‎common/common.h‎

Lines changed: 11 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -384,11 +384,20 @@ struct common_params_speculative {
384384
}
385385

386386
uint32_t need_n_rs_seq() const {
387+
// every speculative type rolls back the target's rejected draft suffix, so any of
388+
// them benefits from the recurrent-state snapshot ring on rollback-capable archs;
389+
// without it, recurrent/hybrid models fall back to full per-step state
390+
// serialization through host memory (SEQ_RM_TYPE_FULL), which is catastrophically
391+
// slow (measured ~750 ms/step on qwen4exp)
387392
bool needs_rs_seq = std::any_of(types.begin(), types.end(), [&](auto t) {
388-
return t == COMMON_SPECULATIVE_TYPE_DRAFT_MTP || t == COMMON_SPECULATIVE_TYPE_DRAFT_EAGLE3 || t == COMMON_SPECULATIVE_TYPE_DRAFT_DFLASH || t == COMMON_SPECULATIVE_TYPE_DRAFT_DSPARK;
393+
return t != COMMON_SPECULATIVE_TYPE_NONE;
389394
});
390395

391-
return needs_rs_seq ? draft.n_max : 0u;
396+
// n_max + 1: a verify batch holds the previously sampled token plus up to n_max
397+
// drafts, and the checkpoint+replay path (used when the draft context cannot roll
398+
// back) rewinds the target across the whole batch, one deeper than the rejected
399+
// draft suffix alone
400+
return needs_rs_seq ? draft.n_max + 1 : 0u;
392401
}
393402
};
394403

‎common/speculative.cpp‎

Lines changed: 12 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -2391,8 +2391,14 @@ common_speculative_init_result::common_speculative_init_result(
23912391
// the draft context holds as many tokens per sequence as the target context
23922392
cparams.n_ctx = llama_n_ctx(ctx_tgt);
23932393

2394-
// note: for small models maybe we can set this to the maximum possible draft from all speculative types
2395-
// the extra memory for small models is likely negligible?
2394+
// note: do NOT give the draft context a rollback ring (n_rs_seq > 0). the draft decodes
2395+
// one token per ubatch while drafting, and the delta-net snapshot store clamps banks
2396+
// older than the current ubatch to the batch-start state, so a rollback deeper than one
2397+
// token restores a state that is too new. the draft then drafts from a desynced prefix:
2398+
// measured on a qwen35 4B target + 2B draft pair, acceptance fell to 0.343 with the ring
2399+
// vs 0.84 with the (slower, serialize-per-step) FULL path, while outputs stayed exact
2400+
// because verification is target-authoritative. the target is unaffected: its rollbacks
2401+
// always land inside its own n_draft+1-token verify batch, whose banks are all fresh.
23962402
cparams.n_rs_seq = 0;
23972403
cparams.ctx_other = ctx_tgt;
23982404

@@ -2401,7 +2407,10 @@ common_speculative_init_result::common_speculative_init_result(
24012407
model_path = params.speculative.draft.mparams.path;
24022408
LOG_INF("%s: loading draft model '%s'\n", __func__, model_path.c_str());
24032409

2404-
llama_model * model_dft = llama_model_load_from_file(params.model.path.c_str(), mparams);
2410+
// load the draft path that was just logged - loading params.model.path here loaded
2411+
// the main model a second time, which only worked by accident for models whose
2412+
// nextn tensors live inside the main GGUF, and never for mtp- sidecar files
2413+
llama_model * model_dft = llama_model_load_from_file(model_path.c_str(), mparams);
24052414
if (model_dft == NULL) {
24062415
LOG_ERR("%s: failed to load draft model, '%s'\n", __func__, model_path.c_str());
24072416
return;

‎conversion/qwen4exp.py‎

Lines changed: 29 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -26,9 +26,9 @@ class Qwen4ExpTextModel(_Qwen35MRopeMixin, _LinearAttentionVReorderBase):
2626

2727
model_arch = gguf.MODEL_ARCH.QWEN4EXP
2828

29-
# the MTP block is a separate draft head; vLLM drops it too
30-
supports_mtp_export = False
31-
no_mtp = True
29+
# the MTP block loads through the shared _QwenMtpMixin remap; the qwen4exp-specific
30+
# glue (fc_embedding/fc_hidden and the head's hyper-connection mixer) is rewritten in
31+
# filter_tensors below
3232

3333
def __init__(self, *args, **kwargs):
3434
super().__init__(*args, **kwargs)
@@ -40,6 +40,25 @@ def __init__(self, *args, **kwargs):
4040
self._ple_map: np.memmap | None = None
4141
self._ple_path: Path | None = None
4242

43+
@classmethod
44+
def filter_tensors(cls, item):
45+
name = item[0]
46+
if name.startswith("model.mtp."):
47+
name = name.replace("model.", "", 1)
48+
item = (name, item[1])
49+
if name.startswith("mtp.") and not cls.no_mtp:
50+
obc = cls._original_block_count
51+
parts = name.split(".")
52+
# separate embedding/hidden projections in place of qwen35's fused eh_proj
53+
if len(parts) == 3 and parts[1] == "fc_embedding":
54+
return f"model.layers.{obc}.fc_embd_mtp.weight", item[1]
55+
if len(parts) == 3 and parts[1] == "fc_hidden":
56+
return f"model.layers.{obc}.fc_hidden_mtp.weight", item[1]
57+
# the head's own hyper-connection mixer in place of shared_head.norm
58+
if len(parts) == 4 and parts[1] == "hyper_connection_mixer":
59+
return f"model.layers.{obc}.hc_mixer_mtp.{parts[2]}.{parts[3]}", item[1]
60+
return super().filter_tensors(item)
61+
4362
def _read_hash_constants(self, suffix: str) -> list[int]:
4463
"""Read an int64 PLE constant straight from the checkpoint.
4564
@@ -68,14 +87,17 @@ def set_gguf_parameters(self):
6887
self.gguf_writer.add_indexer_top_k(hp["indexer_budget"])
6988
ratio = hp["indexer_compress_ratio"]
7089
layer_types = hp["layer_types"]
71-
self.gguf_writer.add_attention_compress_ratios(
72-
[ratio if layer_types[i] == "full_attention" else 0 for i in range(n_layer)]
73-
)
90+
# the MTP block(s) beyond the trunk run dense: pad the per-layer array to block_count
91+
ratios = [ratio if layer_types[i] == "full_attention" else 0 for i in range(n_layer)]
92+
ratios += [0] * (self.block_count - n_layer)
93+
self.gguf_writer.add_attention_compress_ratios(ratios)
7494

7595
# ple_layer_ids is 1-based in the HF config; empty means no n-gram table,
7696
# so emit no PLE keys rather than optional ones
7797
ple_layers = [i - 1 for i in hp["ple_layer_ids"]]
78-
if not ple_layers:
98+
# an mtp- sidecar ships no n-gram table and its filter drops the PLE constants,
99+
# so emit no PLE keys at all: the loader then treats the file as PLE-free
100+
if not ple_layers or self.mtp_only:
79101
return
80102
self.gguf_writer.add_ple_layers(ple_layers)
81103
self.gguf_writer.add_ple_ngram_size(hp["ngram_size"])

‎gguf-py/gguf/constants.py‎

Lines changed: 17 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1170,6 +1170,11 @@ class MODEL_TENSOR(IntEnum):
11701170
NEXTN_HNORM = auto()
11711171
NEXTN_SHARED_HEAD_HEAD = auto()
11721172
NEXTN_SHARED_HEAD_NORM = auto()
1173+
NEXTN_FC_EMBD = auto()
1174+
NEXTN_FC_HIDDEN = auto()
1175+
NEXTN_HC_NORM = auto()
1176+
NEXTN_HC_DOWN = auto()
1177+
NEXTN_HC_UP = auto()
11731178
# eagle3
11741179
FC = auto() # feature fusion layer
11751180
D2T = auto() # draft to target vocabulary mapping
@@ -1940,6 +1945,11 @@ class MODEL_TENSOR(IntEnum):
19401945
MODEL_TENSOR.NEXTN_HNORM: "blk.{bid}.nextn.hnorm",
19411946
MODEL_TENSOR.NEXTN_SHARED_HEAD_HEAD: "blk.{bid}.nextn.shared_head_head",
19421947
MODEL_TENSOR.NEXTN_SHARED_HEAD_NORM: "blk.{bid}.nextn.shared_head_norm",
1948+
MODEL_TENSOR.NEXTN_FC_EMBD: "blk.{bid}.nextn.fc_embd",
1949+
MODEL_TENSOR.NEXTN_FC_HIDDEN: "blk.{bid}.nextn.fc_hidden",
1950+
MODEL_TENSOR.NEXTN_HC_NORM: "blk.{bid}.nextn.hc_norm",
1951+
MODEL_TENSOR.NEXTN_HC_DOWN: "blk.{bid}.nextn.hc_down",
1952+
MODEL_TENSOR.NEXTN_HC_UP: "blk.{bid}.nextn.hc_up",
19431953
MODEL_TENSOR.FC: "fc",
19441954
MODEL_TENSOR.DSPARK_MARKOV_W1: "markov_w1",
19451955
MODEL_TENSOR.DSPARK_MARKOV_W2: "markov_w2",
@@ -2895,6 +2905,13 @@ class MODEL_TENSOR(IntEnum):
28952905
MODEL_TENSOR.PLE_NORM_QUERY,
28962906
MODEL_TENSOR.PLE_NORM_CONV,
28972907
MODEL_TENSOR.PLE_CONV1D,
2908+
MODEL_TENSOR.NEXTN_FC_EMBD,
2909+
MODEL_TENSOR.NEXTN_FC_HIDDEN,
2910+
MODEL_TENSOR.NEXTN_HC_NORM,
2911+
MODEL_TENSOR.NEXTN_HC_DOWN,
2912+
MODEL_TENSOR.NEXTN_HC_UP,
2913+
MODEL_TENSOR.NEXTN_ENORM,
2914+
MODEL_TENSOR.NEXTN_HNORM,
28982915
],
28992916
MODEL_ARCH.PLAMO: [
29002917
MODEL_TENSOR.TOKEN_EMBD,

‎gguf-py/gguf/tensor_mapping.py‎

Lines changed: 20 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2661,6 +2661,26 @@ class TensorNameMap:
26612661
"model.layers.{bid}.hnorm",
26622662
),
26632663

2664+
MODEL_TENSOR.NEXTN_FC_EMBD: (
2665+
"model.layers.{bid}.fc_embd_mtp",
2666+
),
2667+
2668+
MODEL_TENSOR.NEXTN_FC_HIDDEN: (
2669+
"model.layers.{bid}.fc_hidden_mtp",
2670+
),
2671+
2672+
MODEL_TENSOR.NEXTN_HC_NORM: (
2673+
"model.layers.{bid}.hc_mixer_mtp.hc_norm",
2674+
),
2675+
2676+
MODEL_TENSOR.NEXTN_HC_DOWN: (
2677+
"model.layers.{bid}.hc_mixer_mtp.input_mix_weight_down",
2678+
),
2679+
2680+
MODEL_TENSOR.NEXTN_HC_UP: (
2681+
"model.layers.{bid}.hc_mixer_mtp.input_mix_weight_up",
2682+
),
2683+
26642684
MODEL_TENSOR.NEXTN_SHARED_HEAD_HEAD: (
26652685
"model.layers.{bid}.shared_head.head",
26662686
),

‎src/llama-arch.cpp‎

Lines changed: 14 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -568,6 +568,11 @@ static const std::map<llm_tensor, const char *> LLM_TENSOR_NAMES = {
568568
{ LLM_TENSOR_NEXTN_HNORM, "blk.%d.nextn.hnorm" },
569569
{ LLM_TENSOR_NEXTN_SHARED_HEAD_HEAD, "blk.%d.nextn.shared_head_head" },
570570
{ LLM_TENSOR_NEXTN_SHARED_HEAD_NORM, "blk.%d.nextn.shared_head_norm" },
571+
{ LLM_TENSOR_NEXTN_FC_EMBD, "blk.%d.nextn.fc_embd" },
572+
{ LLM_TENSOR_NEXTN_FC_HIDDEN, "blk.%d.nextn.fc_hidden" },
573+
{ LLM_TENSOR_NEXTN_HC_NORM, "blk.%d.nextn.hc_norm" },
574+
{ LLM_TENSOR_NEXTN_HC_DOWN, "blk.%d.nextn.hc_down" },
575+
{ LLM_TENSOR_NEXTN_HC_UP, "blk.%d.nextn.hc_up" },
571576
{ LLM_TENSOR_ATTN_SUB_NORM, "blk.%d.attn_sub_norm" },
572577
{ LLM_TENSOR_FFN_SUB_NORM, "blk.%d.ffn_sub_norm" },
573578
{ LLM_TENSOR_DEC_OUTPUT_NORM, "dec.output_norm" },
@@ -949,6 +954,11 @@ static const std::map<llm_tensor, llm_tensor_info> LLM_TENSOR_INFOS = {
949954
{LLM_TENSOR_NEXTN_HNORM, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL}},
950955
{LLM_TENSOR_NEXTN_SHARED_HEAD_HEAD, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}},
951956
{LLM_TENSOR_NEXTN_SHARED_HEAD_NORM, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL}},
957+
{LLM_TENSOR_NEXTN_FC_EMBD, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}},
958+
{LLM_TENSOR_NEXTN_FC_HIDDEN, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}},
959+
{LLM_TENSOR_NEXTN_HC_NORM, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL}},
960+
{LLM_TENSOR_NEXTN_HC_DOWN, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}},
961+
{LLM_TENSOR_NEXTN_HC_UP, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}},
952962
// Nemotron 3 Super
953963
// latent projections feed ggml_mul_mat, the buft probe must use MUL_MAT to keep them on GPU
954964
{LLM_TENSOR_FFN_LATENT_DOWN, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}},
@@ -1080,6 +1090,10 @@ bool llm_arch_supports_rs_rollback(const llm_arch & arch) {
10801090
switch (arch) {
10811091
case LLM_ARCH_QWEN35:
10821092
case LLM_ARCH_QWEN35MOE:
1093+
// qwen4exp shares the delta-net base + build_rs path with qwen35, and its extra
1094+
// caches already mirror rollback: the indexer cache is addressed by the attention
1095+
// cells and the PLE history/conv state go through seq_rm/build_rs like the SSM state.
1096+
case LLM_ARCH_QWEN4EXP:
10831097
case LLM_ARCH_DEEPSEEK4:
10841098
case LLM_ARCH_NEMOTRON_H:
10851099
case LLM_ARCH_NEMOTRON_H_MOE:

‎src/llama-arch.h‎

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -681,6 +681,13 @@ enum llm_tensor {
681681
LLM_TENSOR_NEXTN_HNORM,
682682
LLM_TENSOR_NEXTN_SHARED_HEAD_HEAD,
683683
LLM_TENSOR_NEXTN_SHARED_HEAD_NORM,
684+
// qwen4exp MTP glue: separate embedding/hidden projections instead of eh_proj,
685+
// and the head's own hyper-connection mixer in place of shared_head_norm
686+
LLM_TENSOR_NEXTN_FC_EMBD,
687+
LLM_TENSOR_NEXTN_FC_HIDDEN,
688+
LLM_TENSOR_NEXTN_HC_NORM,
689+
LLM_TENSOR_NEXTN_HC_DOWN,
690+
LLM_TENSOR_NEXTN_HC_UP,
684691
LLM_TENSOR_MASKED_EMBD_CENTROIDS,
685692
LLM_TENSOR_MASKED_EMBD_ORDERING,
686693
LLM_TENSOR_FC,

‎src/llama-memory-hybrid-idx.cpp‎

Lines changed: 18 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -56,7 +56,13 @@ llama_memory_hybrid_idx::llama_memory_hybrid_idx(
5656
model, hparams_idx, type_k, type_v, v_trans, offload, unified,
5757
kv_size, n_seq_max, n_pad, n_swa, swa_type,
5858
nullptr, filter_idx, nullptr, nullptr, "idx_");
59-
}()) {}
59+
}()) {
60+
// keep enough PLE history that a rollback of up to n_rs_seq tokens still leaves the
61+
// full (ple_ngram_size - 1) hash window of true predecessors; see ple_hist_keep
62+
const uint32_t n_win = model.hparams.ple_ngram_size > 0 ? model.hparams.ple_ngram_size - 1 : 0;
63+
64+
ple_hist_keep_n = n_win + n_rs_seq;
65+
}
6066

6167
llama_memory_context_ptr llama_memory_hybrid_idx::init_batch(llama_batch_allocr & balloc, uint32_t n_ubatch, bool embd_all) {
6268
// note: this repeats llama_memory_hybrid::init_batch because the indexer cache needs the
@@ -437,9 +443,11 @@ void llama_memory_hybrid_idx::ple_hist_state_read(llama_io_read_i & io, llama_se
437443
io.read(&next_pos, sizeof(next_pos));
438444
io.read(&n_toks, sizeof(n_toks));
439445

440-
// the window is never longer than ple_ngram_size - 1; anything else is a corrupt or
441-
// mismatched blob, and reading it would size an allocation from the file
442-
if (n_toks > LLAMA_MAX_PLE_NGRAM - 1) {
446+
// the history holds the hash window plus up to ple_hist_keep_n rollback slack;
447+
// anything far beyond that is a corrupt or mismatched blob, and reading it would
448+
// size an allocation from the file. allow generous slack for blobs written with a
449+
// larger ring than the reader's.
450+
if (n_toks > (uint32_t) LLAMA_MAX_PLE_NGRAM - 1 + std::max<uint32_t>(ple_hist_keep_n, 64)) {
443451
throw std::runtime_error("qwen4exp PLE history: implausible token count in state blob");
444452
}
445453

@@ -573,6 +581,12 @@ llama_memory_hybrid_idx::ple_history & llama_memory_hybrid_idx_context::get_ple_
573581
return mem->ple_hist_get(seq_id);
574582
}
575583

584+
uint32_t llama_memory_hybrid_idx_context::get_ple_hist_keep() const {
585+
GGML_ASSERT(mem != nullptr);
586+
587+
return mem->ple_hist_keep_n;
588+
}
589+
576590
void llama_memory_hybrid_idx_context::set_input_qsa(
577591
ggml_tensor * cell_blk,
578592
ggml_tensor * blk_cells,

‎src/llama-memory-hybrid-idx.h‎

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -95,6 +95,12 @@ class llama_memory_hybrid_idx : public llama_memory_hybrid {
9595
// const because set_input updates it through a const memory context
9696
ple_history & ple_hist_get(llama_seq_id seq_id) const;
9797

98+
// how many history tokens to keep per sequence: (ple_ngram_size - 1) for the hash
99+
// window, plus n_rs_seq so a ring rollback can truncate without losing the true
100+
// predecessors (a shorter history EOS-pads the window after a rewind, which corrupts
101+
// the n-gram rows for the first tokens decoded after every speculative rollback)
102+
uint32_t ple_hist_keep_n = 0;
103+
98104
private:
99105
// the indexer cache holds one key head per layer, so it needs its own hparams:
100106
// llama_kv_cache keeps a reference to what it is given
@@ -161,6 +167,8 @@ class llama_memory_hybrid_idx_context : public llama_memory_hybrid_context {
161167
// [TAG_PLE_HISTORY] the per-sequence n-gram history of the owning memory, for set_input
162168
llama_memory_hybrid_idx::ple_history & get_ple_hist(llama_seq_id seq_id) const;
163169

170+
uint32_t get_ple_hist_keep() const;
171+
164172
// block-compressed sparse attention (qwen4exp QSA) over the cells of the indexer cache.
165173
// Blocks cut the position line, not the cell array, so no caller assumes a contiguous layout:
166174
// cell_blk I32 [n_kv, ns] block each cell belongs to

0 commit comments

Comments
 (0)