File size: 2,984 Bytes
f2ee801
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
from __future__ import annotations
import argparse, json, os, sys, time, gc
from pathlib import Path

os.environ.setdefault("ATTN_BACKEND", "sdpa")

import torch

MODEL_ID = "VAST-AI/AniGen"


class SSDecoderExport(torch.nn.Module):
    def __init__(self, model):
        super().__init__()
        self.model = model

    def forward(self, z, z_skl):
        occ, occ_skl = self.model(z, z_skl)
        return occ, occ_skl


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--out", default=os.environ.get("CF_ONNX_OUT", "/tmp/cf-ss-decoder"))
    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"))
    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-decoder"
    out_dir.mkdir(parents=True, exist_ok=True)
    sys.path.insert(0, str(app_root))

    from huggingface_hub import snapshot_download
    snapshot_download(
        MODEL_ID,
        token=os.environ.get("HF_TOKEN"),
        local_dir=model_root,
        allow_patterns=[
            "ckpts/anigen/ss_dae/config.json",
            "ckpts/anigen/ss_dae/ckpts/decoder_final.pt",
        ],
    )
    os.chdir(model_root)

    from anigen.utils.model_utils import load_decoder
    model = load_decoder(str(model_root / "ckpts/anigen/ss_dae"), "final", "cuda").eval()
    wrapper = SSDecoderExport(model).eval()

    z = torch.zeros((1, 8, 16, 16, 16), device="cuda", dtype=torch.float32)
    z_skl = torch.zeros((1, 4, 16, 16, 16), device="cuda", dtype=torch.float32)
    path = out_dir / "model.onnx"

    started = time.time()
    with torch.inference_mode():
        torch.onnx.export(
            wrapper,
            (z, z_skl),
            str(path),
            input_names=["z", "z_skl"],
            output_names=["occupancy", "occupancy_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-decoder",
        "source": MODEL_ID,
        "checkpoint": "ckpts/anigen/ss_dae/ckpts/decoder_final.pt",
        "opset": 23,
        "static_profile": {
            "z": [1, 8, 16, 16, 16],
            "z_skl": [1, 4, 16, 16, 16],
            "occupancy": [1, 1, 64, 64, 64],
            "occupancy_skl": [1, 1, 64, 64, 64],
        },
        "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_DECODER_EXPORTED", json.dumps(meta), flush=True)

    del wrapper, model, z, z_skl
    gc.collect(); torch.cuda.empty_cache()


if __name__ == "__main__":
    main()