"""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()