File size: 9,830 Bytes
8595af2 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 | """
Hugging Face Spaces / ZeroGPU AoTI helper for Qwen image pipelines.
Usage:
from optimization_optimized import optimize_pipeline_
optimize_pipeline_(
pipe,
image=[Image.new("RGB", (1024, 1024)), Image.new("RGB", (1024, 1024))],
prompt="warmup prompt",
num_inference_steps=4,
true_cfg_scale=1.0,
height=1024,
width=1024,
)
Design goals:
- Keep the same in-place API as your original optimize_pipeline_.
- Avoid crashing the Space if AoTI export/compile fails.
- Only mark transformer dimensions as dynamic when the captured call actually contains them.
- Keep FP8 quantization optional, because it can change quality and may break export on some setups.
"""
from __future__ import annotations
import logging
from contextlib import nullcontext
from typing import Any, Callable, Mapping, Optional, ParamSpec
import torch
from torch.utils._pytree import tree_map
try:
import spaces
except Exception: # Local/dev environment without Hugging Face Spaces.
spaces = None # type: ignore[assignment]
try:
from torchao.quantization import Float8DynamicActivationFloat8WeightConfig, quantize_
except Exception:
Float8DynamicActivationFloat8WeightConfig = None # type: ignore[assignment]
quantize_ = None # type: ignore[assignment]
P = ParamSpec("P")
LOGGER = logging.getLogger(__name__)
# Qwen Image transformer forward usually has:
# hidden_states: [batch, image_seq, channels]
# encoder_hidden_states: [batch, text_seq, channels]
# encoder_hidden_states_mask: [batch, text_seq]
# image_rotary_emb: tuple(image_rotary_emb, text_rotary_emb)
TRANSFORMER_IMAGE_SEQ_LENGTH_DIM = torch.export.Dim("image_seq_length", min=1)
TRANSFORMER_TEXT_SEQ_LENGTH_DIM = torch.export.Dim("text_seq_length", min=1)
DEFAULT_TRANSFORMER_DYNAMIC_SHAPES: dict[str, Any] = {
"hidden_states": {1: TRANSFORMER_IMAGE_SEQ_LENGTH_DIM},
"encoder_hidden_states": {1: TRANSFORMER_TEXT_SEQ_LENGTH_DIM},
"encoder_hidden_states_mask": {1: TRANSFORMER_TEXT_SEQ_LENGTH_DIM},
"image_rotary_emb": (
{0: TRANSFORMER_IMAGE_SEQ_LENGTH_DIM},
{0: TRANSFORMER_TEXT_SEQ_LENGTH_DIM},
),
}
# Conservative-but-fast defaults for ZeroGPU/H200 AoTI.
# You can override any key through optimize_pipeline_(..., inductor_configs={...}).
DEFAULT_INDUCTOR_CONFIGS: dict[str, Any] = {
"conv_1x1_as_mm": True,
"epilogue_fusion": False,
"coordinate_descent_tuning": True,
"coordinate_descent_check_all_directions": True,
"max_autotune": True,
"triton.cudagraphs": True,
}
def _has_zero_gpu_aoti() -> bool:
return (
spaces is not None
and hasattr(spaces, "GPU")
and hasattr(spaces, "aoti_capture")
and hasattr(spaces, "aoti_compile")
and hasattr(spaces, "aoti_apply")
)
def _supports_float8() -> bool:
"""FP8 is mainly useful on Hopper-class GPUs such as H100/H200."""
if not torch.cuda.is_available():
return False
major, _minor = torch.cuda.get_device_capability()
return major >= 9
def _maybe_disable_progress_bar(pipeline: Callable[..., Any]) -> None:
"""Avoid tqdm/progress side effects during graph capture."""
setter = getattr(pipeline, "set_progress_bar_config", None)
if callable(setter):
try:
setter(disable=True)
except Exception:
pass
def _build_dynamic_shapes(
captured_kwargs: Mapping[str, Any],
shape_specs: Optional[Mapping[str, Any]] = None,
) -> dict[str, Any]:
"""
Build a dynamic_shapes dict matching the captured transformer kwargs.
torch.export expects the dynamic_shapes tree to mirror args/kwargs. The original
code unconditionally OR'ed every known Qwen key into the captured tree. This
version only injects a dynamic spec when that key is present in the actual
captured transformer call, which makes the helper more robust across Qwen /
Diffusers versions.
"""
dynamic_shapes = tree_map(lambda _leaf: None, dict(captured_kwargs))
specs = dict(shape_specs or DEFAULT_TRANSFORMER_DYNAMIC_SHAPES)
for name, spec in specs.items():
if name not in captured_kwargs:
continue
# image_rotary_emb is expected to be a pair. Skip if a future pipeline
# changes its structure instead of giving torch.export a mismatched tree.
if name == "image_rotary_emb":
value = captured_kwargs[name]
if not isinstance(value, (tuple, list)) or len(value) != 2:
LOGGER.warning(
"Skipping dynamic shape for image_rotary_emb: expected a 2-item tuple/list, got %s",
type(value).__name__,
)
continue
dynamic_shapes[name] = spec
return dynamic_shapes
def _maybe_quantize_transformer(transformer: torch.nn.Module, enable_float8: bool) -> None:
"""
Optional in-place FP8 quantization.
Keep disabled by default:
- It can affect image quality.
- It requires torchao.
- It is most beneficial on H100/H200-class GPUs.
- If export fails after quantization, the module has still been mutated.
"""
if not enable_float8:
return
if quantize_ is None or Float8DynamicActivationFloat8WeightConfig is None:
raise RuntimeError("enable_float8=True requires torchao to be installed.")
if not _supports_float8():
raise RuntimeError("enable_float8=True requires a Hopper-class CUDA GPU such as H100/H200.")
LOGGER.info("Applying experimental FP8 dynamic activation + FP8 weight quantization.")
quantize_(transformer, Float8DynamicActivationFloat8WeightConfig())
def optimize_pipeline_(
pipeline: Callable[P, Any],
*args: P.args,
duration: int = 1500,
inductor_configs: Optional[Mapping[str, Any]] = None,
dynamic_shape_specs: Optional[Mapping[str, Any]] = None,
enable_float8: bool = False,
strict_export: bool = False,
capture_autocast_dtype: Optional[torch.dtype] = None,
fail_silently: bool = True,
**kwargs: P.kwargs,
) -> Callable[P, Any]:
"""
Compile and apply AoTI to pipeline.transformer in-place.
Parameters:
pipeline:
Diffusers-style pipeline with a `.transformer` module.
*args, **kwargs:
Warmup call passed to the pipeline during AoTI capture.
Use representative image size, prompt length, step count, and guidance.
duration:
ZeroGPU allocation duration for compilation.
inductor_configs:
Optional overrides for DEFAULT_INDUCTOR_CONFIGS.
dynamic_shape_specs:
Optional overrides for dynamic-shape specs.
enable_float8:
Optional experimental torchao FP8 quantization before export.
strict_export:
Passed to torch.export.export(..., strict=...).
False is more permissive for complex model code.
capture_autocast_dtype:
Set to torch.bfloat16 or torch.float16 if your real inference path uses autocast.
Leave None when the pipeline/model is already loaded in the desired dtype.
fail_silently:
If True, log and keep the original transformer on compile failure.
If False, re-raise the error.
Returns:
The same pipeline object, mutated in-place if compile/apply succeeds.
"""
if not _has_zero_gpu_aoti():
LOGGER.warning("Hugging Face Spaces AoTI APIs are unavailable; skipping optimization.")
return pipeline
transformer = getattr(pipeline, "transformer", None)
if transformer is None:
message = "pipeline has no `.transformer` attribute; cannot apply transformer AoTI."
if fail_silently:
LOGGER.warning(message)
return pipeline
raise AttributeError(message)
if isinstance(transformer, torch.nn.Module):
transformer.eval()
_maybe_disable_progress_bar(pipeline)
configs = dict(DEFAULT_INDUCTOR_CONFIGS)
if inductor_configs:
configs.update(inductor_configs)
# Define inside optimize_pipeline_ so it closes over the actual pipeline and warmup args.
@spaces.GPU(duration=duration) # type: ignore[union-attr]
def compile_transformer():
LOGGER.info("Capturing transformer call for AoTI export.")
autocast_ctx = (
torch.autocast(device_type="cuda", dtype=capture_autocast_dtype)
if capture_autocast_dtype is not None and torch.cuda.is_available()
else nullcontext()
)
with torch.inference_mode(), autocast_ctx:
with spaces.aoti_capture(transformer) as call: # type: ignore[union-attr]
pipeline(*args, **kwargs)
dynamic_shapes = _build_dynamic_shapes(call.kwargs, dynamic_shape_specs)
_maybe_quantize_transformer(transformer, enable_float8)
LOGGER.info("Exporting transformer with torch.export; strict=%s", strict_export)
exported = torch.export.export(
mod=transformer,
args=call.args,
kwargs=call.kwargs,
dynamic_shapes=dynamic_shapes,
strict=strict_export,
)
LOGGER.info("Compiling exported transformer with AoTI.")
return spaces.aoti_compile(exported, configs) # type: ignore[union-attr]
try:
compiled_transformer = compile_transformer()
spaces.aoti_apply(compiled_transformer, transformer) # type: ignore[union-attr]
LOGGER.info("AoTI transformer optimization applied successfully.")
except Exception as exc:
LOGGER.exception("AoTI transformer optimization failed.")
if not fail_silently:
raise
LOGGER.warning("Continuing with the original uncompiled transformer. Error: %s", exc)
return pipeline
|