sglang
mtp
speculative-decoding
draft-head
qwen3_5_moe
occamy-1.0-MTP / assemble_head.py
Eang's picture
Preserve BF16 MTP head in NVFP4 assembly and document community vLLM setup
2fd68d3 verified
Raw History Blame Contribute Delete
3.41 kB
"""Assemble a separate Qwen3.5-MoE-family checkpoint with an independent MTP head.
Does not modify base files. Requires safetensors. Runtime support must be tested.
"""
import argparse, hashlib, json, math
from pathlib import Path
from safetensors import safe_open
def preserve_bf16_head(config, quant_config):
"""Exclude the grafted BF16 head from a ModelOpt NVFP4 recipe."""
quant = (quant_config or {}).get('quantization', {})
embedded = config.get('quantization_config', {})
if quant.get('quant_algo') != 'NVFP4' and embedded.get('quant_algo') != 'NVFP4':
return
for section, key in ((quant, 'exclude_modules'), (embedded, 'ignore')):
values = section.setdefault(key, [])
if not isinstance(values, list):
raise ValueError(f'{key} must be a list')
for pattern in ('mtp.layers.0*', 'mtp*'):
if pattern not in values:
values.append(pattern)
config['quantization_config'] = embedded
def main():
p=argparse.ArgumentParser()
p.add_argument('--base',required=True,type=Path)
p.add_argument('--head',required=True,type=Path)
p.add_argument('--out',required=True,type=Path)
a=p.parse_args();base=a.base.resolve();head=a.head.resolve()
config=json.loads((base/'config.json').read_text())
if config.get('text_config',{}).get('model_type')!='qwen3_5_moe_text':
raise ValueError('Only the tested Qwen3.5 MoE model family is supported')
index=json.loads((base/'model.safetensors.index.json').read_text())
if any(k.startswith('mtp.') for k in index['weight_map']):
raise ValueError('Base already has MTP weights; refusing to replace them')
with safe_open(head,framework='pt') as f:
keys=list(f.keys())
if not keys or not all(k.startswith('mtp.') for k in keys):raise ValueError('Not an MTP-only head')
if not all(f.get_slice(k).get_dtype()=='BF16' for k in keys):raise ValueError('Expected BF16 MTP head')
size=sum(math.prod(f.get_slice(k).get_shape())*2 for k in keys)
quant_path = base/'hf_quant_config.json'
quant_config = json.loads(quant_path.read_text()) if quant_path.exists() else None
preserve_bf16_head(config, quant_config)
a.out.mkdir(parents=True,exist_ok=False)
skip={'config.json','hf_quant_config.json','model.safetensors.index.json','README.md','CANDIDATE_STATUS.json','VALIDATION.json','SHA256SUMS'}
for f in base.iterdir():
if f.is_file() and f.name not in skip:(a.out/f.name).symlink_to(f)
name='mtp-trained.safetensors';(a.out/name).symlink_to(head)
config['text_config']['mtp_num_hidden_layers']=1
for k in keys:index['weight_map'][k]=name
index.setdefault('metadata',{})['total_size']=index.get('metadata',{}).get('total_size',0)+size
(a.out/'config.json').write_text(json.dumps(config,indent=2))
if quant_config is not None:
(a.out/'hf_quant_config.json').write_text(json.dumps(quant_config,indent=2))
(a.out/'model.safetensors.index.json').write_text(json.dumps(index,indent=2))
digest=hashlib.sha256()
with head.open('rb') as f:
for block in iter(lambda:f.read(8*1024*1024),b''):digest.update(block)
(a.out/'MTP_ASSEMBLY.json').write_text(json.dumps({'head_sha256':digest.hexdigest(),'head_tensors':len(keys),'base_files_modified':False,'assembly_only_not_a_validation_claim':True},indent=2))
if __name__=='__main__':main()