from __future__ import annotations import argparse, json, os, sys, time, math, gc from pathlib import Path # The wrapper below uses standard torch SDPA and dense tensors; sparse IO shells # stay native (spconv) around this core. os.environ.setdefault("ATTN_BACKEND", "sdpa") os.environ.setdefault("SPARSE_ATTN_BACKEND", "flash_attn") os.environ.setdefault("SPCONV_ALGO", "native") import torch import torch.nn as nn import torch.nn.functional as F MODEL_ID = "VAST-AI/AniGen" def ln32(mod: nn.LayerNorm, x: torch.Tensor) -> torch.Tensor: w = mod.weight.float() if mod.weight is not None else None b = mod.bias.float() if mod.bias is not None else None return F.layer_norm(x.float(), mod.normalized_shape, w, b, mod.eps).to(x.dtype) def rms(mod, x: torch.Tensor) -> torch.Tensor: # x [..., H, D] dtype = x.dtype y = F.normalize(x.float(), dim=-1) gamma = mod.gamma.float() return (y * gamma * mod.scale).to(dtype) def attention(mod, x: torch.Tensor, context: torch.Tensor | None = None) -> torch.Tensor: # Batch=1 production path. Standard SDPA exports to ONNX and TRT can fuse it. if mod._type == "self": qkv = F.linear(x, mod.to_qkv.weight, mod.to_qkv.bias) B, N, _ = qkv.shape qkv = qkv.reshape(B, N, 3, mod.num_heads, -1) q, k, v = qkv.unbind(dim=2) else: q = F.linear(x, mod.to_q.weight, mod.to_q.bias) kv = F.linear(context, mod.to_kv.weight, mod.to_kv.bias) B, N, _ = q.shape q = q.reshape(B, N, mod.num_heads, -1) kv = kv.reshape(B, kv.shape[1], 2, mod.num_heads, -1) k, v = kv.unbind(dim=2) if mod.qk_rms_norm: q = rms(mod.q_rms_norm, q) k = rms(mod.k_rms_norm, k) q = q.permute(0, 2, 1, 3) k = k.permute(0, 2, 1, 3) v = v.permute(0, 2, 1, 3) y = F.scaled_dot_product_attention(q, k, v) y = y.permute(0, 2, 1, 3).reshape(x.shape[0], x.shape[1], mod.channels) return F.linear(y, mod.to_out.weight, mod.to_out.bias) def mlp(sparse_ffn, x: torch.Tensor) -> torch.Tensor: l1 = sparse_ffn.mlp[0] act = sparse_ffn.mlp[1] l2 = sparse_ffn.mlp[2] y = F.linear(x, l1.weight, l1.bias) y = F.gelu(y, approximate=getattr(act, "approximate", "none")) return F.linear(y, l2.weight, l2.bias) def mod_cross(block, x: torch.Tensor, modvec: torch.Tensor, context: torch.Tensor) -> torch.Tensor: if block.share_mod: parts = modvec.chunk(6, dim=1) else: parts = block.adaLN_modulation(modvec).chunk(6, dim=1) shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = parts h = ln32(block.norm1, x) h = h * (1 + scale_msa[:, None, :]) + shift_msa[:, None, :] h = attention(block.self_attn, h) x = x + h * gate_msa[:, None, :] h = ln32(block.norm2, x) if block.norm_for_context: context = ln32(block.context_norm, context) h = attention(block.cross_attn, h, context) x = x + h h = ln32(block.norm3, x) h = h * (1 + scale_mlp[:, None, :]) + shift_mlp[:, None, :] h = mlp(block.mlp, h) return x + h * gate_mlp[:, None, :] class SLatFlowCore(nn.Module): def __init__(self, model): super().__init__() # Register the original blocks so ONNX sees all weights as initializers. self.blocks = model.blocks self.blocks_skin = model.blocks_vert_skin self.blocks_skl = model.blocks_skl self.adapters = model.adapter_geo_to_skin def forward(self, geo, skin, skl, mod_geo, mod_skin, mod_skl, cond): for b_geo, b_skin, b_skl, adapter in zip(self.blocks, self.blocks_skin, self.blocks_skl, self.adapters): f_geo, f_skin, f_skl = geo, skin, skl geo = mod_cross(b_geo, f_geo, mod_geo, cond) skin = mod_cross(b_skin, f_skin, mod_skin, f_skl) + F.linear(f_geo, adapter.weight, adapter.bias) skl = mod_cross(b_skl, f_skl, mod_skl, f_skin) return geo, skin, skl def main(): ap = argparse.ArgumentParser() ap.add_argument("--out", default=os.environ.get("CF_ONNX_OUT", "/tmp/cf-slat-core")) ap.add_argument("--model-root", default=os.environ.get("ANIGEN_MODEL_ROOT", "/tmp/anigen-model")) ap.add_argument("--app-root", default=os.environ.get("ANIGEN_APP_ROOT", "/home/user/app")) a = ap.parse_args() out = Path(a.out); root = Path(a.model_root); app = Path(a.app_root) out_dir = out / "onnx/anigen/slat-flow-core"; out_dir.mkdir(parents=True, exist_ok=True) sys.path.insert(0, str(app)) from huggingface_hub import snapshot_download snapshot_download(MODEL_ID, token=os.environ.get("HF_TOKEN"), local_dir=root, allow_patterns=["ckpts/anigen/slat_flow_auto/config.json", "ckpts/anigen/slat_flow_auto/ckpts/**"]) os.chdir(root) from anigen.utils.model_utils import load_model_from_path model, cfg = load_model_from_path(str(root / "ckpts/anigen/slat_flow_auto"), model_name_in_config="denoiser", device="cuda") model.eval() core = SLatFlowCore(model).cuda().eval() # Small sample lengths for export; sequence axes are dynamic in the ONNX graph. geo = torch.zeros((1, 256, 1024), device="cuda", dtype=torch.float16) skin = torch.zeros((1, 256, 512), device="cuda", dtype=torch.float16) skl = torch.zeros((1, 128, 512), device="cuda", dtype=torch.float16) mod_geo = torch.zeros((1, 1024), device="cuda", dtype=torch.float16) mod_skin = torch.zeros((1, 512), device="cuda", dtype=torch.float16) mod_skl = torch.zeros((1, 512), device="cuda", dtype=torch.float16) cond = torch.zeros((1, 1374, 1024), device="cuda", dtype=torch.float16) path = out_dir / "model.onnx" started = time.time() # Legacy exporter is used here because dynamic_axes is mature for variable token counts. with torch.inference_mode(): torch.onnx.export( core, (geo, skin, skl, mod_geo, mod_skin, mod_skl, cond), str(path), input_names=["geo", "skin", "skl", "mod_geo", "mod_skin", "mod_skl", "cond"], output_names=["geo_out", "skin_out", "skl_out"], dynamic_axes={ "geo": {1: "n_geo"}, "skin": {1: "n_geo"}, "skl": {1: "n_skl"}, "geo_out": {1: "n_geo"}, "skin_out": {1: "n_geo"}, "skl_out": {1: "n_skl"}, }, opset_version=18, do_constant_folding=True, external_data=True, dynamo=False, ) export_s = time.time() - started import onnx onnx.checker.check_model(str(path)) meta = { "component": "anigen-slat-flow-transformer-core", "source": MODEL_ID, "checkpoint": "ckpts/anigen/slat_flow_auto", "opset": 18, "precision": "fp16", "dynamic": {"n_geo": [128, 4096, 16384], "n_skl": [16, 1024, 8192]}, "cond": [1, 1374, 1024], "native_shell": ["SparseConv3d", "SparseDownsample", "SparseUpsample"], "export_seconds": round(export_s, 3), "torch": torch.__version__, "cuda": torch.version.cuda, "gpu": torch.cuda.get_device_name(0), } (out_dir / "export_meta.json").write_text(json.dumps(meta, indent=2)) print("SLAT_FLOW_CORE_EXPORTED", json.dumps(meta), flush=True) del core, model; gc.collect(); torch.cuda.empty_cache() if __name__ == "__main__": main()