| from __future__ import annotations | |
| import json,torch,gc | |
| from pathlib import Path | |
| from safetensors.torch import load_file | |
| def _strip_core(model): | |
| import torch.nn as nn | |
| model.blocks=nn.ModuleList();model.blocks_vert_skin=nn.ModuleList();model.blocks_skl=nn.ModuleList();model.adapter_geo_to_skin=nn.ModuleList();gc.collect() | |
| return model | |
| def load_slat_shell(runtime_root:Path,device='cuda'): | |
| """Load the 188MB sparse shell without the 2.46GB trained transformer core. | |
| Construct on CPU so ordinary (non-buffer) positional-frequency tensors are | |
| initialized correctly, strip the random core immediately, then load only the | |
| shell weights and move the small shell to CUDA. | |
| """ | |
| cfg=json.loads((runtime_root/'weights/anigen/slat-flow-shell/config.json').read_text());mc=cfg['models']['denoiser'] | |
| from anigen import models | |
| model=getattr(models,mc['name'])(**mc['args']);_strip_core(model) | |
| sd=load_file(str(runtime_root/'weights/anigen/slat-flow-shell/shell.safetensors'),device='cpu') | |
| missing,unexpected=model.load_state_dict(sd,strict=False) | |
| bad=[k for k in missing if not k.startswith(('blocks.','blocks_vert_skin.','blocks_skl.','adapter_geo_to_skin.'))] | |
| if bad or unexpected:raise RuntimeError(f'SLat shell state mismatch: missing={bad[:8]} unexpected={unexpected[:8]}') | |
| model.to(device).eval();model._companion_meta_shell=False | |
| return model,cfg | |