6363
6464
6565ATTENTION_BACKEND_CHOICES = ("default" , * sorted (b .value for b in _HUB_KERNELS_REGISTRY ))
66+ COMPILE_MODE_CHOICES = ("regional" , "full" )
67+ # zlib level for PNG outputs. Pillow's default of 6 takes close to four times as long to encode a
68+ # 1024x1024 image for a file about 7% smaller; the output is lossless either way.
69+ PNG_COMPRESS_LEVEL = 1
70+
71+ # `torch.compile` modes that replay CUDA graphs. Repeated blocks share one compiled graph, so each
72+ # block's output is overwritten when the next block replays it.
73+ _CUDA_GRAPH_COMPILE_MODES = ("max-autotune" , "reduce-overhead" )
6674
6775# Kwarg keys whose string value gets auto-loaded before being passed to the pipeline call.
6876# Images resolve via `diffusers.utils.load_image` → PIL.Image.Image; videos resolve via
@@ -211,13 +219,23 @@ def _add_optimization_arguments(parser: ArgumentParser) -> None:
211219 default = None ,
212220 metavar = "JSON" ,
213221 help = (
214- "torch.compile every denoiser submodule on the pipeline. Accepts an optional JSON "
222+ "torch.compile every denoiser submodule on the pipeline, and the VAE decoder . Accepts an optional JSON "
215223 'object of kwargs forwarded to `torch.compile`, e.g. \' {"mode": "max-autotune", '
216224 '"fullgraph": true}\' . Bare `--compile` uses `fullgraph=true`. Adds a one-time '
217225 "compilation cost on the first step but speeds up every subsequent step — worth it "
218226 "for multi-step generation (50+ steps)."
219227 ),
220228 )
229+ parser .add_argument (
230+ "--compile-mode" ,
231+ choices = COMPILE_MODE_CHOICES ,
232+ default = "regional" ,
233+ help = (
234+ "What --compile compiles. 'regional' (default) compiles only the repeated blocks of a denoiser "
235+ "that declares them, which keeps the first step fast. 'full' compiles the whole denoiser, which "
236+ "`torch.compile` modes that use CUDA graphs (max-autotune, reduce-overhead) require."
237+ ),
238+ )
221239
222240
223241def _add_output_arguments (parser : ArgumentParser ) -> None :
@@ -461,26 +479,44 @@ def _apply_optimizations(pipeline: Any, args: Namespace) -> None:
461479 if args .context_parallel :
462480 logger .warning ("--compile is currently not supported with --context-parallel; skipping compile." )
463481 else :
464- _compile_denoiser (pipeline , args .compile )
482+ _compile_denoiser (pipeline , args .compile , args .compile_mode )
483+ _compile_vae_decoder (pipeline , args .compile )
484+
485+
486+ def _parse_compile_spec (compile_spec : str ) -> dict [str , Any ]:
487+ try :
488+ compile_kwargs = json .loads (compile_spec )
489+ except json .JSONDecodeError as e :
490+ raise SystemExit (f"--compile must be valid JSON: { e } " ) from e
491+ if not isinstance (compile_kwargs , dict ):
492+ raise SystemExit ("--compile must decode to a JSON object." )
493+ return compile_kwargs
465494
466495
467- def _compile_denoiser (pipeline : Any , compile_spec : str ) -> None :
496+ def _compile_vae_decoder (pipeline : Any , compile_spec : str ) -> None :
497+ """Compile the decoder of the pipeline's VAE with the `--compile` kwargs.
498+
499+ The decoder runs once per call, so it has no repeated blocks to compile regionally and is always compiled whole.
500+ """
501+ decoder = getattr (getattr (pipeline , "vae" , None ), "decoder" , None )
502+ if not isinstance (decoder , torch .nn .Module ):
503+ return
504+ pipeline .vae .decoder = torch .compile (decoder , ** _parse_compile_spec (compile_spec ))
505+
506+
507+ def _compile_denoiser (pipeline : Any , compile_spec : str , compile_mode : str = "regional" ) -> None :
468508 """Compile every `transformer*` and `unet*` submodule on the pipeline.
469509
470510 `compile_spec` is the raw JSON string from `--compile` (`"{}"` for bare flag). Decoded into kwargs and forwarded
471511 verbatim to the compile call.
472512
473- Prefers regional compilation via `module.compile_repeated_blocks(**kwargs)` — only compiles the repeated inner
474- blocks (the bulk of the compute), much faster first-step latency than compiling the whole module. Falls back to
475- full `torch.compile` if the model doesn't expose `_repeated_blocks`.
513+ With `compile_mode="regional"`, prefers regional compilation via `module.compile_repeated_blocks(**kwargs)` — only
514+ compiles the repeated inner blocks (the bulk of the compute), much faster first-step latency than compiling the
515+ whole module. Falls back to full `torch.compile` if the model doesn't expose `_repeated_blocks`. With
516+ `compile_mode="full"`, always compiles the whole module.
476517 """
477518
478- try :
479- compile_kwargs = json .loads (compile_spec )
480- except json .JSONDecodeError as e :
481- raise SystemExit (f"--compile must be valid JSON: { e } " ) from e
482- if not isinstance (compile_kwargs , dict ):
483- raise SystemExit ("--compile must decode to a JSON object." )
519+ compile_kwargs = _parse_compile_spec (compile_spec )
484520
485521 for attr in dir (pipeline ):
486522 if not any (attr .startswith (key ) for key in _DENOISER_COMPONENT_KEYS ):
@@ -489,11 +525,17 @@ def _compile_denoiser(pipeline: Any, compile_spec: str) -> None:
489525 if not isinstance (module , torch .nn .Module ):
490526 continue
491527
492- if getattr (module , "_repeated_blocks" , None ):
528+ if compile_mode == "regional" and getattr (module , "_repeated_blocks" , None ):
529+ if compile_kwargs .get ("mode" ) in _CUDA_GRAPH_COMPILE_MODES :
530+ raise SystemExit (
531+ f"--compile mode { compile_kwargs ['mode' ]!r} uses CUDA graphs, which fail when the repeated "
532+ f"blocks of { type (module ).__name__ } are compiled separately. Add `--compile-mode full`, or "
533+ "use mode 'max-autotune-no-cudagraphs'."
534+ )
493535 # Regional compile — only the repeated blocks. Mutates `module` in place.
494536 module .compile_repeated_blocks (** compile_kwargs )
495537 else :
496- # No regional metadata declared; fall back to compiling the whole module .
538+ # Full compile was asked for, or no regional metadata is declared .
497539 setattr (pipeline , attr , torch .compile (module , ** compile_kwargs ))
498540
499541
@@ -535,6 +577,9 @@ def _load_lora(pipeline: Any, args: Namespace) -> None:
535577
536578
537579def _load_pipeline (args : Namespace ) -> Any :
580+ if args .compile is None and args .compile_mode != "regional" :
581+ raise SystemExit (f"--compile-mode { args .compile_mode } only applies together with --compile." )
582+
538583 # Detect modular repos by trying the standard config; `ModularPipeline` repos ship
539584 # `modular_model_index.json` instead of `model_index.json`, so `load_config` OSErrors.
540585 # A repo can also ship a `model_index.json` whose `_class_name` is a modular pipeline
@@ -800,7 +845,7 @@ def _save_videos(videos: list[Any], args: Namespace) -> list[str]:
800845 arr = (np .clip (arr , 0.0 , 1.0 ) * 255 ).round ().astype (np .uint8 )
801846 frame = Image .fromarray (arr )
802847 frame_path = frames_dir / f"{ i :04d} .png"
803- frame .save (frame_path )
848+ frame .save (frame_path , compress_level = PNG_COMPRESS_LEVEL )
804849 saved .append (str (frame_path ))
805850 try :
806851 export_to_video (frames , str (path ), fps = args .fps )
@@ -837,14 +882,14 @@ def _save_output(value: Any, args: Namespace) -> list[str]:
837882 for arr , path in zip (value , paths ):
838883 if arr .dtype != np .uint8 :
839884 arr = (np .clip (arr , 0.0 , 1.0 ) * 255 ).round ().astype (np .uint8 )
840- Image .fromarray (arr ).save (path )
885+ Image .fromarray (arr ).save (path , compress_level = PNG_COMPRESS_LEVEL )
841886 return [str (p ) for p in paths ]
842887
843888 pil_images = _as_pil_list (value )
844889 if pil_images is not None :
845890 paths = _resolve_output_paths (len (pil_images ), args .output , ext = "png" )
846891 for img , path in zip (pil_images , paths ):
847- img .save (path )
892+ img .save (path , compress_level = PNG_COMPRESS_LEVEL )
848893 return [str (p ) for p in paths ]
849894
850895 frames = _as_frame_sequence (value )
0 commit comments