Ass3 / export_to_onnx.py
launch-calcium's picture
Upload 6 files
3c08420 verified
Raw
History Blame Contribute Delete
3.61 kB
# export_to_onnx.py
import os
import torch
from diffusers import StableDiffusionPipeline
import argparse
parser = argparse.ArgumentParser()
parser.add_argument("--model_id", default="runwayml/stable-diffusion-v1-5")
parser.add_argument("--out_dir", default="onnx_models")
parser.add_argument("--device", default="cpu")
parser.add_argument("--opset", type=int, default=17)
args = parser.parse_args()
os.makedirs(args.out_dir, exist_ok=True)
print("Loading pipeline:", args.model_id)
pipe = StableDiffusionPipeline.from_pretrained(args.model_id, torch_dtype=torch.float32)
pipe.to(args.device)
tokenizer = pipe.tokenizer
# 1) Export text encoder
print("Exporting text encoder...")
sample_text = "a photo of a cat"
inputs = tokenizer(sample_text, return_tensors="pt", padding="max_length", max_length=77, truncation=True)
input_ids = inputs["input_ids"]
attention_mask = inputs["attention_mask"]
text_encoder = pipe.text_encoder.eval()
text_encoder_input_names = ["input_ids", "attention_mask"]
text_encoder_output_names = ["last_hidden_state"]
dynamic_axes = {
"input_ids": {0: "batch", 1: "seq"},
"attention_mask": {0: "batch", 1: "seq"},
"last_hidden_state": {0: "batch", 1: "seq"}
}
torch.onnx.export(
text_encoder,
(input_ids, attention_mask),
os.path.join(args.out_dir, "text_encoder.onnx"),
input_names=text_encoder_input_names,
output_names=text_encoder_output_names,
dynamic_axes=dynamic_axes,
opset_version=args.opset,
do_constant_folding=True,
)
# 2) Export UNet
print("Exporting UNet...")
unet = pipe.unet.eval()
# For 512x512 generation latents spatial size is 64x64 (latent downscale factor)
dummy_latents = torch.randn(1, unet.in_channels, 64, 64, dtype=torch.float32)
dummy_timestep = torch.tensor([1], dtype=torch.int64)
dummy_encoder_hidden_states = torch.randn(1, 77, pipe.text_encoder.config.hidden_size, dtype=torch.float32)
unet_input_names = ["latent", "timestep", "encoder_hidden_states"]
unet_output_names = ["sample"]
unet_dynamic_axes = {
"latent": {0: "batch", 2: "h", 3: "w"},
"timestep": {0: "batch"},
"encoder_hidden_states": {0: "batch", 1: "seq"},
"sample": {0: "batch", 2: "h", 3: "w"}
}
torch.onnx.export(
unet,
(dummy_latents, dummy_timestep, dummy_encoder_hidden_states),
os.path.join(args.out_dir, "unet.onnx"),
input_names=unet_input_names,
output_names=unet_output_names,
dynamic_axes=unet_dynamic_axes,
opset_version=args.opset,
do_constant_folding=True,
verbose=False,
)
# 3) Export VAE decoder
print("Exporting VAE decoder...")
vae = pipe.vae
vae_decoder = vae.decode.eval()
dummy_vae_latents = torch.randn(1, vae.latent_channels, 64, 64, dtype=torch.float32)
vae_input_names = ["latent"]
vae_output_names = ["sample"]
vae_dynamic_axes = {"latent": {0: "batch", 2: "h", 3: "w"}, "sample": {0: "batch", 2: "h", 3: "w"}}
import torch.nn as nn
class VaeDecodeWrapper(nn.Module):
def __init__(self, vae):
super().__init__()
self.vae = vae
def forward(self, latents):
out = self.vae.decode(latents)
if isinstance(out, dict):
# Some versions return dict-like; adapt if needed
return out.get("sample", out)
return out
vae_wrapper = VaeDecodeWrapper(vae).eval()
torch.onnx.export(
vae_wrapper,
(dummy_vae_latents,),
os.path.join(args.out_dir, "vae_decoder.onnx"),
input_names=vae_input_names,
output_names=vae_output_names,
dynamic_axes=vae_dynamic_axes,
opset_version=args.opset,
do_constant_folding=True,
)
print("Export finished. Models in", args.out_dir)