File size: 8,185 Bytes
27e9fbe 2f38eb6 27e9fbe 2f38eb6 27e9fbe 50d8cca 27e9fbe 2f38eb6 27e9fbe 043065a 27e9fbe f6f2c6a 27e9fbe 4edaf55 27e9fbe 4edaf55 27e9fbe | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 | from __future__ import annotations
import argparse
import json
import os
import 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'
PLUGIN_NAMESPACE='companionforge'
# Stable legacy-ONNX custom op path for PyTorch 2.8: torch.library + registered symbolic.
_lib = torch.library.Library("companionforge", "DEF")
try:
_lib.define("sparse_window_attention(Tensor qkv, Tensor coords, int window_size, int shift_x, int shift_y, int shift_z) -> Tensor")
except Exception:
pass
def _window_cuda(qkv, coords, window_size:int, shift_x:int, shift_y:int, shift_z:int):
from anigen.modules.sparse import SparseTensor
from anigen.modules.sparse.attention.windowed_attn import sparse_windowed_scaled_dot_product_self_attention
st=SparseTensor(qkv,coords)
return sparse_windowed_scaled_dot_product_self_attention(st,int(window_size),(int(shift_x),int(shift_y),int(shift_z))).feats
try:
_lib.impl("sparse_window_attention", _window_cuda, "CUDA")
except Exception:
pass
def _window_symbolic(g,qkv,coords,window_size,shift_x,shift_y,shift_z):
from torch.onnx.symbolic_helper import _get_const
ws=int(_get_const(window_size,'i','window_size'));sx=int(_get_const(shift_x,'i','shift_x'));sy=int(_get_const(shift_y,'i','shift_y'));sz=int(_get_const(shift_z,'i','shift_z'))
return g.op(f'{DOMAIN}::SparseWindowAttention',qkv,coords,window_size_i=ws,shift_x_i=sx,shift_y_i=sy,shift_z_i=sz,plugin_namespace_s=PLUGIN_NAMESPACE,plugin_version_s='1')
torch.onnx.register_custom_op_symbolic('companionforge::sparse_window_attention',_window_symbolic,18)
class Branch(nn.Module):
def __init__(self, decoder, branch:str):
super().__init__(); self.branch=branch
if branch=='geo':
self.input_layer=decoder.input_layer; self.pos=decoder.pos_embedder; self.blocks=decoder.blocks
self.channels=decoder.model_channels; self.heads=decoder.num_heads
self.kind='plain'
elif branch=='skin':
self.input_layer=decoder.input_layer_skin; self.pos=decoder.pos_embedder_skin; self.blocks=decoder.blocks_skin
self.channels=decoder.model_channels_skin; self.heads=self.blocks[0].self_attn.num_heads
self.kind='multi'
elif branch=='skl':
self.input_layer=decoder.input_layer_skl; self.pos=decoder.pos_embedder_skl; self.blocks=decoder.blocks_skl
self.channels=decoder.model_channels_skl; self.heads=self.blocks[0].self_attn.num_heads
self.kind='multi'
else: raise ValueError(branch)
self.register_buffer('freqs',self.pos.freqs.detach().clone(),persistent=True)
self.eps=1e-6; self.window_size=int(decoder.window_size)
def position(self,coords):
xyz=coords[:,1:].float(); flat=xyz.reshape(-1)
emb=torch.outer(flat,self.freqs)
emb=torch.cat([torch.sin(emb),torch.cos(emb)],dim=-1)
emb=emb.reshape(xyz.shape[0],-1)
if emb.shape[-1] < self.channels:
emb=torch.cat([emb,torch.zeros((emb.shape[0],self.channels-emb.shape[-1]),device=emb.device,dtype=emb.dtype)],dim=-1)
return emb
def _ln(self,x,norm):
w=norm.weight if getattr(norm,'elementwise_affine',False) else None
b=norm.bias if getattr(norm,'elementwise_affine',False) else None
return F.layer_norm(x.float(),tuple(norm.normalized_shape),w,b,self.eps).to(x.dtype)
def _attn(self,attn,x,coords,shift):
qkv=F.linear(x,attn.to_qkv.weight,attn.to_qkv.bias)
n=qkv.shape[0]; d=self.channels//attn.num_heads
qkv=qkv.reshape(n,3,attn.num_heads,d)
if getattr(attn,'qk_rms_norm',False):
q,k,v=qkv.unbind(1)
# Match SparseMultiHeadRMSNorm: F.normalize(float) * gamma * sqrt(dim), cast back.
q=(F.normalize(q.float(),dim=-1)*attn.q_rms_norm.gamma.float()*float(d**0.5)).to(q.dtype)
k=(F.normalize(k.float(),dim=-1)*attn.k_rms_norm.gamma.float()*float(d**0.5)).to(k.dtype)
qkv=torch.stack([q,k,v],dim=1)
y=torch.ops.companionforge.sparse_window_attention(qkv,coords,self.window_size,shift,shift,shift)
y=y.reshape(n,self.channels)
return F.linear(y,attn.to_out.weight,attn.to_out.bias)
def _mlp(self,mlp,x):
# SparseFeedForwardNet.mlp = SparseLinear, SparseGELU(tanh), SparseLinear.
a=mlp.mlp[0]; b=mlp.mlp[2]
x=F.linear(x,a.weight,a.bias)
x=F.gelu(x,approximate='tanh')
return F.linear(x,b.weight,b.bias)
def forward(self,feats,coords):
x=F.linear(feats.float(),self.input_layer.weight.float(),self.input_layer.bias.float() if self.input_layer.bias is not None else None)
x=(x+self.position(coords).to(x.dtype)).to(torch.float16)
for i,blk in enumerate(self.blocks):
shift=(self.window_size//2)*(i%2)
if self.kind=='plain':
h=self._ln(x,blk.norm1); h=self._attn(blk.attn,h,coords,shift); x=x+h
h=self._ln(x,blk.norm2); h=self._mlp(blk.mlp,h); x=x+h
else:
# Current production decoder has context_num=0: self-attention + FFN only.
h=self._ln(x,blk.norm1); h=self._attn(blk.self_attn,h,coords,shift); x=x+h
h=self._ln(x,blk.norm3); h=self._mlp(blk.mlp,h); x=x+h
return x
def export_one(decoder,branch,out:Path,sample_n:int=128):
m=Branch(decoder,branch).cuda().eval()
cin={'geo':decoder.latent_channels,'skin':decoder.latent_channels_vertskin,'skl':decoder.latent_channels_skl}[branch]
feats=torch.randn(sample_n,cin,device='cuda',dtype=torch.float16)
xyz=torch.randint(0,64,(sample_n,3),device='cuda',dtype=torch.int32)
coords=torch.cat([torch.zeros((sample_n,1),device='cuda',dtype=torch.int32),xyz],dim=1)
path=out/f'{branch}.onnx'; path.parent.mkdir(parents=True,exist_ok=True)
try:
torch.onnx.export(
m,(feats,coords),str(path),
input_names=['feats','coords'],output_names=['out_feats'],
dynamic_axes={'feats':{0:'N'},'coords':{0:'N'},'out_feats':{0:'N'}},
opset_version=18,do_constant_folding=True,external_data=True,dynamo=False,
)
except Exception as exc:
print('EXPORT_ERROR',branch,type(exc).__name__,str(exc)[:3000].replace('\n',' | '),flush=True)
raise SystemExit(31)
import onnx
model=onnx.load(str(path),load_external_data=True)
# Legacy exporter may omit the custom domain opset declaration.
if not any(x.domain==DOMAIN for x in model.opset_import): model.opset_import.append(onnx.helper.make_opsetid(DOMAIN,1))
model.producer_name='Companion-Forge';model.producer_version='6.5-slat-dae-custom'
data_name=f'{branch}.onnx.data'
data_path=out/data_name
if data_path.exists(): data_path.unlink()
onnx.save_model(model,str(path),save_as_external_data=True,all_tensors_to_one_file=True,location=data_name,size_threshold=1024)
meta={'branch':branch,'input_channels':cin,'model_channels':m.channels,'heads':m.heads,'blocks':len(m.blocks),'window_size':m.window_size,'custom_op':f'{DOMAIN}::SparseWindowAttention'}
(out/f'{branch}.json').write_text(json.dumps(meta,indent=2))
print('EXPORTED',branch,path,path.stat().st_size,flush=True)
del m; torch.cuda.empty_cache()
def main():
ap=argparse.ArgumentParser();ap.add_argument('--model-root',default='/tmp/anigen-model');ap.add_argument('--out',default='/tmp/slat-dae-branches');args=ap.parse_args()
from huggingface_hub import snapshot_download
root=Path(args.model_root)
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
decoder=load_decoder('ckpts/anigen/slat_dae','final','cuda')
out=Path(args.out)
for b in ('geo','skin','skl'): export_one(decoder,b,out)
if __name__=='__main__': main()
|