Companion-Forge-L4-ONNX / scripts /export_anigen_dense.py
patdev's picture
Fix DSINE ONNX export on CUDA
0e0c466 verified
Raw
History Blame Contribute Delete
2.88 kB
from __future__ import annotations
import json, os, sys, time
from pathlib import Path
import torch
OUT = Path(os.environ.get("CF_ONNX_OUT", "/tmp/cf-onnx-out"))
MODEL_ROOT = Path(os.environ.get("ANIGEN_MODEL_ROOT", "/tmp/anigen-model"))
APP_ROOT = Path(os.environ.get("ANIGEN_APP_ROOT", "/home/user/app"))
OUT.mkdir(parents=True, exist_ok=True)
sys.path.insert(0, str(APP_ROOT))
sys.path.insert(0, str(APP_ROOT / "third_parties" / "dsine"))
def ensure_weights():
from huggingface_hub import snapshot_download
snapshot_download(
repo_id="VAST-AI/AniGen", token=os.environ.get("HF_TOKEN"), local_dir=MODEL_ROOT,
allow_patterns=["ckpts/dinov2/**", "ckpts/dsine/**"],
)
os.chdir(MODEL_ROOT)
class DinoExport(torch.nn.Module):
def __init__(self, model):
super().__init__(); self.model=model
def forward(self, pixel_values):
return self.model(pixel_values, is_training=True)["x_prenorm"]
class DsineExport(torch.nn.Module):
def __init__(self, model):
super().__init__(); self.model=model
def forward(self, image, intrins):
return self.model(image, intrins=intrins)[-1]
def export_one(module, args, path: Path, input_names, output_names):
path.parent.mkdir(parents=True, exist_ok=True)
module.eval()
t=time.time()
with torch.inference_mode():
torch.onnx.export(
module, args, str(path), input_names=input_names, output_names=output_names,
opset_version=18, do_constant_folding=True, dynamo=False,
dynamic_axes=None, external_data=True,
)
import onnx
model=onnx.load(str(path), load_external_data=False)
onnx.checker.check_model(model)
print(f"EXPORTED {path} {time.time()-t:.2f}s", flush=True)
def main():
ensure_weights()
print("loading dinov2", flush=True)
dino = torch.hub.load('./ckpts/dinov2', 'dinov2_vitl14_reg', pretrained=True, source='local').eval().cpu()
x=torch.zeros((1,3,518,518), dtype=torch.float32)
export_one(DinoExport(dino), (x,), OUT/'onnx/dinov2/model.onnx', ['pixel_values'], ['x_prenorm'])
del dino; import gc; gc.collect()
print("loading dsine", flush=True)
from anigen.utils.image_utils import load_dsine
dsine=load_dsine('cuda').eval()
image=torch.zeros((1,3,544,544),dtype=torch.float32,device='cuda')
intrins=torch.tensor([[[471.117,0,259.0],[0,471.117,259.0],[0,0,1.0]]],dtype=torch.float32,device='cuda')
export_one(DsineExport(dsine),(image,intrins),OUT/'onnx/dsine/model.onnx',['image','intrins'],['normal'])
meta={
'torch': torch.__version__, 'opset':18,
'dinov2': {'input':[1,3,518,518], 'output':'x_prenorm'},
'dsine': {'image_input':[1,3,544,544], 'intrinsics_input':[1,3,3], 'output':'normal'},
}
(OUT/'export_meta.json').write_text(json.dumps(meta,indent=2))
if __name__=='__main__': main()