(args)
| 271 | |
| 272 | |
| 273 | def main(args): |
| 274 | original_ckpt = load_original_checkpoint(args) |
| 275 | has_guidance = any("guidance" in k for k in original_ckpt) |
| 276 | |
| 277 | if args.transformer: |
| 278 | num_layers = 19 |
| 279 | num_single_layers = 38 |
| 280 | inner_dim = 3072 |
| 281 | mlp_ratio = 4.0 |
| 282 | converted_transformer_state_dict = convert_flux_transformer_checkpoint_to_diffusers( |
| 283 | original_ckpt, num_layers, num_single_layers, inner_dim, mlp_ratio=mlp_ratio |
| 284 | ) |
| 285 | transformer = FluxTransformer2DModel(guidance_embeds=has_guidance) |
| 286 | transformer.load_state_dict(converted_transformer_state_dict, strict=True) |
| 287 | |
| 288 | print( |
| 289 | f"Saving Flux Transformer in Diffusers format. Variant: {'guidance-distilled' if has_guidance else 'timestep-distilled'}" |
| 290 | ) |
| 291 | transformer.to(dtype).save_pretrained(f"{args.output_path}/transformer") |
| 292 | |
| 293 | if args.vae: |
| 294 | config = AutoencoderKL.load_config("stabilityai/stable-diffusion-3-medium-diffusers", subfolder="vae") |
| 295 | vae = AutoencoderKL.from_config(config, scaling_factor=0.3611, shift_factor=0.1159).to(torch.bfloat16) |
| 296 | |
| 297 | converted_vae_state_dict = convert_ldm_vae_checkpoint(original_ckpt, vae.config) |
| 298 | vae.load_state_dict(converted_vae_state_dict, strict=True) |
| 299 | vae.to(dtype).save_pretrained(f"{args.output_path}/vae") |
| 300 | |
| 301 | |
| 302 | if __name__ == "__main__": |
no test coverage detected