Probe fp8_dynamic per conv dimensionality and apply bench levers before placement

The fp8_dynamic conv smoke probe only exercised Conv2d, so a torchao build
whose Conv3d kernel path is missing or broken would pass the probe for a
video VAE (HunyuanVideo-1.5), report it quantized, and crash at the first
decode. The probe now runs per (device, conv ndim) and an explicit request
must pass it for every conv dimensionality the target VAE contains; the
auto ladder gate inspects the VAE the same way when one is provided.

The video benchmark moved the fully dense pipeline to CUDA before applying
the configured quant/optimisation levers, the reverse of the production
loader (video.py quantizes before apply_memory_plan). Dense-oversized
configs could OOM where the shipped quantized path loads fine, and
load_peak_gb recorded the dense placement. The bench now builds on CPU,
applies the levers, then places on CUDA and captures the load peak.
This commit is contained in:
Daniel Han 2026-07-10 08:01:24 +00:00
commit 6f887afb2b
3 changed files with 125 additions and 34 deletions

View file

@ -298,7 +298,10 @@ def _build_pipe(repo: str, force_fp32_vae: bool):
if force_fp32_vae:
torch_dtype = {"vae": torch.float32, "default": torch.bfloat16}
pipe = diffusers.DiffusionPipeline.from_pretrained(repo, torch_dtype = torch_dtype)
pipe = pipe.to("cuda")
# Stays on CPU: the caller applies the configured levers FIRST and only then places on
# CUDA, mirroring the loader (video.py quantizes before apply_memory_plan). Placing the
# dense pipeline first would OOM configs whose quantized form fits but whose dense form
# does not, and would record a dense load peak for a quantized row.
if force_fp32_vae and getattr(pipe, "vae", None) is not None:
pipe.vae.to(torch.float32) # belt-and-suspenders; a no-op on the primary path above
return pipe
@ -613,8 +616,10 @@ def _run_config(
torch._dynamo.reset()
except Exception:
pass
# Loader order (video.py): build on CPU, apply quant/optim levers, THEN place on CUDA.
# Placing dense first would OOM configs whose quantized form fits but dense does not, and
# load_peak_gb would record the dense placement instead of the measured config.
pipe = _build_pipe(repo, force_fp32)
load_peak = _peak_gb()
engaged = _apply_levers(
pipe,
cfg,
@ -624,6 +629,9 @@ def _run_config(
default_steps = default_steps,
logger = logger,
)
pipe = pipe.to("cuda")
_sync()
load_peak = _peak_gb()
_empty()
weights_gb = _alloc_gb()
dit_active = engaged["dit"] is not None

View file

@ -81,8 +81,9 @@ _VAE_FAMILY_SCHEME_DENY: dict[str, frozenset[str]] = {
"ltx-2": frozenset({VAE_QUANT_FP8_DYNAMIC}),
}
# Cache of device -> bool for the fp8_dynamic conv smoke probe (run once per device).
_VAE_DYNAMIC_PROBE_CACHE: dict[str, bool] = {}
# Cache of (device, conv ndim) -> bool for the fp8_dynamic conv smoke probe (run once per
# device and dimensionality: the 2D and 3D torchao conv kernels are separate code paths).
_VAE_DYNAMIC_PROBE_CACHE: dict[tuple[str, int], bool] = {}
# ``auto`` only quantises a VAE big enough for the saving to be worth the fp8-decode overhead.
# A decoded-image speed/memory sweep showed the split is by kind: image AutoencoderKLs are
@ -148,17 +149,20 @@ def _vae_param_bytes(vae: Any) -> int:
return _VAE_AUTO_MIN_BYTES
def _vae_fp8_dynamic_probe(device: str) -> bool:
"""True iff torchao's fp8_dynamic CONV path runs on this build: quantise a tiny
Conv2d(16, 16, 3, padding=1) (channels a multiple of 16 so torchao does not skip it;
a SPATIAL 3x3 kernel, NOT 1x1 -- torchao 0.17's fp8 conv kernel rejects pointwise convs,
which is the path _cast_vae_fp8_dynamic actually casts) with the PerTensor fp8 config and
run one forward. Cached per device. This makes an explicit fp8_dynamic request robust to a
torchao build whose Float8 config lacks the conv path (the Linear-only fp8 the transformer
probes does not prove the conv path works) -- it fails here and the request stays dense
rather than crashing at the first decode."""
if device in _VAE_DYNAMIC_PROBE_CACHE:
return _VAE_DYNAMIC_PROBE_CACHE[device]
def _vae_fp8_dynamic_probe(device: str, ndim: int = 2) -> bool:
"""True iff torchao's fp8_dynamic CONV path runs on this build for ``ndim``-D convolutions:
quantise a tiny Conv2d/Conv3d(16, 16, 3, padding=1) (channels a multiple of 16 so torchao
does not skip it; a SPATIAL 3x3(x3) kernel, NOT 1x1 -- torchao 0.17's fp8 conv kernel
rejects pointwise convs, which is the path _cast_vae_fp8_dynamic actually casts) with the
PerTensor fp8 config and run one forward. Cached per (device, ndim): the 2D and 3D conv
kernels are separate torchao code paths, so a working Conv2d does not prove Conv3d (the
video decoders) works. This makes an explicit fp8_dynamic request robust to a torchao build
whose Float8 config lacks the conv path (the Linear-only fp8 the transformer probes does not
prove the conv path works) -- it fails here and the request stays dense rather than crashing
at the first decode."""
key = (device, ndim)
if key in _VAE_DYNAMIC_PROBE_CACHE:
return _VAE_DYNAMIC_PROBE_CACHE[key]
ok = False
try:
import torch
@ -169,23 +173,42 @@ def _vae_fp8_dynamic_probe(device: str) -> bool:
quantize_,
)
conv = nn.Conv2d(16, 16, 3, padding = 1).to(device = device, dtype = torch.bfloat16)
conv_cls = nn.Conv3d if ndim == 3 else nn.Conv2d
conv = conv_cls(16, 16, 3, padding = 1).to(device = device, dtype = torch.bfloat16)
quantize_(
conv,
Float8DynamicActivationFloat8WeightConfig(granularity = PerTensor()),
filter_fn = lambda m, fqn = "": isinstance(m, nn.Conv2d),
filter_fn = lambda m, fqn = "": isinstance(m, conv_cls),
)
x = torch.randn(1, 16, 4, 4, device = device, dtype = torch.bfloat16)
x = torch.randn(1, 16, *([4] * ndim), device = device, dtype = torch.bfloat16)
with torch.no_grad():
conv(x)
torch.cuda.synchronize()
ok = True
except Exception:
ok = False
_VAE_DYNAMIC_PROBE_CACHE[device] = ok
_VAE_DYNAMIC_PROBE_CACHE[key] = ok
return ok
def _vae_conv_ndims(vae: Any) -> tuple[int, ...]:
"""The conv dimensionalities ``_cast_vae_fp8_dynamic`` would quantise in this VAE: (2,) for
image AutoencoderKLs, (3,) or (2, 3) for the video Conv3d decoders. Falls back to (2,) when
the VAE cannot be inspected so the gate never silently widens."""
ndims: set[int] = set()
try:
from torch import nn
for module in vae.modules():
if isinstance(module, nn.Conv3d):
ndims.add(3)
elif isinstance(module, nn.Conv2d):
ndims.add(2)
except Exception:
ndims = set()
return tuple(sorted(ndims)) or (2,)
def select_vae_quant_scheme(
target: Any,
requested: Optional[str],
@ -193,6 +216,7 @@ def select_vae_quant_scheme(
family: Optional[str] = None,
offload_active: bool = False,
force_fp32: bool = False,
vae: Any = None,
) -> Optional[str]:
"""Resolve the concrete VAE scheme to apply, or None to stay dense bf16.
@ -202,7 +226,8 @@ def select_vae_quant_scheme(
comment) and returns the first scheme that: survives the active offload policy (torchao
fp8_dynamic tensors reject the Module.to() an offload hook uses -> only layerwise fp8 under
offload), is not family-denied, is hardware-supported, and (fp8_dynamic only) passes a real
conv smoke probe. Returns None when nothing qualifies (e.g. no CUDA / pre-Ada)."""
conv smoke probe for every conv dimensionality the ``vae`` contains (Conv2d when the VAE is
not provided). Returns None when nothing qualifies (e.g. no CUDA / pre-Ada)."""
requested = normalize_vae_quant(requested)
if requested is None or force_fp32:
return None
@ -222,8 +247,13 @@ def select_vae_quant_scheme(
continue
if not vae_quant_supported(target, scheme):
continue
# fp8_dynamic additionally needs the torchao CONV fp8 kernel to actually run on this build.
if scheme == VAE_QUANT_FP8_DYNAMIC and not _vae_fp8_dynamic_probe(device):
# fp8_dynamic additionally needs the torchao CONV fp8 kernel to actually run on this
# build, for every conv dimensionality the VAE contains (Conv3d video decoders are a
# separate kernel path; Conv2d assumed when no VAE was provided to inspect).
if scheme == VAE_QUANT_FP8_DYNAMIC and not all(
_vae_fp8_dynamic_probe(device, ndim)
for ndim in (_vae_conv_ndims(vae) if vae is not None else (2,))
):
continue
return scheme
return None
@ -267,6 +297,7 @@ def quantize_vae(
family = family,
offload_active = offload_active,
force_fp32 = force_fp32,
vae = vae,
)
if mode is None:
return None
@ -290,12 +321,15 @@ def quantize_vae(
if not vae_quant_supported(target, mode):
return None
# fp8_dynamic additionally needs torchao's fp8 CONV kernel to actually run on this build
# (a spatial-conv smoke probe); otherwise the cast succeeds but the first decode crashes.
if mode == VAE_QUANT_FP8_DYNAMIC and not _vae_fp8_dynamic_probe(
str(getattr(target, "device", "cuda"))
):
_note(logger, "vae 'fp8_dynamic' skipped: torchao build lacks a working fp8 conv path")
return None
# for EVERY conv dimensionality this VAE contains (a Conv3d video decoder is a separate
# kernel path from Conv2d); otherwise the cast succeeds but the first decode crashes.
if mode == VAE_QUANT_FP8_DYNAMIC:
device = str(getattr(target, "device", "cuda"))
if not all(_vae_fp8_dynamic_probe(device, ndim) for ndim in _vae_conv_ndims(vae)):
_note(
logger, "vae 'fp8_dynamic' skipped: torchao build lacks a working fp8 conv path"
)
return None
try:
if mode == VAE_QUANT_FP8_DYNAMIC:
_cast_vae_fp8_dynamic(vae, target)

View file

@ -165,7 +165,7 @@ def test_select_datacenter_uses_layerwise_fp8(monkeypatch):
_stub_capability(monkeypatch, (10, 0))
_allow_vae(monkeypatch, {VAE_QUANT_FP8_DYNAMIC, VAE_QUANT_FP8})
monkeypatch.setattr(
vq, "_vae_fp8_dynamic_probe", lambda device: pytest.fail("auto must not probe fp8_dynamic")
vq, "_vae_fp8_dynamic_probe", lambda *a: pytest.fail("auto must not probe fp8_dynamic")
)
assert select_vae_quant_scheme(_target(), "auto", family = "flux.1") == VAE_QUANT_FP8
@ -176,7 +176,7 @@ def test_select_offload_uses_layerwise_fp8(monkeypatch):
_stub_capability(monkeypatch, (10, 0))
_allow_vae(monkeypatch, {VAE_QUANT_FP8_DYNAMIC, VAE_QUANT_FP8})
monkeypatch.setattr(
vq, "_vae_fp8_dynamic_probe", lambda device: pytest.fail("probe must not run under offload")
vq, "_vae_fp8_dynamic_probe", lambda *a: pytest.fail("probe must not run under offload")
)
assert select_vae_quant_scheme(_target(), "auto", offload_active = True) == VAE_QUANT_FP8
@ -194,7 +194,7 @@ def test_select_family_deny_skips_scheme(monkeypatch):
# dense (None). (This is the SDXL case in the real deny list: fp8 marginal -> stay dense.)
_stub_capability(monkeypatch, (10, 0))
_allow_vae(monkeypatch, {VAE_QUANT_FP8_DYNAMIC, VAE_QUANT_FP8})
monkeypatch.setattr(vq, "_vae_fp8_dynamic_probe", lambda device: True)
monkeypatch.setattr(vq, "_vae_fp8_dynamic_probe", lambda device, ndim = 2: True)
monkeypatch.setattr(vq, "_VAE_FAMILY_SCHEME_DENY", {"badfam": frozenset({VAE_QUANT_FP8})})
assert select_vae_quant_scheme(_target(), "auto", family = "badfam") is None
@ -217,7 +217,7 @@ def test_select_auto_never_uses_fp8_dynamic(monkeypatch):
# layerwise fp8: fp8_dynamic is deliberately kept out of the auto ladder (explicit opt-in).
_stub_capability(monkeypatch, (10, 0))
_allow_vae(monkeypatch, {VAE_QUANT_FP8_DYNAMIC, VAE_QUANT_FP8})
monkeypatch.setattr(vq, "_vae_fp8_dynamic_probe", lambda device: True)
monkeypatch.setattr(vq, "_vae_fp8_dynamic_probe", lambda device, ndim = 2: True)
assert select_vae_quant_scheme(_target(), "auto") == VAE_QUANT_FP8
@ -452,13 +452,13 @@ def test_quantize_vae_explicit_fp8_dynamic_probe_gates(monkeypatch):
# working fp8 conv path (probe False) stays dense; a passing probe applies the caster.
_stub_torch(monkeypatch, cc = (10, 0))
_allow_vae(monkeypatch, {VAE_QUANT_FP8_DYNAMIC, VAE_QUANT_FP8})
monkeypatch.setattr(vq, "_vae_fp8_dynamic_probe", lambda device: False)
monkeypatch.setattr(vq, "_vae_fp8_dynamic_probe", lambda device, ndim = 2: False)
monkeypatch.setattr(
vq, "_cast_vae_fp8_dynamic", lambda v, t: pytest.fail("probe failed: must not cast")
)
pipe = types.SimpleNamespace(vae = object())
assert quantize_vae(pipe, _target(), mode = "fp8_dynamic", family = "flux.2-klein") is None
monkeypatch.setattr(vq, "_vae_fp8_dynamic_probe", lambda device: True)
monkeypatch.setattr(vq, "_vae_fp8_dynamic_probe", lambda device, ndim = 2: True)
calls: list = []
monkeypatch.setattr(vq, "_cast_vae_fp8_dynamic", lambda v, t: calls.append(v))
assert (
@ -491,3 +491,52 @@ def test_real_family_deny_list_policy(monkeypatch):
select_vae_quant_scheme(_target(), "fp8_dynamic", family = "hunyuanvideo-1.5")
== VAE_QUANT_FP8_DYNAMIC
)
# ── conv-dimensionality probe gating ─────────────────────────────────────────────
def _conv_vae(torch, *ndims):
"""A fake VAE whose .modules() yields stub Conv2d/Conv3d instances for ``ndims``."""
mods = [(torch.nn.Conv3d if n == 3 else torch.nn.Conv2d)() for n in ndims]
return types.SimpleNamespace(modules = lambda: iter(mods))
def test_vae_conv_ndims(monkeypatch):
torch = _stub_torch(monkeypatch)
assert vq._vae_conv_ndims(_conv_vae(torch, 2)) == (2,)
assert vq._vae_conv_ndims(_conv_vae(torch, 3)) == (3,)
assert vq._vae_conv_ndims(_conv_vae(torch, 2, 3)) == (2, 3)
# A VAE that cannot be inspected (no .modules()) falls back to the 2D probe.
assert vq._vae_conv_ndims(object()) == (2,)
def test_quantize_vae_fp8_dynamic_probes_conv3d_for_video_vae(monkeypatch):
# A Conv3d (video) VAE must pass the 3D conv probe: a torchao build whose Conv2d fp8 path
# works but whose Conv3d path is broken would otherwise report the VAE quantised and crash
# at the first video decode. Probe ok for 2D but not 3D -> the request stays dense.
torch = _stub_torch(monkeypatch, cc = (10, 0))
_allow_vae(monkeypatch, {VAE_QUANT_FP8_DYNAMIC, VAE_QUANT_FP8})
probed: list = []
def _probe(device, ndim = 2):
probed.append(ndim)
return ndim == 2
monkeypatch.setattr(vq, "_vae_fp8_dynamic_probe", _probe)
monkeypatch.setattr(
vq, "_cast_vae_fp8_dynamic", lambda v, t: pytest.fail("3D probe failed: must not cast")
)
pipe = types.SimpleNamespace(vae = _conv_vae(torch, 2, 3))
assert quantize_vae(pipe, _target(), mode = "fp8_dynamic", family = "hunyuanvideo-1.5") is None
assert 3 in probed
# A build whose 3D path also works casts the video VAE.
monkeypatch.setattr(vq, "_vae_fp8_dynamic_probe", lambda device, ndim = 2: True)
pipe = types.SimpleNamespace(vae = _conv_vae(torch, 2, 3))
calls: list = []
monkeypatch.setattr(vq, "_cast_vae_fp8_dynamic", lambda v, t: calls.append(v))
assert (
quantize_vae(pipe, _target(), mode = "fp8_dynamic", family = "hunyuanvideo-1.5")
== VAE_QUANT_FP8_DYNAMIC
)
assert len(calls) == 1