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:
parent
14eba1c887
commit
6f887afb2b
3 changed files with 125 additions and 34 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue