"""Export the trained EMA U-Net to ONNX for browser inference. Run: python diffusion/export_onnx.py Writes diffusion/web_static/unet.onnx and verifies it against PyTorch. """ import os import numpy as np import onnxruntime as ort import torch from train_diffusion import UNet, Diffusion, T HERE = os.path.dirname(os.path.abspath(__file__)) CKPT = os.path.join(HERE, "out_big", "checkpoint.pt") OUT = os.path.join(HERE, "web_static", "unet.onnx") os.makedirs(os.path.dirname(OUT), exist_ok=True) model = UNet(base=128).eval() model.load_state_dict(torch.load(CKPT, map_location="cpu", weights_only=True)["ema"]) x = torch.randn(1, 3, 32, 32) t = torch.tensor([500], dtype=torch.int64) torch.onnx.export( model, (x, t), OUT, input_names=["x", "t"], output_names=["eps"], opset_version=17, dynamo=False, ) print(f"exported: {os.path.getsize(OUT)/1e6:.1f} MB") sess = ort.InferenceSession(OUT, providers=["CPUExecutionProvider"]) with torch.no_grad(): ref = model(x, t).numpy() out = sess.run(None, {"x": x.numpy(), "t": t.numpy()})[0] err = np.abs(ref - out).max() print(f"max abs diff vs pytorch: {err:.2e}") assert err < 1e-3, "ONNX output mismatch" print("verified OK")