from __future__ import annotations import argparse,json,os,sys from pathlib import Path import torch import torch.nn as nn import torch.nn.functional as F APP_ROOT=Path(os.environ.get('ANIGEN_APP_ROOT','/home/user/app')) if str(APP_ROOT) not in sys.path:sys.path.insert(0,str(APP_ROOT)) DOMAIN='com.companionforge';NS='companionforge' _lib=torch.library.Library("companionforge_heads","DEF") try: _lib.define("sparse_subdivide(Tensor feats, Tensor coords) -> (Tensor, Tensor)") _lib.define("sparse_subm_conv3d(Tensor feats, Tensor coords, Tensor weight, Tensor bias, int out_channels, int kernel) -> (Tensor, Tensor, Tensor)") except Exception:pass def _sub_cuda(feats,coords): off=torch.tensor([[0,x,y,z] for x in (0,1) for y in (0,1) for z in (0,1)],device=coords.device,dtype=coords.dtype);oc=coords.clone();oc[:,1:]*=2;oc=(oc[:,None,:]+off[None]).flatten(0,1);of=feats[:,None,:].expand(feats.shape[0],8,feats.shape[1]).flatten(0,1);return of,oc def _conv_cuda(feats,coords,weight,bias,out_channels:int,kernel:int): return torch.zeros((feats.shape[0],int(out_channels)),device=feats.device,dtype=feats.dtype),coords,torch.tensor(feats.shape[0],device=feats.device,dtype=torch.int32) try: _lib.impl("sparse_subdivide",_sub_cuda,"CUDA");_lib.impl("sparse_subm_conv3d",_conv_cuda,"CUDA") except Exception:pass def _sub_sym(g,feats,coords):return g.op(f'{DOMAIN}::SparseSubdivide',feats,coords,outputs=2,plugin_namespace_s=NS,plugin_version_s='1') def _conv_sym(g,feats,coords,weight,bias,out_channels,kernel): from torch.onnx.symbolic_helper import _get_const oc=int(_get_const(out_channels,'i','out_channels'));k=int(_get_const(kernel,'i','kernel')) return g.op(f'{DOMAIN}::SparseConv3D',feats,coords,weight,bias,out_channels_i=oc,kernel_size_i=k,stride_i=1,dilation_i=1,padding_i=0,subm_i=1,spatial_x_i=256,spatial_y_i=256,spatial_z_i=256,batch_size_i=1,plugin_namespace_s=NS,plugin_version_s='1',outputs=3) torch.onnx.register_custom_op_symbolic('companionforge_heads::sparse_subdivide',_sub_sym,18) torch.onnx.register_custom_op_symbolic('companionforge_heads::sparse_subm_conv3d',_conv_sym,18) def gn(x,norm): # SparseGroupNorm32 for batch=1: [N,C] -> [1,C,N]. y=x.float().transpose(0,1).unsqueeze(0) y=F.group_norm(y,norm.num_groups,norm.weight.float() if norm.weight is not None else None,norm.bias.float() if norm.bias is not None else None,norm.eps) return y.squeeze(0).transpose(0,1).to(x.dtype) def conv(feats,coords,module): c=module.conv; b=c.bias if c.bias is not None else torch.empty(0,device=feats.device,dtype=feats.dtype); k=int(c.kernel_size[0] if isinstance(c.kernel_size,(tuple,list)) else c.kernel_size); return torch.ops.companionforge_heads.sparse_subm_conv3d(feats,coords,c.weight,b,int(c.out_channels),k)[:2] class UpBlock(nn.Module): def __init__(self,b): super().__init__();self.gn0=b.act_layers[0];self.conv1=b.out_layers[0];self.gn1=b.out_layers[1];self.conv2=b.out_layers[3];self.skip=b.skip_connection;self.skip_identity=isinstance(self.skip,nn.Identity) def forward(self,feats,coords): h=F.silu(gn(feats,self.gn0));h,hc=torch.ops.companionforge_heads.sparse_subdivide(h,coords);x,xc=torch.ops.companionforge_heads.sparse_subdivide(feats,coords) h,hc=conv(h,hc,self.conv1);h=F.silu(gn(h,self.gn1));h,hc=conv(h,hc,self.conv2) if self.skip_identity:s=x else:s,_=conv(x,xc,self.skip) return h+s,hc class UpsampleHead(nn.Module): def __init__(self,decoder,kind): super().__init__();self.kind=kind if kind=='geo':self.blocks=decoder.upsample;self.out=decoder.out_layer elif kind=='skin':self.blocks=decoder.upsample_skin_net;self.out=decoder.out_layer_skin else:raise ValueError(kind) self.wrapped=nn.ModuleList([UpBlock(x) for x in self.blocks]) def forward(self,feats,coords): for b in self.wrapped:feats,coords=b(feats,coords) # Output SparseLinear weights remain fp32 upstream; emulate autocast result. y=F.linear(feats.float(),self.out.weight.float(),self.out.bias.float() if self.out.bias is not None else None) return y,coords class SkeletonHead(nn.Module): def __init__(self,decoder):super().__init__();self.out=decoder.out_layer_skl;self.skin=decoder.out_layer_skl_skin def forward(self,feats): a=F.linear(feats.float(),self.out.weight.float(),self.out.bias.float() if self.out.bias is not None else None) b=F.linear(feats.float(),self.skin.weight.float(),self.skin.bias.float() if self.skin.bias is not None else None) return torch.cat([a,b],dim=-1) def export_up(decoder,kind,out): m=UpsampleHead(decoder,kind).cuda().eval();cin=decoder.model_channels if kind=='geo' else decoder.model_channels_skin f=torch.randn(64,cin,device='cuda',dtype=torch.float16);xyz=torch.randint(0,64,(64,3),device='cuda',dtype=torch.int32);c=torch.cat([torch.zeros((64,1),device='cuda',dtype=torch.int32),xyz],1) p=out/f'{kind}-head.onnx';torch.onnx.export(m,(f,c),str(p),input_names=['feats','coords'],output_names=['out_feats','out_coords'],dynamic_axes={'feats':{0:'N'},'coords':{0:'N'},'out_feats':{0:'N4'},'out_coords':{0:'N4'}},opset_version=18,do_constant_folding=True,external_data=True,dynamo=False) fix(p);print('EXPORTED',p,flush=True) def export_skl(decoder,out): m=SkeletonHead(decoder).cuda().eval();f=torch.randn(32,decoder.model_channels_skl,device='cuda',dtype=torch.float16);p=out/'skl-head.onnx';torch.onnx.export(m,(f,),str(p),input_names=['feats'],output_names=['out_feats'],dynamic_axes={'feats':{0:'N'},'out_feats':{0:'N'}},opset_version=18,do_constant_folding=True,external_data=True,dynamo=False);fix(p);print('EXPORTED',p,flush=True) def fix(p): import onnx m=onnx.load(str(p),load_external_data=False) if not any(x.domain==DOMAIN for x in m.opset_import):m.opset_import.append(onnx.helper.make_opsetid(DOMAIN,1)) m.producer_name='Companion-Forge';m.producer_version='6.5-slat-dae-custom';onnx.save_model(m,str(p),save_as_external_data=True,all_tensors_to_one_file=True,location=p.name+'.data',size_threshold=1024) def main(): ap=argparse.ArgumentParser();ap.add_argument('--out',default='/tmp/slat-dae-heads');args=ap.parse_args();out=Path(args.out);out.mkdir(parents=True,exist_ok=True) from huggingface_hub import snapshot_download root=Path('/tmp/anigen-model');snapshot_download('VAST-AI/AniGen',token=os.environ.get('HF_TOKEN'),local_dir=root,allow_patterns=['ckpts/anigen/slat_dae/config.json','ckpts/anigen/slat_dae/ckpts/decoder_final.pt']);os.chdir(root) from anigen.utils.model_utils import load_decoder d=load_decoder('ckpts/anigen/slat_dae','final','cuda');export_up(d,'geo',out);export_up(d,'skin',out);export_skl(d,out) meta={'geo_in':d.model_channels,'skin_in':d.model_channels_skin,'skl_in':d.model_channels_skl,'resolution':d.resolution,'ops':['SparseSubdivide','SparseConv3D']};(out/'meta.json').write_text(json.dumps(meta,indent=2)) if __name__=='__main__':main()