Spaces:
Paused
Paused
| # 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) |