@@ -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 (
0 commit comments