Skip to content

Commit c061df1

Browse files
am17anggerganov
andauthored
Qwen4Exp: add MTP (#29761)
* Qwen4Exp: add MTP * remove has_state member, check via ctx_bufs being non-empty * consistent naming + less verbose comments * cont : clean-up recurrent memory * cont : clean-up comments --------- Co-authored-by: Georgi Gerganov <ggerganov@gmail.com>
1 parent 66e0c17 commit c061df1

13 files changed

Lines changed: 289 additions & 55 deletions

‎common/speculative.cpp‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -2539,7 +2539,7 @@ common_speculative_init_result::common_speculative_init_result(
25392539
model_path = params.speculative.draft.mparams.path;
25402540
LOG_INF("%s: loading draft model '%s'\n", __func__, model_path.c_str());
25412541

2542-
llama_model * model_dft = llama_model_load_from_file(params.model.path.c_str(), mparams);
2542+
llama_model * model_dft = llama_model_load_from_file(model_path.c_str(), mparams);
25432543
if (model_dft == NULL) {
25442544
LOG_ERR("%s: failed to load draft model, '%s'\n", __func__, model_path.c_str());
25452545
return;

‎conversion/qwen4exp.py‎

Lines changed: 36 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -25,15 +25,34 @@ class Qwen4ExpTextModel(_Qwen35MRopeMixin, _LinearAttentionVReorderBase):
2525

2626
model_arch = gguf.MODEL_ARCH.QWEN4EXP
2727

28-
# the MTP block is a separate draft head; vLLM drops it too
29-
supports_mtp_export = False
30-
no_mtp = True
28+
# the MTP head: one full-attention QSA block after the trunk, fed by the trunk's hc-wide residual
29+
supports_mtp_export = True
30+
31+
# MTP tensors the shared Qwen remapper does not know
32+
_MTP_EXTRA = {
33+
"fc_embedding": "nextn_fc_embedding",
34+
"fc_hidden": "nextn_fc_hidden",
35+
"hyper_connection_mixer": "nextn_hc_head",
36+
}
3137

3238
def __init__(self, *args, **kwargs):
3339
super().__init__(*args, **kwargs)
3440
# only the shard names, so the table itself is never held
3541
self._ple_shards: dict[int, str] = {}
3642
self._ple_row_dim: int | None = None
43+
self._mtp_fc: dict[str, Tensor] = {}
44+
45+
@classmethod
46+
def filter_tensors(cls, item):
47+
name, gen = item
48+
part = name.split(".")[1] if name.startswith("mtp.") else None
49+
if part in cls._MTP_EXTRA:
50+
if cls.no_mtp:
51+
return None
52+
assert cls._original_block_count is not None
53+
rest = name.split(".", 2)[2]
54+
return f"model.layers.{cls._original_block_count}.{cls._MTP_EXTRA[part]}.{rest}", gen
55+
return super().filter_tensors(item)
3756

3857
def _read_hash_constants(self, suffix: str) -> list[int]:
3958
"""Read an int64 PLE constant straight from the checkpoint.
@@ -63,14 +82,17 @@ def set_gguf_parameters(self):
6382
self.gguf_writer.add_indexer_top_k(hp["indexer_budget"])
6483
ratio = hp["indexer_compress_ratio"]
6584
layer_types = hp["layer_types"]
85+
# the MTP block is a full-attention QSA layer too
6686
self.gguf_writer.add_attention_compress_ratios(
6787
[ratio if layer_types[i] == "full_attention" else 0 for i in range(n_layer)]
88+
+ [ratio] * (self.block_count - n_layer)
6889
)
6990

7091
# ple_layer_ids is 1-based in the HF config; empty means no n-gram table,
7192
# so emit no PLE keys rather than optional ones
93+
# the MTP head never reads PLE, so an MTP-only file carries none of it
7294
ple_layers = [i - 1 for i in hp["ple_layer_ids"]]
73-
if not ple_layers:
95+
if not ple_layers or self.mtp_only:
7496
return
7597
self.gguf_writer.add_ple_layers(ple_layers)
7698
self.gguf_writer.add_ple_ngram_size(hp["ngram_size"])
@@ -120,6 +142,14 @@ def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iter
120142
if ".ngram_embedding.shard_" in name:
121143
return self._place_ple_shard(data_torch, name)
122144

145+
# eh_proj([e ; h_s]) = fc_embedding(e) + fc_hidden(h_s) for every hc stream s
146+
if name.endswith((".nextn_fc_embedding.weight", ".nextn_fc_hidden.weight")):
147+
self._mtp_fc[name.rsplit(".", 2)[1]] = data_torch
148+
if len(self._mtp_fc) < 2:
149+
return []
150+
eh = torch.cat([self._mtp_fc.pop("nextn_fc_embedding"), self._mtp_fc.pop("nextn_fc_hidden")], dim=1)
151+
return [(self.format_tensor_name(gguf.MODEL_TENSOR.NEXTN_EH_PROJ, bid, ".weight"), eh)]
152+
123153
# one projection feeds indexer q and k; split it, as minimax-m3 does
124154
if ".indexer.index_qk_proj.weight" in name:
125155
n_q = self.hparams["indexer_n_heads"] * self.hparams["indexer_head_dim"]
@@ -182,6 +212,8 @@ def load() -> np.ndarray:
182212

183213
def prepare_tensors(self):
184214
super().prepare_tensors()
215+
if self._mtp_fc:
216+
raise ValueError(f"MTP projection missing its other half: {sorted(self._mtp_fc)}")
185217
n_parts = self.hparams.get("split_ngram_parts", 0)
186218
if self._ple_shards and len(self._ple_shards) != n_parts:
187219
raise ValueError(

‎gguf-py/gguf/constants.py‎

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1205,6 +1205,9 @@ class MODEL_TENSOR(IntEnum):
12051205
NEXTN_HNORM = auto()
12061206
NEXTN_SHARED_HEAD_HEAD = auto()
12071207
NEXTN_SHARED_HEAD_NORM = auto()
1208+
NEXTN_HC_HEAD_NORM = auto()
1209+
NEXTN_HC_HEAD_DOWN = auto()
1210+
NEXTN_HC_HEAD_UP = auto()
12081211
# eagle3
12091212
FC = auto() # feature fusion layer
12101213
D2T = auto() # draft to target vocabulary mapping
@@ -1995,6 +1998,9 @@ class MODEL_TENSOR(IntEnum):
19951998
MODEL_TENSOR.NEXTN_HNORM: "blk.{bid}.nextn.hnorm",
19961999
MODEL_TENSOR.NEXTN_SHARED_HEAD_HEAD: "blk.{bid}.nextn.shared_head_head",
19972000
MODEL_TENSOR.NEXTN_SHARED_HEAD_NORM: "blk.{bid}.nextn.shared_head_norm",
2001+
MODEL_TENSOR.NEXTN_HC_HEAD_NORM: "blk.{bid}.nextn.hc_head_norm",
2002+
MODEL_TENSOR.NEXTN_HC_HEAD_DOWN: "blk.{bid}.nextn.hc_head_down",
2003+
MODEL_TENSOR.NEXTN_HC_HEAD_UP: "blk.{bid}.nextn.hc_head_up",
19982004
MODEL_TENSOR.FC: "fc",
19992005
MODEL_TENSOR.DSPARK_MARKOV_W1: "markov_w1",
20002006
MODEL_TENSOR.DSPARK_MARKOV_W2: "markov_w2",
@@ -2992,6 +2998,13 @@ class MODEL_TENSOR(IntEnum):
29922998
MODEL_TENSOR.PLE_NORM_QUERY,
29932999
MODEL_TENSOR.PLE_NORM_CONV,
29943000
MODEL_TENSOR.PLE_CONV1D,
3001+
# MTP block: [fc_embedding | fc_hidden] as eh_proj, its own hyper-connection mixer as the head
3002+
MODEL_TENSOR.NEXTN_EH_PROJ,
3003+
MODEL_TENSOR.NEXTN_ENORM,
3004+
MODEL_TENSOR.NEXTN_HNORM,
3005+
MODEL_TENSOR.NEXTN_HC_HEAD_NORM,
3006+
MODEL_TENSOR.NEXTN_HC_HEAD_DOWN,
3007+
MODEL_TENSOR.NEXTN_HC_HEAD_UP,
29953008
],
29963009
MODEL_ARCH.PLAMO: [
29973010
MODEL_TENSOR.TOKEN_EMBD,

‎gguf-py/gguf/tensor_mapping.py‎

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2817,6 +2817,16 @@ class TensorNameMap:
28172817
MODEL_TENSOR.HC_HEAD_UP: (
28182818
"model.hyper_connection_mixer.input_mix_weight_up",
28192819
),
2820+
# the MTP block's own mixer, renamed to its layer by the converter
2821+
MODEL_TENSOR.NEXTN_HC_HEAD_NORM: (
2822+
"model.layers.{bid}.nextn_hc_head.hc_norm",
2823+
),
2824+
MODEL_TENSOR.NEXTN_HC_HEAD_DOWN: (
2825+
"model.layers.{bid}.nextn_hc_head.input_mix_weight_down",
2826+
),
2827+
MODEL_TENSOR.NEXTN_HC_HEAD_UP: (
2828+
"model.layers.{bid}.nextn_hc_head.input_mix_weight_up",
2829+
),
28202830
MODEL_TENSOR.INDEXER_Q_NORM: (
28212831
"model.layers.{bid}.self_attn.indexer.q_layernorm",
28222832
),

‎src/llama-arch.cpp‎

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -591,6 +591,9 @@ static const std::map<llm_tensor, const char *> LLM_TENSOR_NAMES = {
591591
{ LLM_TENSOR_NEXTN_HNORM, "blk.%d.nextn.hnorm" },
592592
{ LLM_TENSOR_NEXTN_SHARED_HEAD_HEAD, "blk.%d.nextn.shared_head_head" },
593593
{ LLM_TENSOR_NEXTN_SHARED_HEAD_NORM, "blk.%d.nextn.shared_head_norm" },
594+
{ LLM_TENSOR_NEXTN_HC_HEAD_NORM, "blk.%d.nextn.hc_head_norm" },
595+
{ LLM_TENSOR_NEXTN_HC_HEAD_DOWN, "blk.%d.nextn.hc_head_down" },
596+
{ LLM_TENSOR_NEXTN_HC_HEAD_UP, "blk.%d.nextn.hc_head_up" },
594597
{ LLM_TENSOR_ATTN_SUB_NORM, "blk.%d.attn_sub_norm" },
595598
{ LLM_TENSOR_FFN_SUB_NORM, "blk.%d.ffn_sub_norm" },
596599
{ LLM_TENSOR_DEC_OUTPUT_NORM, "dec.output_norm" },
@@ -985,6 +988,9 @@ static const std::map<llm_tensor, llm_tensor_info> LLM_TENSOR_INFOS = {
985988
{LLM_TENSOR_NEXTN_HNORM, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL}},
986989
{LLM_TENSOR_NEXTN_SHARED_HEAD_HEAD, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}},
987990
{LLM_TENSOR_NEXTN_SHARED_HEAD_NORM, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL}},
991+
{LLM_TENSOR_NEXTN_HC_HEAD_NORM, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL}},
992+
{LLM_TENSOR_NEXTN_HC_HEAD_DOWN, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}},
993+
{LLM_TENSOR_NEXTN_HC_HEAD_UP, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}},
988994
// Nemotron 3 Super
989995
// latent projections feed ggml_mul_mat, the buft probe must use MUL_MAT to keep them on GPU
990996
{LLM_TENSOR_FFN_LATENT_DOWN, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}},

‎src/llama-arch.h‎

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -704,6 +704,9 @@ enum llm_tensor {
704704
LLM_TENSOR_NEXTN_HNORM,
705705
LLM_TENSOR_NEXTN_SHARED_HEAD_HEAD,
706706
LLM_TENSOR_NEXTN_SHARED_HEAD_NORM,
707+
LLM_TENSOR_NEXTN_HC_HEAD_NORM, // qwen4exp: the MTP block's own hyper-connection mixer
708+
LLM_TENSOR_NEXTN_HC_HEAD_DOWN,
709+
LLM_TENSOR_NEXTN_HC_HEAD_UP,
707710
LLM_TENSOR_MASKED_EMBD_CENTROIDS,
708711
LLM_TENSOR_MASKED_EMBD_ORDERING,
709712
LLM_TENSOR_HRM_Z_L_INIT,

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

Lines changed: 0 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -14,8 +14,6 @@
1414
// llama_memory_hybrid plus a third cache with one indexer key per token, for block-sparse attention (qwen4exp QSA)
1515
// the indexer is a side buffer over the attention cells: same size, padding, streams and slots, so cell j is one token in both
1616

17-
// TODO: this memory module is pending complete reimplementation - do not use for model other than Qwen4
18-
1917
class llama_memory_hybrid_idx : public llama_memory_hybrid {
2018
public:
2119
llama_memory_hybrid_idx(

‎src/llama-memory-recurrent.cpp‎

Lines changed: 17 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -125,6 +125,13 @@ llama_memory_recurrent::llama_memory_recurrent(
125125
ctxs_bufs.emplace_back(std::move(ctx), buf);
126126
}
127127

128+
if (is_empty()) {
129+
if (n_rs_seq > 0) {
130+
n_rs_seq = 0;
131+
LLAMA_LOG_INFO("%s: disabling rollback snapshots because the memory module is empty\n", __func__);
132+
}
133+
}
134+
128135
{
129136
const size_t memory_size_r = size_r_bytes();
130137
const size_t memory_size_s = size_s_bytes();
@@ -192,6 +199,11 @@ bool llama_memory_recurrent::seq_rm(llama_seq_id seq_id, llama_pos p0, llama_pos
192199

193200
// partial rollback via per-token snapshot index (bounded by n_rs_seq)
194201
if (0 < p0 && p0 <= cell.pos && p1 > cell.pos) {
202+
// the filter kept no layer (e.g. an MTP draft context), so only the position moves back
203+
if (is_empty()) {
204+
cell.pos = p0 - 1;
205+
return true;
206+
}
195207
const llama_pos rollback = cell.pos - (p0 - 1);
196208
// pending rollback is single-use
197209
const bool pending = rs_idx[seq_id] != 0;
@@ -718,6 +730,11 @@ bool llama_memory_recurrent::get_can_shift() const {
718730
return true;
719731
}
720732

733+
bool llama_memory_recurrent::is_empty() const {
734+
assert(total_size() == 0);
735+
return ctxs_bufs.empty();
736+
}
737+
721738
size_t llama_memory_recurrent::total_size() const {
722739
size_t size = 0;
723740
for (const auto & [_, buf] : ctxs_bufs) {

‎src/llama-memory-recurrent.h‎

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -123,6 +123,9 @@ class llama_memory_recurrent : public llama_memory_i {
123123
// ggml contexts for the KV cache along with the allocated backend buffers:
124124
std::vector<std::pair<ggml_context_ptr, ggml_backend_buffer_ptr>> ctxs_bufs;
125125

126+
// true if no layers - can happen if the layer filter removes all layers
127+
bool is_empty() const;
128+
126129
size_t total_size() const;
127130

128131
size_t size_r_bytes() const;

‎src/llama-model.cpp‎

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2709,6 +2709,15 @@ llama_memory_i * llama_model::create_memory(const llama_memory_params & params,
27092709
return il < hparams.n_layer() && !hparams.is_recr(il);
27102710
};
27112711
}
2712+
2713+
// the MTP draft context holds the MTP block alone: its attention and indexer, no recurrent layer
2714+
if (arch == LLM_ARCH_QWEN4EXP && params.ctx_type == LLAMA_CONTEXT_TYPE_MTP) {
2715+
filter_attn = [&](uint32_t il) { return il >= hparams.n_layer(); };
2716+
filter_recr = [&](uint32_t) { return false; };
2717+
if (filter_idx) {
2718+
filter_idx = [&](uint32_t il) { return il >= hparams.n_layer(); };
2719+
}
2720+
}
27122721
}
27132722

27142723
if (hparams.swa_type != LLAMA_SWA_TYPE_NONE) {

0 commit comments

Comments
 (0)