Companion-Forge-L4-ONNX / scripts /export_ss_flow.py
patdev's picture
Add full AniGen SS Flow ONNX exporter
2e605e2 verified
Raw
History Blame Contribute Delete
3.71 kB
from __future__ import annotations
import argparse, json, os, sys, time, gc
from pathlib import Path
# Export with PyTorch SDPA instead of flash-attn custom calls.
os.environ.setdefault("ATTN_BACKEND", "sdpa")
import torch
MODEL_ID = "VAST-AI/AniGen"
OUT_DEFAULT = "/tmp/cf-ss-flow"
MODEL_ROOT_DEFAULT = "/tmp/anigen-model"
APP_ROOT_DEFAULT = "/home/user/app"
def ensure_model(root: Path):
from huggingface_hub import snapshot_download
snapshot_download(
MODEL_ID,
token=os.environ.get("HF_TOKEN"),
local_dir=root,
allow_patterns=[
"ckpts/anigen/ss_flow_solo/config.json",
"ckpts/anigen/ss_flow_solo/ckpts/**",
],
)
class SSFlowExport(torch.nn.Module):
def __init__(self, model):
super().__init__()
self.model = model
def forward(self, x, x_skl, timestep, cond):
y, y_skl = self.model(x, x_skl, timestep, cond)
return y, y_skl
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--out", default=os.environ.get("CF_ONNX_OUT", OUT_DEFAULT))
ap.add_argument("--model-root", default=os.environ.get("ANIGEN_MODEL_ROOT", MODEL_ROOT_DEFAULT))
ap.add_argument("--app-root", default=os.environ.get("ANIGEN_APP_ROOT", APP_ROOT_DEFAULT))
args = ap.parse_args()
out = Path(args.out)
model_root = Path(args.model_root)
app_root = Path(args.app_root)
out_dir = out / "onnx/anigen/ss-flow-solo"
out_dir.mkdir(parents=True, exist_ok=True)
sys.path.insert(0, str(app_root))
ensure_model(model_root)
os.chdir(model_root)
from anigen.utils.model_utils import load_model_from_path
model, cfg = load_model_from_path(
str(model_root / "ckpts/anigen/ss_flow_solo"),
model_name_in_config="denoiser",
device="cuda",
use_ema=False,
)
model.eval()
wrapper = SSFlowExport(model).eval()
# Production profile is batch=1. SS latent is fixed 16^3; DINO ViT-L/14-reg
# on 518x518 yields 1369 patch + cls + 4 register = 1374 conditioning tokens.
x = torch.zeros((1, 8, 16, 16, 16), device="cuda", dtype=torch.float32)
x_skl = torch.zeros((1, 4, 16, 16, 16), device="cuda", dtype=torch.float32)
timestep = torch.tensor([500.0], device="cuda", dtype=torch.float32)
cond = torch.zeros((1, 1374, 1024), device="cuda", dtype=torch.float32)
path = out_dir / "model.onnx"
started = time.time()
with torch.inference_mode():
torch.onnx.export(
wrapper,
(x, x_skl, timestep, cond),
str(path),
input_names=["x", "x_skl", "timestep", "cond"],
output_names=["velocity", "velocity_skl"],
opset_version=23,
dynamo=True,
external_data=True,
)
export_s = time.time() - started
import onnx
onnx.checker.check_model(str(path))
meta = {
"component": "anigen-ss-flow-solo",
"source": MODEL_ID,
"checkpoint": "ckpts/anigen/ss_flow_solo",
"opset": 23,
"precision_source": "fp16 torso / fp32 io",
"static_profile": {
"x": [1, 8, 16, 16, 16],
"x_skl": [1, 4, 16, 16, 16],
"timestep": [1],
"cond": [1, 1374, 1024],
},
"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("SS_FLOW_EXPORTED", json.dumps(meta), flush=True)
del wrapper, model, x, x_skl, timestep, cond
gc.collect(); torch.cuda.empty_cache()
if __name__ == "__main__":
main()