Companion-Forge-L4-ONNX / scripts /export_slat_flow_core.py
patdev's picture
Add AniGen SLat transformer core ONNX exporter
226fe2e verified
Raw
History Blame Contribute Delete
7.29 kB
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()