Commit c061df198 for llama.cpp
commit c061df19838ff60970faf54fd7e414953590125d
Author: Aman Gupta <amangupta052@gmail.com>
Date: Thu Oct 1 19:13:27 2026 +0800
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>
diff --git a/common/speculative.cpp b/common/speculative.cpp
index b1244d9a6..5c36c9ca5 100644
--- a/common/speculative.cpp
+++ b/common/speculative.cpp
@@ -2539,7 +2539,7 @@ common_speculative_init_result::common_speculative_init_result(
model_path = params.speculative.draft.mparams.path;
LOG_INF("%s: loading draft model '%s'\n", __func__, model_path.c_str());
- llama_model * model_dft = llama_model_load_from_file(params.model.path.c_str(), mparams);
+ llama_model * model_dft = llama_model_load_from_file(model_path.c_str(), mparams);
if (model_dft == NULL) {
LOG_ERR("%s: failed to load draft model, '%s'\n", __func__, model_path.c_str());
return;
diff --git a/conversion/qwen4exp.py b/conversion/qwen4exp.py
index 168796d61..7ca26c2a3 100644
--- a/conversion/qwen4exp.py
+++ b/conversion/qwen4exp.py
@@ -25,15 +25,34 @@ class Qwen4ExpTextModel(_Qwen35MRopeMixin, _LinearAttentionVReorderBase):
model_arch = gguf.MODEL_ARCH.QWEN4EXP
- # the MTP block is a separate draft head; vLLM drops it too
- supports_mtp_export = False
- no_mtp = True
+ # the MTP head: one full-attention QSA block after the trunk, fed by the trunk's hc-wide residual
+ supports_mtp_export = True
+
+ # MTP tensors the shared Qwen remapper does not know
+ _MTP_EXTRA = {
+ "fc_embedding": "nextn_fc_embedding",
+ "fc_hidden": "nextn_fc_hidden",
+ "hyper_connection_mixer": "nextn_hc_head",
+ }
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
# only the shard names, so the table itself is never held
self._ple_shards: dict[int, str] = {}
self._ple_row_dim: int | None = None
+ self._mtp_fc: dict[str, Tensor] = {}
+
+ @classmethod
+ def filter_tensors(cls, item):
+ name, gen = item
+ part = name.split(".")[1] if name.startswith("mtp.") else None
+ if part in cls._MTP_EXTRA:
+ if cls.no_mtp:
+ return None
+ assert cls._original_block_count is not None
+ rest = name.split(".", 2)[2]
+ return f"model.layers.{cls._original_block_count}.{cls._MTP_EXTRA[part]}.{rest}", gen
+ return super().filter_tensors(item)
def _read_hash_constants(self, suffix: str) -> list[int]:
"""Read an int64 PLE constant straight from the checkpoint.
@@ -63,14 +82,17 @@ class Qwen4ExpTextModel(_Qwen35MRopeMixin, _LinearAttentionVReorderBase):
self.gguf_writer.add_indexer_top_k(hp["indexer_budget"])
ratio = hp["indexer_compress_ratio"]
layer_types = hp["layer_types"]
+ # the MTP block is a full-attention QSA layer too
self.gguf_writer.add_attention_compress_ratios(
[ratio if layer_types[i] == "full_attention" else 0 for i in range(n_layer)]
+ + [ratio] * (self.block_count - n_layer)
)
# ple_layer_ids is 1-based in the HF config; empty means no n-gram table,
# so emit no PLE keys rather than optional ones
+ # the MTP head never reads PLE, so an MTP-only file carries none of it
ple_layers = [i - 1 for i in hp["ple_layer_ids"]]
- if not ple_layers:
+ if not ple_layers or self.mtp_only:
return
self.gguf_writer.add_ple_layers(ple_layers)
self.gguf_writer.add_ple_ngram_size(hp["ngram_size"])
@@ -120,6 +142,14 @@ class Qwen4ExpTextModel(_Qwen35MRopeMixin, _LinearAttentionVReorderBase):
if ".ngram_embedding.shard_" in name:
return self._place_ple_shard(data_torch, name)
+ # eh_proj([e ; h_s]) = fc_embedding(e) + fc_hidden(h_s) for every hc stream s
+ if name.endswith((".nextn_fc_embedding.weight", ".nextn_fc_hidden.weight")):
+ self._mtp_fc[name.rsplit(".", 2)[1]] = data_torch
+ if len(self._mtp_fc) < 2:
+ return []
+ eh = torch.cat([self._mtp_fc.pop("nextn_fc_embedding"), self._mtp_fc.pop("nextn_fc_hidden")], dim=1)
+ return [(self.format_tensor_name(gguf.MODEL_TENSOR.NEXTN_EH_PROJ, bid, ".weight"), eh)]
+
# one projection feeds indexer q and k; split it, as minimax-m3 does
if ".indexer.index_qk_proj.weight" in name:
n_q = self.hparams["indexer_n_heads"] * self.hparams["indexer_head_dim"]
@@ -182,6 +212,8 @@ class Qwen4ExpTextModel(_Qwen35MRopeMixin, _LinearAttentionVReorderBase):
def prepare_tensors(self):
super().prepare_tensors()
+ if self._mtp_fc:
+ raise ValueError(f"MTP projection missing its other half: {sorted(self._mtp_fc)}")
n_parts = self.hparams.get("split_ngram_parts", 0)
if self._ple_shards and len(self._ple_shards) != n_parts:
raise ValueError(
diff --git a/gguf-py/gguf/constants.py b/gguf-py/gguf/constants.py
index 2daedfda9..8075e82a3 100644
--- a/gguf-py/gguf/constants.py
+++ b/gguf-py/gguf/constants.py
@@ -1205,6 +1205,9 @@ class MODEL_TENSOR(IntEnum):
NEXTN_HNORM = auto()
NEXTN_SHARED_HEAD_HEAD = auto()
NEXTN_SHARED_HEAD_NORM = auto()
+ NEXTN_HC_HEAD_NORM = auto()
+ NEXTN_HC_HEAD_DOWN = auto()
+ NEXTN_HC_HEAD_UP = auto()
# eagle3
FC = auto() # feature fusion layer
D2T = auto() # draft to target vocabulary mapping
@@ -1995,6 +1998,9 @@ TENSOR_NAMES: dict[MODEL_TENSOR, str] = {
MODEL_TENSOR.NEXTN_HNORM: "blk.{bid}.nextn.hnorm",
MODEL_TENSOR.NEXTN_SHARED_HEAD_HEAD: "blk.{bid}.nextn.shared_head_head",
MODEL_TENSOR.NEXTN_SHARED_HEAD_NORM: "blk.{bid}.nextn.shared_head_norm",
+ MODEL_TENSOR.NEXTN_HC_HEAD_NORM: "blk.{bid}.nextn.hc_head_norm",
+ MODEL_TENSOR.NEXTN_HC_HEAD_DOWN: "blk.{bid}.nextn.hc_head_down",
+ MODEL_TENSOR.NEXTN_HC_HEAD_UP: "blk.{bid}.nextn.hc_head_up",
MODEL_TENSOR.FC: "fc",
MODEL_TENSOR.DSPARK_MARKOV_W1: "markov_w1",
MODEL_TENSOR.DSPARK_MARKOV_W2: "markov_w2",
@@ -2992,6 +2998,13 @@ MODEL_TENSORS: dict[MODEL_ARCH, list[MODEL_TENSOR]] = {
MODEL_TENSOR.PLE_NORM_QUERY,
MODEL_TENSOR.PLE_NORM_CONV,
MODEL_TENSOR.PLE_CONV1D,
+ # MTP block: [fc_embedding | fc_hidden] as eh_proj, its own hyper-connection mixer as the head
+ MODEL_TENSOR.NEXTN_EH_PROJ,
+ MODEL_TENSOR.NEXTN_ENORM,
+ MODEL_TENSOR.NEXTN_HNORM,
+ MODEL_TENSOR.NEXTN_HC_HEAD_NORM,
+ MODEL_TENSOR.NEXTN_HC_HEAD_DOWN,
+ MODEL_TENSOR.NEXTN_HC_HEAD_UP,
],
MODEL_ARCH.PLAMO: [
MODEL_TENSOR.TOKEN_EMBD,
diff --git a/gguf-py/gguf/tensor_mapping.py b/gguf-py/gguf/tensor_mapping.py
index 16df24911..ed7f1c2be 100644
--- a/gguf-py/gguf/tensor_mapping.py
+++ b/gguf-py/gguf/tensor_mapping.py
@@ -2817,6 +2817,16 @@ class TensorNameMap:
MODEL_TENSOR.HC_HEAD_UP: (
"model.hyper_connection_mixer.input_mix_weight_up",
),
+ # the MTP block's own mixer, renamed to its layer by the converter
+ MODEL_TENSOR.NEXTN_HC_HEAD_NORM: (
+ "model.layers.{bid}.nextn_hc_head.hc_norm",
+ ),
+ MODEL_TENSOR.NEXTN_HC_HEAD_DOWN: (
+ "model.layers.{bid}.nextn_hc_head.input_mix_weight_down",
+ ),
+ MODEL_TENSOR.NEXTN_HC_HEAD_UP: (
+ "model.layers.{bid}.nextn_hc_head.input_mix_weight_up",
+ ),
MODEL_TENSOR.INDEXER_Q_NORM: (
"model.layers.{bid}.self_attn.indexer.q_layernorm",
),
diff --git a/src/llama-arch.cpp b/src/llama-arch.cpp
index 7ed8f8c86..cc361c7f2 100644
--- a/src/llama-arch.cpp
+++ b/src/llama-arch.cpp
@@ -591,6 +591,9 @@ static const std::map<llm_tensor, const char *> LLM_TENSOR_NAMES = {
{ LLM_TENSOR_NEXTN_HNORM, "blk.%d.nextn.hnorm" },
{ LLM_TENSOR_NEXTN_SHARED_HEAD_HEAD, "blk.%d.nextn.shared_head_head" },
{ LLM_TENSOR_NEXTN_SHARED_HEAD_NORM, "blk.%d.nextn.shared_head_norm" },
+ { LLM_TENSOR_NEXTN_HC_HEAD_NORM, "blk.%d.nextn.hc_head_norm" },
+ { LLM_TENSOR_NEXTN_HC_HEAD_DOWN, "blk.%d.nextn.hc_head_down" },
+ { LLM_TENSOR_NEXTN_HC_HEAD_UP, "blk.%d.nextn.hc_head_up" },
{ LLM_TENSOR_ATTN_SUB_NORM, "blk.%d.attn_sub_norm" },
{ LLM_TENSOR_FFN_SUB_NORM, "blk.%d.ffn_sub_norm" },
{ LLM_TENSOR_DEC_OUTPUT_NORM, "dec.output_norm" },
@@ -985,6 +988,9 @@ static const std::map<llm_tensor, llm_tensor_info> LLM_TENSOR_INFOS = {
{LLM_TENSOR_NEXTN_HNORM, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL}},
{LLM_TENSOR_NEXTN_SHARED_HEAD_HEAD, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}},
{LLM_TENSOR_NEXTN_SHARED_HEAD_NORM, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL}},
+ {LLM_TENSOR_NEXTN_HC_HEAD_NORM, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL}},
+ {LLM_TENSOR_NEXTN_HC_HEAD_DOWN, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}},
+ {LLM_TENSOR_NEXTN_HC_HEAD_UP, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}},
// Nemotron 3 Super
// latent projections feed ggml_mul_mat, the buft probe must use MUL_MAT to keep them on GPU
{LLM_TENSOR_FFN_LATENT_DOWN, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}},
diff --git a/src/llama-arch.h b/src/llama-arch.h
index 3b437840d..612b6797a 100644
--- a/src/llama-arch.h
+++ b/src/llama-arch.h
@@ -704,6 +704,9 @@ enum llm_tensor {
LLM_TENSOR_NEXTN_HNORM,
LLM_TENSOR_NEXTN_SHARED_HEAD_HEAD,
LLM_TENSOR_NEXTN_SHARED_HEAD_NORM,
+ LLM_TENSOR_NEXTN_HC_HEAD_NORM, // qwen4exp: the MTP block's own hyper-connection mixer
+ LLM_TENSOR_NEXTN_HC_HEAD_DOWN,
+ LLM_TENSOR_NEXTN_HC_HEAD_UP,
LLM_TENSOR_MASKED_EMBD_CENTROIDS,
LLM_TENSOR_MASKED_EMBD_ORDERING,
LLM_TENSOR_HRM_Z_L_INIT,
diff --git a/src/llama-memory-hybrid-idx.h b/src/llama-memory-hybrid-idx.h
index b954d9f7a..1e63a4099 100644
--- a/src/llama-memory-hybrid-idx.h
+++ b/src/llama-memory-hybrid-idx.h
@@ -14,8 +14,6 @@
// llama_memory_hybrid plus a third cache with one indexer key per token, for block-sparse attention (qwen4exp QSA)
// the indexer is a side buffer over the attention cells: same size, padding, streams and slots, so cell j is one token in both
-// TODO: this memory module is pending complete reimplementation - do not use for model other than Qwen4
-
class llama_memory_hybrid_idx : public llama_memory_hybrid {
public:
llama_memory_hybrid_idx(
diff --git a/src/llama-memory-recurrent.cpp b/src/llama-memory-recurrent.cpp
index 528c90c41..8353a3093 100644
--- a/src/llama-memory-recurrent.cpp
+++ b/src/llama-memory-recurrent.cpp
@@ -125,6 +125,13 @@ llama_memory_recurrent::llama_memory_recurrent(
ctxs_bufs.emplace_back(std::move(ctx), buf);
}
+ if (is_empty()) {
+ if (n_rs_seq > 0) {
+ n_rs_seq = 0;
+ LLAMA_LOG_INFO("%s: disabling rollback snapshots because the memory module is empty\n", __func__);
+ }
+ }
+
{
const size_t memory_size_r = size_r_bytes();
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
// partial rollback via per-token snapshot index (bounded by n_rs_seq)
if (0 < p0 && p0 <= cell.pos && p1 > cell.pos) {
+ // the filter kept no layer (e.g. an MTP draft context), so only the position moves back
+ if (is_empty()) {
+ cell.pos = p0 - 1;
+ return true;
+ }
const llama_pos rollback = cell.pos - (p0 - 1);
// pending rollback is single-use
const bool pending = rs_idx[seq_id] != 0;
@@ -718,6 +730,11 @@ bool llama_memory_recurrent::get_can_shift() const {
return true;
}
+bool llama_memory_recurrent::is_empty() const {
+ assert(total_size() == 0);
+ return ctxs_bufs.empty();
+}
+
size_t llama_memory_recurrent::total_size() const {
size_t size = 0;
for (const auto & [_, buf] : ctxs_bufs) {
diff --git a/src/llama-memory-recurrent.h b/src/llama-memory-recurrent.h
index 25ade10e5..08489a4be 100644
--- a/src/llama-memory-recurrent.h
+++ b/src/llama-memory-recurrent.h
@@ -123,6 +123,9 @@ private:
// ggml contexts for the KV cache along with the allocated backend buffers:
std::vector<std::pair<ggml_context_ptr, ggml_backend_buffer_ptr>> ctxs_bufs;
+ // true if no layers - can happen if the layer filter removes all layers
+ bool is_empty() const;
+
size_t total_size() const;
size_t size_r_bytes() const;
diff --git a/src/llama-model.cpp b/src/llama-model.cpp
index eff60eea8..404509888 100644
--- a/src/llama-model.cpp
+++ b/src/llama-model.cpp
@@ -2709,6 +2709,15 @@ llama_memory_i * llama_model::create_memory(const llama_memory_params & params,
return il < hparams.n_layer() && !hparams.is_recr(il);
};
}
+
+ // the MTP draft context holds the MTP block alone: its attention and indexer, no recurrent layer
+ if (arch == LLM_ARCH_QWEN4EXP && params.ctx_type == LLAMA_CONTEXT_TYPE_MTP) {
+ filter_attn = [&](uint32_t il) { return il >= hparams.n_layer(); };
+ filter_recr = [&](uint32_t) { return false; };
+ if (filter_idx) {
+ filter_idx = [&](uint32_t il) { return il >= hparams.n_layer(); };
+ }
+ }
}
if (hparams.swa_type != LLAMA_SWA_TYPE_NONE) {
diff --git a/src/llama-model.h b/src/llama-model.h
index 8c4e438a2..25e514869 100644
--- a/src/llama-model.h
+++ b/src/llama-model.h
@@ -233,6 +233,11 @@ struct llama_layer_nextn {
struct ggml_tensor * shared_head_head_s = nullptr;
struct ggml_tensor * shared_head_head_in_s = nullptr;
struct ggml_tensor * shared_head_norm = nullptr;
+
+ // qwen4exp: the MTP block collapses its hyper-connection streams with its own mixer
+ struct ggml_tensor * hc_head_norm = nullptr;
+ struct ggml_tensor * hc_head_down = nullptr;
+ struct ggml_tensor * hc_head_up = nullptr;
};
struct llama_layer_switch_lora {
diff --git a/src/models/models.h b/src/models/models.h
index 898d22f6b..a800d3fc0 100644
--- a/src/models/models.h
+++ b/src/models/models.h
@@ -2395,7 +2395,12 @@ struct llama_model_qwen4exp : public llama_model_base {
struct graph : public llm_build_delta_net_base {
graph(const llama_model & model, const llm_graph_params & params);
- private:
+ protected:
+ // the helpers alone, graph_mtp builds its own body
+ struct no_build {};
+ graph(const llama_model & model, const llm_graph_params & params, no_build) :
+ llm_build_delta_net_base(params), model(model) {}
+
// HC replaces every layer norm: residual is [n_embd, hc, n_tokens]
ggml_tensor * build_hc_mix(
ggml_tensor * x,
@@ -2489,6 +2494,11 @@ struct llama_model_qwen4exp : public llama_model_base {
const llama_model & model;
};
+ // MTP draft head: one QSA block after the trunk, fed by the trunk's hc-wide residual
+ struct graph_mtp : public graph {
+ graph_mtp(const llama_model & model, const llm_graph_params & params);
+ };
+
std::unique_ptr<llm_graph_context> build_arch_graph(const llm_graph_params & params) const override;
};
diff --git a/src/models/qwen4exp.cpp b/src/models/qwen4exp.cpp
index 768ade04d..ab12c1d54 100644
--- a/src/models/qwen4exp.cpp
+++ b/src/models/qwen4exp.cpp
@@ -6,9 +6,6 @@
#include <algorithm>
#include <cinttypes>
-// [TAG_QWEN4_REIMPLEMENT]
-// TODO: this graph implementation is pending complete reimplementation - do not use it as a reference
-
// bad metadata must be catchable: GGML_ASSERT aborts the whole process
static void qwen4exp_require_nonzero(const llama_model_loader & ml, llm_kv kid, uint32_t value) {
if (value == 0) {
@@ -66,7 +63,7 @@ void llama_model_qwen4exp::load_arch_hparams(llama_model_loader & ml) {
// QSA pools the indexer keys of blocks of compress_ratio cells, one block size for the whole model
hparams.indexer_kpool = 0;
- for (uint32_t il = 0; il < hparams.n_layer(); ++il) {
+ for (uint32_t il = 0; il < hparams.n_layer_all; ++il) {
const uint32_t r = hparams.dsv4_compress_ratios[il];
if (r == 0) {
continue;
@@ -178,13 +175,18 @@ void llama_model_qwen4exp::load_arch_tensors(llama_model_loader & ml) {
const int64_t hc_dim = hc * n_embd;
const int64_t hc_lr = hparams.hc_low_rank;
+ // an MTP-only file carries the MTP block, the embeddings and the LM head, but no trunk
+ const bool mtp_only = n_layer_nextn > 0 && ml.get_weight(tn(LLM_TENSOR_HC_ATTN_NORM, "weight", 0).str().c_str()) == nullptr;
+ const int trunk_flags = mtp_only ? TENSOR_NOT_REQUIRED : 0;
+ const int mtp_flags = ml.load_mtp ? 0 : TENSOR_SKIP;
+
tok_embd = create_tensor(tn(LLM_TENSOR_TOKEN_EMBD, "weight"), { n_embd, n_vocab }, 0);
// there is no output_norm: the final hyper-connection mixer carries it
// the gammas load as [n_embd, hc] so the grouped norm multiplies them without a graph reshape
- hc_head_norm = create_tensor(tn(LLM_TENSOR_HC_HEAD_NORM, "weight"), { n_embd, hc }, TENSOR_ALLOW_RESHAPE);
- hc_head_down = create_tensor(tn(LLM_TENSOR_HC_HEAD_DOWN, "weight"), { hc_dim, hc_lr }, 0);
- hc_head_up = create_tensor(tn(LLM_TENSOR_HC_HEAD_UP, "weight"), { hc_lr, hc_dim }, 0);
+ hc_head_norm = create_tensor(tn(LLM_TENSOR_HC_HEAD_NORM, "weight"), { n_embd, hc }, trunk_flags | TENSOR_ALLOW_RESHAPE);
+ hc_head_down = create_tensor(tn(LLM_TENSOR_HC_HEAD_DOWN, "weight"), { hc_dim, hc_lr }, trunk_flags);
+ hc_head_up = create_tensor(tn(LLM_TENSOR_HC_HEAD_UP, "weight"), { hc_lr, hc_dim }, trunk_flags);
output = create_tensor(tn(LLM_TENSOR_OUTPUT, "weight"), { n_embd, n_vocab }, TENSOR_NOT_REQUIRED);
if (output == NULL) {
@@ -213,7 +215,7 @@ void llama_model_qwen4exp::load_arch_tensors(llama_model_loader & ml) {
{ hparams.ple_head_dim, ple_rows }, TENSOR_READ_LAZY);
}
- for (int il = 0; il < n_layer; ++il) {
+ auto load_block = [&](int il, int flags) {
auto & layer = layers[il];
const int64_t n_ff_exp = hparams.n_ff_exp() ? hparams.n_ff_exp() : n_ff / n_expert_used;
@@ -228,61 +230,82 @@ void llama_model_qwen4exp::load_arch_tensors(llama_model_loader & ml) {
const int64_t conv_dim = key_dim * 2 + value_dim;
// two HC modules per layer: before the token mixer, before the MoE
- layer.hc_attn_norm = create_tensor(tn(LLM_TENSOR_HC_ATTN_NORM, "weight", il), { n_embd, hc }, TENSOR_ALLOW_RESHAPE);
- layer.hc_attn_down = create_tensor(tn(LLM_TENSOR_HC_ATTN_DOWN, "weight", il), { hc_dim, hc_lr }, 0);
- layer.hc_attn_up = create_tensor(tn(LLM_TENSOR_HC_ATTN_UP, "weight", il), { hc_lr, hc_dim }, 0);
- layer.hc_attn_inject = create_tensor(tn(LLM_TENSOR_HC_ATTN_INJECT, "weight", il), { hc_dim, hc }, 0);
- layer.hc_ffn_norm = create_tensor(tn(LLM_TENSOR_HC_FFN_NORM, "weight", il), { n_embd, hc }, TENSOR_ALLOW_RESHAPE);
- layer.hc_ffn_down = create_tensor(tn(LLM_TENSOR_HC_FFN_DOWN, "weight", il), { hc_dim, hc_lr }, 0);
- layer.hc_ffn_up = create_tensor(tn(LLM_TENSOR_HC_FFN_UP, "weight", il), { hc_lr, hc_dim }, 0);
- layer.hc_ffn_inject = create_tensor(tn(LLM_TENSOR_HC_FFN_INJECT, "weight", il), { hc_dim, hc }, 0);
+ layer.hc_attn_norm = create_tensor(tn(LLM_TENSOR_HC_ATTN_NORM, "weight", il), { n_embd, hc }, flags | TENSOR_ALLOW_RESHAPE);
+ layer.hc_attn_down = create_tensor(tn(LLM_TENSOR_HC_ATTN_DOWN, "weight", il), { hc_dim, hc_lr }, flags);
+ layer.hc_attn_up = create_tensor(tn(LLM_TENSOR_HC_ATTN_UP, "weight", il), { hc_lr, hc_dim }, flags);
+ layer.hc_attn_inject = create_tensor(tn(LLM_TENSOR_HC_ATTN_INJECT, "weight", il), { hc_dim, hc }, flags);
+ layer.hc_ffn_norm = create_tensor(tn(LLM_TENSOR_HC_FFN_NORM, "weight", il), { n_embd, hc }, flags | TENSOR_ALLOW_RESHAPE);
+ layer.hc_ffn_down = create_tensor(tn(LLM_TENSOR_HC_FFN_DOWN, "weight", il), { hc_dim, hc_lr }, flags);
+ layer.hc_ffn_up = create_tensor(tn(LLM_TENSOR_HC_FFN_UP, "weight", il), { hc_lr, hc_dim }, flags);
+ layer.hc_ffn_inject = create_tensor(tn(LLM_TENSOR_HC_FFN_INJECT, "weight", il), { hc_dim, hc }, flags);
if (!hparams.is_recr(il)) {
// full attention: wq holds [q|gate] interleaved per head
- create_tensor_qkv(layer, il, n_embd, n_embd_head_k * n_head * 2, n_embd_k_gqa, n_embd_v_gqa, 0);
- layer.wo = create_tensor(tn(LLM_TENSOR_ATTN_OUT, "weight", il), { n_embd_head_k * n_head, n_embd }, 0);
+ create_tensor_qkv(layer, il, n_embd, n_embd_head_k * n_head * 2, n_embd_k_gqa, n_embd_v_gqa, flags);
+ layer.wo = create_tensor(tn(LLM_TENSOR_ATTN_OUT, "weight", il), { n_embd_head_k * n_head, n_embd }, flags);
- layer.attn_q_norm = create_tensor(tn(LLM_TENSOR_ATTN_Q_NORM, "weight", il), { n_embd_head_k }, 0);
- layer.attn_k_norm = create_tensor(tn(LLM_TENSOR_ATTN_K_NORM, "weight", il), { n_embd_head_k }, 0);
+ layer.attn_q_norm = create_tensor(tn(LLM_TENSOR_ATTN_Q_NORM, "weight", il), { n_embd_head_k }, flags);
+ layer.attn_k_norm = create_tensor(tn(LLM_TENSOR_ATTN_K_NORM, "weight", il), { n_embd_head_k }, flags);
const int64_t idx_dim = hparams.indexer_head_size;
- layer.index_q_proj = create_tensor(tn(LLM_TENSOR_INDEXER_Q_PROJ, "weight", il), { n_embd, hparams.indexer_n_head * idx_dim }, 0);
- layer.index_k_proj = create_tensor(tn(LLM_TENSOR_INDEXER_K_PROJ, "weight", il), { n_embd, idx_dim }, 0);
- layer.index_q_norm = create_tensor(tn(LLM_TENSOR_INDEXER_Q_NORM, "weight", il), { idx_dim }, 0);
- layer.index_k_norm = create_tensor(tn(LLM_TENSOR_INDEXER_K_NORM, "weight", il), { idx_dim }, 0);
+ layer.index_q_proj = create_tensor(tn(LLM_TENSOR_INDEXER_Q_PROJ, "weight", il), { n_embd, hparams.indexer_n_head * idx_dim }, flags);
+ layer.index_k_proj = create_tensor(tn(LLM_TENSOR_INDEXER_K_PROJ, "weight", il), { n_embd, idx_dim }, flags);
+ layer.index_q_norm = create_tensor(tn(LLM_TENSOR_INDEXER_Q_NORM, "weight", il), { idx_dim }, flags);
+ layer.index_k_norm = create_tensor(tn(LLM_TENSOR_INDEXER_K_NORM, "weight", il), { idx_dim }, flags);
} else {
- layer.wqkv = create_tensor(tn(LLM_TENSOR_ATTN_QKV, "weight", il), { n_embd, key_dim * 2 + value_dim }, 0);
- layer.wqkv_gate = create_tensor(tn(LLM_TENSOR_ATTN_GATE, "weight", il), { n_embd, value_dim }, 0);
- layer.ssm_conv1d = create_tensor(tn(LLM_TENSOR_SSM_CONV1D, "weight", il), { hparams.ssm_d_conv, conv_dim }, 0);
- layer.ssm_dt = create_tensor(tn(LLM_TENSOR_SSM_DT, "bias", il), { hparams.ssm_dt_rank }, 0);
- layer.ssm_a = create_tensor(tn(LLM_TENSOR_SSM_A_NOSCAN, il), { hparams.ssm_dt_rank }, 0);
- layer.ssm_beta = create_tensor(tn(LLM_TENSOR_SSM_BETA, "weight", il), { n_embd, n_v_heads }, 0);
- layer.ssm_alpha = create_tensor(tn(LLM_TENSOR_SSM_ALPHA, "weight", il), { n_embd, n_v_heads }, 0);
- layer.ssm_norm = create_tensor(tn(LLM_TENSOR_SSM_NORM, "weight", il), { head_v_dim }, 0);
- layer.ssm_out = create_tensor(tn(LLM_TENSOR_SSM_OUT, "weight", il), { value_dim, n_embd }, 0);
+ layer.wqkv = create_tensor(tn(LLM_TENSOR_ATTN_QKV, "weight", il), { n_embd, key_dim * 2 + value_dim }, flags);
+ layer.wqkv_gate = create_tensor(tn(LLM_TENSOR_ATTN_GATE, "weight", il), { n_embd, value_dim }, flags);
+ layer.ssm_conv1d = create_tensor(tn(LLM_TENSOR_SSM_CONV1D, "weight", il), { hparams.ssm_d_conv, conv_dim }, flags);
+ layer.ssm_dt = create_tensor(tn(LLM_TENSOR_SSM_DT, "bias", il), { hparams.ssm_dt_rank }, flags);
+ layer.ssm_a = create_tensor(tn(LLM_TENSOR_SSM_A_NOSCAN, il), { hparams.ssm_dt_rank }, flags);
+ layer.ssm_beta = create_tensor(tn(LLM_TENSOR_SSM_BETA, "weight", il), { n_embd, n_v_heads }, flags);
+ layer.ssm_alpha = create_tensor(tn(LLM_TENSOR_SSM_ALPHA, "weight", il), { n_embd, n_v_heads }, flags);
+ layer.ssm_norm = create_tensor(tn(LLM_TENSOR_SSM_NORM, "weight", il), { head_v_dim }, flags);
+ layer.ssm_out = create_tensor(tn(LLM_TENSOR_SSM_OUT, "weight", il), { value_dim, n_embd }, flags);
}
if (hparams.is_ple(il)) {
- layer.ple_key = create_tensor(tn(LLM_TENSOR_PLE_KEY, "weight", il), { n_embd, hc_dim }, 0);
- layer.ple_value = create_tensor(tn(LLM_TENSOR_PLE_VALUE, "weight", il), { n_embd, n_embd }, 0);
- layer.ple_norm_key = create_tensor(tn(LLM_TENSOR_PLE_NORM_KEY, "weight", il), { n_embd, hc }, TENSOR_ALLOW_RESHAPE);
- layer.ple_norm_query = create_tensor(tn(LLM_TENSOR_PLE_NORM_QUERY, "weight", il), { n_embd, hc }, TENSOR_ALLOW_RESHAPE);
- layer.ple_norm_conv = create_tensor(tn(LLM_TENSOR_PLE_NORM_CONV, "weight", il), { n_embd, hc }, TENSOR_ALLOW_RESHAPE);
- layer.ple_conv1d = create_tensor(tn(LLM_TENSOR_PLE_CONV1D, "weight", il), { hparams.ple_conv_kernel, hc_dim }, 0);
+ layer.ple_key = create_tensor(tn(LLM_TENSOR_PLE_KEY, "weight", il), { n_embd, hc_dim }, flags);
+ layer.ple_value = create_tensor(tn(LLM_TENSOR_PLE_VALUE, "weight", il), { n_embd, n_embd }, flags);
+ layer.ple_norm_key = create_tensor(tn(LLM_TENSOR_PLE_NORM_KEY, "weight", il), { n_embd, hc }, flags | TENSOR_ALLOW_RESHAPE);
+ layer.ple_norm_query = create_tensor(tn(LLM_TENSOR_PLE_NORM_QUERY, "weight", il), { n_embd, hc }, flags | TENSOR_ALLOW_RESHAPE);
+ layer.ple_norm_conv = create_tensor(tn(LLM_TENSOR_PLE_NORM_CONV, "weight", il), { n_embd, hc }, flags | TENSOR_ALLOW_RESHAPE);
+ layer.ple_conv1d = create_tensor(tn(LLM_TENSOR_PLE_CONV1D, "weight", il), { hparams.ple_conv_kernel, hc_dim }, flags);
}
- layer.ffn_gate_inp = create_tensor(tn(LLM_TENSOR_FFN_GATE_INP, "weight", il), { n_embd, n_expert }, 0);
- layer.ffn_down_exps = create_tensor(tn(LLM_TENSOR_FFN_DOWN_EXPS, "weight", il), { n_ff_exp, n_embd, n_expert }, 0);
- create_tensor_gate_up_exps(layer, il, n_embd, n_ff_exp, n_expert, 0);
+ layer.ffn_gate_inp = create_tensor(tn(LLM_TENSOR_FFN_GATE_INP, "weight", il), { n_embd, n_expert }, flags);
+ layer.ffn_down_exps = create_tensor(tn(LLM_TENSOR_FFN_DOWN_EXPS, "weight", il), { n_ff_exp, n_embd, n_expert }, flags);
+ create_tensor_gate_up_exps(layer, il, n_embd, n_ff_exp, n_expert, flags);
+
+ layer.ffn_gate_inp_shexp = create_tensor(tn(LLM_TENSOR_FFN_GATE_INP_SHEXP, "weight", il), { n_embd }, flags);
+ layer.ffn_gate_shexp = create_tensor(tn(LLM_TENSOR_FFN_GATE_SHEXP, "weight", il), { n_embd, n_ff_shexp }, flags);
+ layer.ffn_up_shexp = create_tensor(tn(LLM_TENSOR_FFN_UP_SHEXP, "weight", il), { n_embd, n_ff_shexp }, flags);
+ layer.ffn_down_shexp = create_tensor(tn(LLM_TENSOR_FFN_DOWN_SHEXP, "weight", il), { n_ff_shexp, n_embd }, flags);
+ };
+
+ for (int il = 0; il < n_layer; ++il) {
+ load_block(il, trunk_flags);
+ }
- layer.ffn_gate_inp_shexp = create_tensor(tn(LLM_TENSOR_FFN_GATE_INP_SHEXP, "weight", il), { n_embd }, 0);
- layer.ffn_gate_shexp = create_tensor(tn(LLM_TENSOR_FFN_GATE_SHEXP, "weight", il), { n_embd, n_ff_shexp }, 0);
- layer.ffn_up_shexp = create_tensor(tn(LLM_TENSOR_FFN_UP_SHEXP, "weight", il), { n_embd, n_ff_shexp }, 0);
- layer.ffn_down_shexp = create_tensor(tn(LLM_TENSOR_FFN_DOWN_SHEXP, "weight", il), { n_ff_shexp, n_embd }, 0);
+ // the MTP block: one full-attention QSA layer fed by [enorm(e) ; hnorm(h)_s] -> eh_proj per hc stream
+ for (int il = n_layer; il < n_layer_all; ++il) {
+ load_block(il, mtp_flags);
+
+ auto & nextn = layers[il].nextn;
+ nextn.eh_proj = create_tensor(tn(LLM_TENSOR_NEXTN_EH_PROJ, "weight", il), { 2*n_embd, n_embd }, mtp_flags);
+ nextn.enorm = create_tensor(tn(LLM_TENSOR_NEXTN_ENORM, "weight", il), { n_embd }, mtp_flags);
+ // RMS per hc stream of the trunk residual, so the gammas load as [n_embd, hc] like the mixer norms
+ nextn.hnorm = create_tensor(tn(LLM_TENSOR_NEXTN_HNORM, "weight", il), { n_embd, hc }, mtp_flags | TENSOR_ALLOW_RESHAPE);
+ nextn.hc_head_norm = create_tensor(tn(LLM_TENSOR_NEXTN_HC_HEAD_NORM, "weight", il), { n_embd, hc }, mtp_flags | TENSOR_ALLOW_RESHAPE);
+ nextn.hc_head_down = create_tensor(tn(LLM_TENSOR_NEXTN_HC_HEAD_DOWN, "weight", il), { hc_dim, hc_lr }, mtp_flags);
+ nextn.hc_head_up = create_tensor(tn(LLM_TENSOR_NEXTN_HC_HEAD_UP, "weight", il), { hc_lr, hc_dim }, mtp_flags);
}
}
std::unique_ptr<llm_graph_context> llama_model_qwen4exp::build_arch_graph(const llm_graph_params & params) const {
+ if (params.gtype == LLM_GRAPH_TYPE_DECODER_MTP) {
+ return std::make_unique<graph_mtp>(*this, params);
+ }
return std::make_unique<graph>(*this, params);
}
@@ -447,7 +470,7 @@ llama_model_qwen4exp::graph::graph(const llama_model & model, const llm_graph_pa
cur = build_layer_attn(inp->get_attn(), mctx_hyb, inp_kpool, cur, inp_pos, sections, il);
}
- if (il == n_layer - 1 && inp_out_ids) {
+ if (il == n_layer - 1 && inp_out_ids && (!cparams.embeddings_nextn || cparams.embeddings_nextn_masked)) {
// everything below is per token, so drop the rows that produce no output
cur = ggml_get_rows(ctx0, cur, inp_out_ids);
inject = ggml_get_rows(ctx0, inject, inp_out_ids);
@@ -475,6 +498,19 @@ llama_model_qwen4exp::graph::graph(const llama_model & model, const llm_graph_pa
cb(res_hc, "l_last", il);
}
+ // the MTP head reads the hc-wide residual, before the final mixer
+ if (cparams.embeddings_nextn) {
+ res->t_h_nextn = ggml_reshape_2d(ctx0, res_hc, n_embd*hc, res_hc->ne[2]);
+ cb(res->t_h_nextn, "h_nextn", -1);
+ ggml_build_forward_expand(gf, res->t_h_nextn);
+ }
+
+ if (cparams.embeddings_nextn && !cparams.embeddings_nextn_masked && inp_out_ids) {
+ res_hc = ggml_reshape_2d(ctx0, res_hc, n_embd*hc, res_hc->ne[2]);
+ res_hc = ggml_get_rows(ctx0, res_hc, inp_out_ids);
+ res_hc = ggml_reshape_3d(ctx0, res_hc, n_embd, hc, res_hc->ne[1]);
+ }
+
// the final mixer is the output norm: there is no separate one
ggml_tensor * cur = build_hc_mix(res_hc,
model.hc_head_norm, model.hc_head_down, model.hc_head_up,
@@ -490,6 +526,97 @@ llama_model_qwen4exp::graph::graph(const llama_model & model, const llm_graph_pa
ggml_build_forward_expand(gf, cur);
}
+llama_model_qwen4exp::graph_mtp::graph_mtp(const llama_model & model, const llm_graph_params & params) :
+ graph(model, params, no_build{}) {
+ GGML_ASSERT(hparams.n_layer_nextn == 1 && "qwen4exp MTP has a single block");
+ GGML_ASSERT(ubatch.token && "qwen4exp MTP requires token input");
+
+ const int64_t hc = hparams.dsv4_hc_mult;
+ GGML_ASSERT(hparams.n_embd_out() == (uint32_t) (n_embd*hc) && "qwen4exp MTP hidden width mismatch");
+
+ const int il = hparams.n_layer();
+ const auto & layer = model.layers[il];
+
+ GGML_ASSERT(layer.nextn.eh_proj && layer.nextn.enorm && layer.nextn.hnorm && layer.nextn.hc_head_norm &&
+ "MTP block missing, load the model with MTP enabled");
+
+ int sections[4];
+ std::copy(std::begin(hparams.rope_sections), std::begin(hparams.rope_sections) + 4, sections);
+
+ auto inp = std::make_unique<llm_graph_input_embd_h>(hparams.n_embd_out());
+
+ inp->tokens = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, n_tokens);
+ ggml_set_input(inp->tokens);
+
+ inp->embd = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, hparams.n_embd_out(), n_tokens);
+ ggml_set_input(inp->embd);
+
+ inp->h = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, hparams.n_embd_out(), n_tokens);
+ ggml_set_input(inp->h);
+ ggml_set_name(inp->h, "mtp_h_input");
+
+ ggml_tensor * tok_embd = ggml_get_rows(ctx0, model.tok_embd, inp->tokens);
+ cb(tok_embd, "mtp_tok_embd", il);
+
+ ggml_tensor * h = inp->h;
+
+ res->add_input(std::move(inp));
+
+ auto * inp_hyb = build_inp_mem_hybrid();
+ const auto * mctx_hyb = static_cast<const llama_memory_hybrid_idx_context *>(inp_hyb->mctx);
+
+ // the draft memory has no recurrent layer, but its input still has to be allocated
+ ggml_build_forward_expand(gf, inp_hyb->get_recr()->s_copy);
+
+ llm_graph_input_kpool * inp_kpool = nullptr;
+ if (mctx_hyb->get_idx() && hparams.indexer_kpool > 0) {
+ GGML_ASSERT(mctx_hyb->get_idx()->get_n_kv() == mctx_hyb->get_attn()->get_n_kv() &&
+ "the indexer cache must track the attention cache cell for cell");
+ inp_kpool = build_inp_kpool(mctx_hyb);
+ }
+
+ ggml_tensor * inp_pos = build_inp_pos();
+ ggml_tensor * inp_out_ids = build_inp_out_ids();
+
+ ggml_tensor * h_norm = build_norm(ggml_reshape_3d(ctx0, h, n_embd, hc, n_tokens), layer.nextn.hnorm, nullptr, LLM_NORM_RMS, il);
+ cb(h_norm, "mtp_hnorm", il);
+
+ ggml_tensor * e_norm = build_norm(tok_embd, layer.nextn.enorm, nullptr, LLM_NORM_RMS, il);
+ e_norm = ggml_repeat_4d(ctx0, ggml_reshape_3d(ctx0, e_norm, n_embd, 1, n_tokens), n_embd, hc, n_tokens, 1);
+ cb(e_norm, "mtp_enorm", il);
+
+ ggml_tensor * res_hc = build_lora_mm(layer.nextn.eh_proj, ggml_concat(ctx0, e_norm, h_norm, 0)); // [n_embd, hc, n_tokens]
+ cb(res_hc, "mtp_eh_proj", il);
+
+ ggml_tensor * inject = nullptr;
+ ggml_tensor * cur = build_hc_mix(res_hc, layer.hc_attn_norm, layer.hc_attn_down, layer.hc_attn_up, layer.hc_attn_inject, &inject, il);
+ cur = build_layer_attn(inp_hyb->get_attn(), mctx_hyb, inp_kpool, cur, inp_pos, sections, il);
+ res_hc = build_hc_combine(res_hc, cur, inject, il);
+
+ cur = build_hc_mix(res_hc, layer.hc_ffn_norm, layer.hc_ffn_down, layer.hc_ffn_up, layer.hc_ffn_inject, &inject, il);
+ cur = build_layer_ffn(cur, il);
+ res_hc = build_hc_combine(res_hc, cur, inject, il);
+
+ // the next draft step reads this residual as its h
+ ggml_tensor * flat = ggml_reshape_2d(ctx0, res_hc, n_embd*hc, n_tokens);
+ ggml_tensor * flat_out = inp_out_ids ? ggml_get_rows(ctx0, flat, inp_out_ids) : flat;
+ res->t_h_nextn = cparams.embeddings_nextn_masked ? flat_out : flat;
+ cb(res->t_h_nextn, "h_nextn", il);
+ ggml_build_forward_expand(gf, res->t_h_nextn);
+
+ cur = build_hc_mix(ggml_reshape_3d(ctx0, flat_out, n_embd, hc, flat_out->ne[1]),
+ layer.nextn.hc_head_norm, layer.nextn.hc_head_down, layer.nextn.hc_head_up,
+ nullptr, nullptr, il);
+ cb(cur, "result_norm", -1);
+ res->t_embd = cur;
+
+ cur = build_lora_mm(model.output, cur, model.output_s);
+ cb(cur, "result_output", -1);
+ res->t_logits = cur;
+
+ ggml_build_forward_expand(gf, cur);
+}
+
std::pair<ggml_tensor *, ggml_tensor *> llama_model_qwen4exp::graph::build_qkvz(
ggml_tensor * input,
int il) {
@@ -713,6 +840,7 @@ ggml_tensor * llama_model_qwen4exp::graph::build_qsa_sel(
const int64_t n_sel = sel_idx->ne[0];
GGML_ASSERT(n_sel == inp_kpool->n_sel);
+ // TODO: figure out to reduce the large copmute buffer that this creates
// scatter zeros for the selected cells into an all -inf row, the extra row n_kv takes the sentinels
// seeding from sel_idx ties the scatter storage lifetime to this layer
const int64_t n_kv = inp_kpool->n_kv;