| from __future__ import annotations |
| import argparse, json, os, sys, time, math, gc |
| from pathlib import Path |
|
|
| |
| |
| 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: |
| |
| 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: |
| |
| 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__() |
| |
| 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() |
|
|
| |
| 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() |
| |
| 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() |
|
|