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