Companion-Forge-L4-ONNX / scripts /extract_slat_shell.py
patdev's picture
Add AniGen SLat sparse-shell checkpoint extractor
5fb8c5d verified
Raw
History Blame Contribute Delete
1.47 kB
from __future__ import annotations
import json,os
from pathlib import Path
import torch
from safetensors.torch import save_file
from huggingface_hub import snapshot_download
MODEL='VAST-AI/AniGen'
root=Path('/tmp/anigen-model');out=Path('/tmp/slat-shell');out.mkdir(parents=True,exist_ok=True)
snapshot_download(MODEL,token=os.environ.get('HF_TOKEN'),local_dir=root,allow_patterns=['ckpts/anigen/slat_flow_auto/config.json','ckpts/anigen/slat_flow_auto/ckpts/**'])
config=json.loads((root/'ckpts/anigen/slat_flow_auto/config.json').read_text())
ckpts=sorted((root/'ckpts/anigen/slat_flow_auto/ckpts').glob('*.pt'))
ckpt=[p for p in ckpts if 'ema' not in p.name and 'misc' not in p.name][-1]
sd=torch.load(ckpt,map_location='cpu')
if next(iter(sd)).startswith('module.'): sd={k[7:]:v for k,v in sd.items()}
CORE_PREFIXES=('blocks.','blocks_vert_skin.','blocks_skl.','adapter_geo_to_skin.')
shell={k:v.contiguous() for k,v in sd.items() if not k.startswith(CORE_PREFIXES)}
core={k:v.contiguous() for k,v in sd.items() if k.startswith(CORE_PREFIXES)}
save_file(shell,str(out/'shell.safetensors'))
(out/'config.json').write_text(json.dumps(config,indent=2))
meta={'source_checkpoint':ckpt.name,'shell_tensors':len(shell),'core_tensors_removed':len(core),'shell_bytes':(out/'shell.safetensors').stat().st_size,'original_bytes':ckpt.stat().st_size,'removed_prefixes':list(CORE_PREFIXES)}
(out/'meta.json').write_text(json.dumps(meta,indent=2));print(json.dumps(meta,indent=2))