Download assemble_head.py from Accio-Lab/occamy-1.0-MTP: direct link, hf CLI and curl.
- Browser
- Download file 3.41 kB
-
https://huggingface.co/Accio-Lab/occamy-1.0-MTP/resolve/main/assemble_head.py
- Command line
-
hf download hf://Accio-Lab/occamy-1.0-MTP/assemble_head.py
-
curl -L -o assemble_head.py https://huggingface.co/Accio-Lab/occamy-1.0-MTP/resolve/main/assemble_head.py
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() | |