# 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)