"""Export the cached official SDXL Refiner UNet as unfused opset-14 ONNX.""" from __future__ import annotations import argparse from pathlib import Path import torch from diffusers import UNet2DConditionModel class RefinerUnetWrapper(torch.nn.Module): def __init__(self, unet: UNet2DConditionModel) -> None: super().__init__() self.unet = unet def forward( self, sample: torch.Tensor, timestep: torch.Tensor, encoder_hidden_states: torch.Tensor, text_embeds: torch.Tensor, time_ids: torch.Tensor, ) -> torch.Tensor: return self.unet( sample=sample, timestep=timestep, encoder_hidden_states=encoder_hidden_states, added_cond_kwargs={ "text_embeds": text_embeds, "time_ids": time_ids, }, return_dict=False, )[0] def main() -> None: parser = argparse.ArgumentParser() parser.add_argument("snapshot", type=Path) parser.add_argument("output", type=Path) args = parser.parse_args() unet = UNet2DConditionModel.from_pretrained( args.snapshot, subfolder="unet", torch_dtype=torch.float32, local_files_only=True, ) unet.eval() wrapper = RefinerUnetWrapper(unet).eval() sample = torch.randn(1, 4, 64, 64, dtype=torch.float32) timestep = torch.tensor(151.0, dtype=torch.float32) encoder_hidden_states = torch.randn( 1, 77, 1280, dtype=torch.float32 ) text_embeds = torch.randn(1, 1280, dtype=torch.float32) time_ids = torch.tensor( [[512.0, 512.0, 0.0, 0.0, 6.0]], dtype=torch.float32, ) args.output.parent.mkdir(parents=True, exist_ok=True) with torch.inference_mode(): torch.onnx.export( wrapper, ( sample, timestep, encoder_hidden_states, text_embeds, time_ids, ), str(args.output), input_names=[ "sample", "timestep", "encoder_hidden_states", "text_embeds", "time_ids", ], output_names=["out_sample"], dynamic_axes={ "sample": { 0: "batch_size", 2: "height", 3: "width", }, "encoder_hidden_states": { 0: "batch_size", 1: "sequence_length", }, "text_embeds": {0: "batch_size"}, "time_ids": {0: "batch_size"}, "out_sample": { 0: "batch_size", 2: "height", 3: "width", }, }, opset_version=14, do_constant_folding=True, external_data=True, dynamo=False, ) print(f"Exported SDXL Refiner UNet to {args.output}") if __name__ == "__main__": main()