Gate VAE fp8 quant by decoded-image accuracy (auto layerwise fp8, fp8_dynamic opt-in)
A B200 decoded-image LPIPS/SSIM sweep vs the dense bf16 VAE (new
scripts/quant_accuracy_sweep.py) settles the two VAE schemes:
- Layerwise fp8 (storage-only) holds across families (SSIM >= 0.977 on all but
SDXL), so auto now engages layerwise fp8 ONLY. For a VAE decode (a few percent
of end to end) fp8_dynamic's fp8-matmul speedup over storage fp8 is negligible,
so auto never takes the accuracy risk.
- fp8_dynamic (torchao PerTensor conv compute) is in-bar on only FLUX.2 and
Hunyuan and out-of-bar or catastrophic elsewhere (Qwen-Image SSIM 0.46), so it
is now an explicit opt-in, re-gated by a per-family deny list derived from the
sweep. SDXL denies both schemes (its small VAE stays dense).
Also fixes a real decode-time crash: torchao 0.17's fp8 conv kernel rejects
pointwise (1x1 / 1x1x1) convs ("Activation and filter channels must match"), so
an explicit fp8_dynamic request cast fine then threw at the first decode on most
families. The conv filter now keeps 1x1 convs dense, the smoke probe uses a
spatial 3x3 conv (so it exercises the path that actually runs), and an explicit
fp8_dynamic request runs that probe before casting.
Tests updated for the fp8-only auto ladder, the 1x1 exclusion, the explicit
probe gate, and the shipped deny list.
This commit is contained in:
parent
c0a4d38eb6
commit
b12a177113
3 changed files with 749 additions and 52 deletions
616
scripts/quant_accuracy_sweep.py
Normal file
616
scripts/quant_accuracy_sweep.py
Normal file
|
|
@ -0,0 +1,616 @@
|
|||
# SPDX-License-Identifier: AGPL-3.0-only
|
||||
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
||||
|
||||
"""Decoded-image accuracy sweep for the auto VAE (and end-to-end) quantisation.
|
||||
|
||||
The VAE is the most quality-sensitive stage of a diffusion pipeline: it turns the DiT
|
||||
latent into RGB pixels, so a coarse fp8 grid on its convs can BAND the output. The new
|
||||
default on fp8-GEMM silicon casts the VAE to torchao PerTensor ``fp8_dynamic`` (Conv2d /
|
||||
Conv3d) or diffusers layerwise ``fp8``. This harness measures exactly what ships: it
|
||||
loads a family's VAE, decodes a FIXED seeded latent set through the dense bf16 VAE
|
||||
(reference) and through the same VAE quantised by the repo's own casters
|
||||
(``core.inference.diffusion_vae_quant.quantize_vae``), then reports decoded-image
|
||||
LPIPS(AlexNet) / PSNR / SSIM of quantised vs dense. A (family, scheme) that exceeds the
|
||||
bar (LPIPS <= 0.05, SSIM >= 0.95) belongs in ``_VAE_FAMILY_SCHEME_DENY``.
|
||||
|
||||
``--mode e2e`` instead runs a full pipeline generate dense-bf16 vs everything-auto
|
||||
(auto transformer + auto text encoder + auto VAE) and reports mean LPIPS over a prompt
|
||||
set, the PyTorch-blog "nearly indistinguishable" (~0.1) composed-defaults check.
|
||||
|
||||
The LPIPS AlexNet net is kept on CPU (or a separate --lpips-device) so it never holds
|
||||
memory on the measured GPU. torch / torchao / diffusers / lpips are imported lazily so
|
||||
``--help`` works on a host without them.
|
||||
|
||||
Examples:
|
||||
python scripts/quant_accuracy_sweep.py --family sdxl flux.1 qwen-image
|
||||
python scripts/quant_accuracy_sweep.py --family ltx-2 --latent-t 3 --latent-hw 32
|
||||
python scripts/quant_accuracy_sweep.py --mode e2e --family flux.1 --e2e-model \\
|
||||
black-forest-labs/FLUX.1-schnell
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import Any, Optional
|
||||
|
||||
# ── env: the ancient bitsandbytes in this venv cannot build for CUDA 13 and hard-raises when
|
||||
# diffusers lazily imports its bnb quantiser. We never use bnb here (quant is torchao / layerwise),
|
||||
# so tell diffusers bnb is unavailable BEFORE any VAE class import, and silence the bnb welcome.
|
||||
os.environ.setdefault("BITSANDBYTES_NOWELCOME", "1")
|
||||
os.environ.setdefault("HF_HOME", "/mnt/disks/unslothai/ubuntu/workspace_81/hf_cache")
|
||||
|
||||
_REPO_ROOT = Path(__file__).resolve().parent.parent
|
||||
_BACKEND_ROOT = _REPO_ROOT / "studio" / "backend"
|
||||
for _p in (str(_BACKEND_ROOT), str(_REPO_ROOT / "scripts")):
|
||||
if _p not in sys.path:
|
||||
sys.path.insert(0, _p)
|
||||
|
||||
|
||||
# ── VAE-only accuracy bars (decoder is the most sensitive stage) ──────────────
|
||||
LPIPS_BAR = 0.05
|
||||
SSIM_BAR = 0.95
|
||||
# End-to-end composed-defaults bar (PyTorch blog "nearly indistinguishable").
|
||||
E2E_LPIPS_BAR = 0.10
|
||||
|
||||
# family -> the diffusers base repo whose ``vae`` subfolder we decode with. Only the VAE
|
||||
# subfolder is fetched; the class is resolved from its config by diffusers AutoModel.
|
||||
_VAE_FAMILIES: dict[str, dict[str, Any]] = {
|
||||
"sdxl": {"repo": "stabilityai/stable-diffusion-xl-base-1.0"},
|
||||
"flux.1": {"repo": "black-forest-labs/FLUX.1-schnell"},
|
||||
"qwen-image": {"repo": "Qwen/Qwen-Image"},
|
||||
"flux.2-klein": {"repo": "black-forest-labs/FLUX.2-klein-4B"},
|
||||
"ltx-2": {"repo": "Lightricks/LTX-2"},
|
||||
"hunyuanvideo-1.5": {"repo": "hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v"},
|
||||
}
|
||||
|
||||
|
||||
def _check_deps() -> None:
|
||||
import importlib.util as ilu
|
||||
|
||||
missing = [m for m in ("torch", "torchao", "diffusers", "lpips", "numpy", "PIL") if not ilu.find_spec(m)]
|
||||
if missing:
|
||||
print(
|
||||
"missing deps: " + ", ".join(missing) + "\n"
|
||||
" uv pip install torch torchao diffusers lpips numpy pillow",
|
||||
file=sys.stderr,
|
||||
flush=True,
|
||||
)
|
||||
raise SystemExit(2)
|
||||
|
||||
|
||||
def _import_diffusers():
|
||||
"""Import diffusers with the bnb quantiser disabled (see env note at top)."""
|
||||
import torch # noqa: F401 (torch/torchao first so their extensions register)
|
||||
import torchao # noqa: F401
|
||||
import diffusers.utils.import_utils as iu
|
||||
|
||||
iu._bitsandbytes_available = False
|
||||
import diffusers
|
||||
|
||||
return diffusers
|
||||
|
||||
|
||||
# ── VAE loading + latent shape introspection ─────────────────────────────────
|
||||
|
||||
|
||||
def _load_vae(repo: str, subfolder: str, device: str):
|
||||
import torch
|
||||
|
||||
diffusers = _import_diffusers()
|
||||
vae = diffusers.AutoModel.from_pretrained(repo, subfolder=subfolder, torch_dtype=torch.bfloat16)
|
||||
vae = vae.to(device).eval()
|
||||
return vae
|
||||
|
||||
|
||||
def _first_decoder_conv(vae: Any):
|
||||
"""Return the decoder's input conv (its in_channels == latent channels, its ndim tells
|
||||
2D vs 3D). Falls back to the first conv anywhere."""
|
||||
from torch import nn
|
||||
|
||||
dec = getattr(vae, "decoder", None)
|
||||
for mod in (dec, vae):
|
||||
if mod is None:
|
||||
continue
|
||||
for m in mod.modules():
|
||||
if isinstance(m, (nn.Conv2d, nn.Conv3d)):
|
||||
return m
|
||||
return None
|
||||
|
||||
|
||||
def _latent_spec(vae: Any) -> tuple[int, bool]:
|
||||
"""(latent_channels, is_3d) for a VAE, from its decoder input conv (robust across
|
||||
AutoencoderKL / QwenImage / Flux2 / LTX2 / HunyuanVideo naming)."""
|
||||
from torch import nn
|
||||
|
||||
conv = _first_decoder_conv(vae)
|
||||
is_3d = isinstance(conv, nn.Conv3d)
|
||||
channels = None
|
||||
for key in ("latent_channels", "z_dim", "in_channels"):
|
||||
v = getattr(getattr(vae, "config", object()), key, None)
|
||||
if isinstance(v, int):
|
||||
channels = v
|
||||
break
|
||||
if conv is not None:
|
||||
channels = conv.in_channels # authoritative: what decode actually consumes
|
||||
return int(channels), bool(is_3d)
|
||||
|
||||
|
||||
def _ref_images(args: argparse.Namespace, size: int) -> list:
|
||||
"""Natural reference photos (resized to size x size) for the encode round-trip."""
|
||||
import glob
|
||||
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
|
||||
ref_dir = Path(args.ref_image_dir)
|
||||
files = sorted(glob.glob(str(ref_dir / "*.jpg")) + glob.glob(str(ref_dir / "*.png")))[: args.num_samples]
|
||||
imgs = []
|
||||
for f in files:
|
||||
im = Image.open(f).convert("RGB").resize((size, size), Image.BICUBIC)
|
||||
imgs.append(np.asarray(im, dtype=np.uint8))
|
||||
return imgs
|
||||
|
||||
|
||||
def _encode_latent(vae: Any, x: Any):
|
||||
"""Encode a preprocessed pixel tensor to a deterministic latent (mode of the posterior
|
||||
when the VAE exposes one), robust across AutoencoderKL / QwenImage / Flux2 / LTX2 / HV15."""
|
||||
import torch
|
||||
|
||||
with torch.no_grad():
|
||||
enc = vae.encode(x)
|
||||
dist = getattr(enc, "latent_dist", None)
|
||||
if dist is not None:
|
||||
return dist.mode() if hasattr(dist, "mode") else dist.sample()
|
||||
if hasattr(enc, "latent"):
|
||||
return enc.latent
|
||||
if isinstance(enc, (tuple, list)):
|
||||
first = enc[0]
|
||||
return first.mode() if hasattr(first, "mode") else first
|
||||
return enc
|
||||
|
||||
|
||||
def _make_latents(vae: Any, args: argparse.Namespace, device: str):
|
||||
"""A fixed latent batch at the family's latent shape (one per sample). Default: encode
|
||||
natural photos through the dense VAE (in-distribution, natural decoded content -- the
|
||||
regime the LPIPS/SSIM bars are calibrated for). ``--latents random`` uses seeded N(0,1)."""
|
||||
import torch
|
||||
|
||||
channels, is_3d = _latent_spec(vae)
|
||||
lat = []
|
||||
if args.latents == "encode":
|
||||
size = args.enc_hw_3d if is_3d else args.enc_hw
|
||||
imgs = _ref_images(args, size)
|
||||
for arr in imgs:
|
||||
x = torch.from_numpy(arr).float().permute(2, 0, 1).unsqueeze(0).div(127.5).sub(1.0)
|
||||
if is_3d:
|
||||
x = x.unsqueeze(2).repeat(1, 1, args.enc_frames, 1, 1) # static clip [1,3,T,H,W]
|
||||
x = x.to(device=device, dtype=torch.bfloat16)
|
||||
lat.append(_encode_latent(vae, x))
|
||||
if lat:
|
||||
return lat, is_3d
|
||||
print(" (no ref images found; falling back to random latents)", flush=True)
|
||||
for seed in range(args.num_samples):
|
||||
g = torch.Generator().manual_seed(1000 + seed)
|
||||
shape = (
|
||||
(1, channels, args.latent_t, args.latent_hw_3d, args.latent_hw_3d)
|
||||
if is_3d
|
||||
else (1, channels, args.latent_hw, args.latent_hw)
|
||||
)
|
||||
z = torch.randn(shape, generator=g, dtype=torch.float32)
|
||||
lat.append(z.to(device=device, dtype=torch.bfloat16))
|
||||
return lat, is_3d
|
||||
|
||||
|
||||
def _decode(vae: Any, z: Any):
|
||||
"""Decode one latent, returning a list of HxWx3 uint8 numpy frames (>1 for a video VAE)."""
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
with torch.no_grad():
|
||||
try:
|
||||
out = vae.decode(z)
|
||||
except TypeError:
|
||||
out = vae.decode(z, return_dict=True)
|
||||
sample = out.sample if hasattr(out, "sample") else out[0]
|
||||
sample = sample.float().clamp(-1, 1)
|
||||
# [B,C,H,W] (image) or [B,C,T,H,W] (video). Emit one frame per temporal slot.
|
||||
frames = []
|
||||
if sample.dim() == 5:
|
||||
b, c, t, h, w = sample.shape
|
||||
for ti in range(t):
|
||||
frames.append(sample[0, :, ti])
|
||||
else:
|
||||
frames.append(sample[0])
|
||||
imgs = []
|
||||
for f in frames:
|
||||
arr = ((f.permute(1, 2, 0).cpu().numpy() + 1.0) * 127.5).round().clip(0, 255).astype(np.uint8)
|
||||
if arr.shape[2] == 1:
|
||||
arr = np.repeat(arr, 3, axis=2)
|
||||
imgs.append(arr)
|
||||
return imgs # list of HxWx3 uint8
|
||||
|
||||
|
||||
# ── metrics ──────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class _Lpips:
|
||||
"""AlexNet LPIPS kept off the measured GPU. Inputs are HxWx3 uint8 arrays mapped to [-1,1]."""
|
||||
|
||||
def __init__(self, device: str = "cpu") -> None:
|
||||
import lpips
|
||||
import torch
|
||||
|
||||
self.torch = torch
|
||||
self.device = device
|
||||
self.fn = lpips.LPIPS(net="alex", verbose=False).to(device).eval()
|
||||
|
||||
def __call__(self, a: Any, b: Any) -> float:
|
||||
t = self.torch
|
||||
|
||||
def to_t(x):
|
||||
return t.from_numpy(x).float().permute(2, 0, 1).unsqueeze(0).div(127.5).sub(1.0).to(self.device)
|
||||
|
||||
with t.no_grad():
|
||||
return float(self.fn(to_t(a), to_t(b)).item())
|
||||
|
||||
|
||||
def _metrics(ref_frames: list, q_frames: list, lp: "_Lpips") -> dict[str, float]:
|
||||
from diffusion_quality import psnr, ssim # pure-numpy PSNR/SSIM
|
||||
from PIL import Image
|
||||
|
||||
ls, ps, ss = [], [], []
|
||||
for a, b in zip(ref_frames, q_frames):
|
||||
ls.append(lp(a, b))
|
||||
ps.append(psnr(Image.fromarray(a), Image.fromarray(b)))
|
||||
ss.append(ssim(Image.fromarray(a), Image.fromarray(b)))
|
||||
|
||||
def _m(xs):
|
||||
fin = [x for x in xs if x != float("inf")]
|
||||
base = fin if fin else xs
|
||||
return round(sum(base) / len(base), 4) if base else None
|
||||
|
||||
return {"lpips": _m(ls), "psnr": _m(ps), "ssim": _m(ss)}
|
||||
|
||||
|
||||
# ── VAE isolation sweep ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _apply_fp8_dynamic_no1x1(vae_q: Any) -> None:
|
||||
"""Diagnostic caster: the shipped PerTensor fp8_dynamic config, but the conv filter
|
||||
ALSO excludes pointwise (1x1 / 1x1x1) convs. torchao 0.17's f8f8bf16_conv kernel
|
||||
rejects pointwise convs ("Activation and filter channels must match"), so the shipped
|
||||
fp8_dynamic caster crashes at decode on any VAE that has a 1x1 conv with %16 channels.
|
||||
This variant isolates the fp8_dynamic MATH accuracy on the convs that DO run, to show
|
||||
whether excluding 1x1 (a recommended caster fix) keeps fp8_dynamic in-bar."""
|
||||
from torch import nn
|
||||
from torchao.quantization import Float8DynamicActivationFloat8WeightConfig, PerTensor, quantize_
|
||||
|
||||
from core.inference.diffusion_vae_quant import _VAE_KEEP_DENSE_TOKENS
|
||||
|
||||
def filter_fn(module: Any, fqn: str = "") -> bool:
|
||||
if not isinstance(module, (nn.Linear, nn.Conv2d, nn.Conv3d)):
|
||||
return False
|
||||
w = getattr(module, "weight", None)
|
||||
if w is None or w.dim() < 2 or w.shape[0] % 16 or w.shape[1] % 16:
|
||||
return False
|
||||
ks = getattr(module, "kernel_size", None)
|
||||
if isinstance(ks, tuple) and all(k == 1 for k in ks): # pointwise conv -> torchao kernel fails
|
||||
return False
|
||||
name = fqn.lower() if fqn else ""
|
||||
return not any(tok in name for tok in _VAE_KEEP_DENSE_TOKENS)
|
||||
|
||||
quantize_(vae_q, Float8DynamicActivationFloat8WeightConfig(granularity=PerTensor()), filter_fn=filter_fn)
|
||||
|
||||
|
||||
def _sweep_vae(args: argparse.Namespace, lp: "_Lpips", out_dir: Path) -> list[dict]:
|
||||
import copy
|
||||
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
|
||||
from core.inference import diffusion_vae_quant as vq
|
||||
|
||||
class _Target:
|
||||
def __init__(self):
|
||||
import torch
|
||||
|
||||
self.device = "cuda"
|
||||
self.dtype = torch.bfloat16
|
||||
|
||||
target = _Target()
|
||||
rows: list[dict] = []
|
||||
for family in args.family:
|
||||
repo = _VAE_FAMILIES[family]["repo"]
|
||||
print(f"\n=== VAE {family} ({repo}) ===", flush=True)
|
||||
t0 = time.time()
|
||||
try:
|
||||
vae = _load_vae(repo, "vae", "cuda")
|
||||
except Exception as exc: # noqa: BLE001
|
||||
print(f" load FAILED: {type(exc).__name__}: {str(exc)[:200]}", flush=True)
|
||||
rows.append({"family": family, "scheme": "-", "error": f"load: {exc}"})
|
||||
continue
|
||||
ch, is_3d = _latent_spec(vae)
|
||||
cls = type(vae).__name__
|
||||
print(f" {cls} latent_ch={ch} {'3D' if is_3d else '2D'} loaded {time.time()-t0:.0f}s", flush=True)
|
||||
|
||||
latents, _ = _make_latents(vae, args, "cuda")
|
||||
ref_by_sample = [_decode(vae, z) for z in latents]
|
||||
|
||||
fam_dir = out_dir / family
|
||||
fam_dir.mkdir(parents=True, exist_ok=True)
|
||||
Image.fromarray(ref_by_sample[0][0]).save(fam_dir / "dense_s0.png")
|
||||
|
||||
for scheme in args.scheme:
|
||||
try:
|
||||
vae_q = copy.deepcopy(vae)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
print(f" [{scheme}] deepcopy FAILED: {exc}", flush=True)
|
||||
continue
|
||||
pipe = type("P", (), {"vae": vae_q})()
|
||||
# "fp8_dynamic_no1x1" is a diagnostic that bypasses quantize_vae to apply the
|
||||
# PerTensor fp8 config with pointwise convs excluded (torchao 0.17 kernel gap).
|
||||
if scheme == "fp8_dynamic_no1x1":
|
||||
try:
|
||||
_apply_fp8_dynamic_no1x1(vae_q)
|
||||
engaged = scheme
|
||||
except Exception as exc: # noqa: BLE001
|
||||
print(f" [{scheme}] apply FAILED: {str(exc)[:120]}", flush=True)
|
||||
del vae_q
|
||||
_empty_cache()
|
||||
continue
|
||||
else:
|
||||
engaged = vq.quantize_vae(
|
||||
pipe, target, mode=scheme, family=family, offload_active=False, force_fp32=False
|
||||
)
|
||||
if engaged != scheme:
|
||||
print(f" [{scheme}] NOT engaged (returned {engaged}); skipping", flush=True)
|
||||
rows.append({"family": family, "vae_class": cls, "scheme": scheme, "verdict": "NOT_ENGAGED"})
|
||||
del vae_q
|
||||
_empty_cache()
|
||||
continue
|
||||
try:
|
||||
q_by_sample = [_decode(vae_q, z) for z in latents]
|
||||
except Exception as exc: # noqa: BLE001 — the shipped caster produced a VAE that crashes at decode
|
||||
emsg = f"{type(exc).__name__}: {str(exc)[:100]}"
|
||||
print(f" [{scheme}] DECODE CRASH: {emsg}", flush=True)
|
||||
rows.append(
|
||||
{"family": family, "vae_class": cls, "scheme": scheme, "verdict": "CRASH", "error": emsg}
|
||||
)
|
||||
del vae_q
|
||||
_empty_cache()
|
||||
continue
|
||||
all_ref = [f for frames in ref_by_sample for f in frames]
|
||||
all_q = [f for frames in q_by_sample for f in frames]
|
||||
m = _metrics(all_ref, all_q, lp)
|
||||
Image.fromarray(q_by_sample[0][0]).save(fam_dir / f"{scheme}_s0.png")
|
||||
lp_pass = m["lpips"] is not None and m["lpips"] <= LPIPS_BAR
|
||||
ss_pass = m["ssim"] is not None and m["ssim"] >= SSIM_BAR
|
||||
verdict = "PASS" if (lp_pass and ss_pass) else "FAIL"
|
||||
row = {
|
||||
"family": family,
|
||||
"vae_class": cls,
|
||||
"scheme": scheme,
|
||||
"n_frames": len(all_ref),
|
||||
**m,
|
||||
"verdict": verdict,
|
||||
}
|
||||
rows.append(row)
|
||||
print(
|
||||
f" [{scheme}] LPIPS={m['lpips']} PSNR={m['psnr']} SSIM={m['ssim']} -> {verdict}",
|
||||
flush=True,
|
||||
)
|
||||
del vae_q
|
||||
_empty_cache()
|
||||
del vae
|
||||
_empty_cache()
|
||||
return rows
|
||||
|
||||
|
||||
def _empty_cache() -> None:
|
||||
try:
|
||||
import torch
|
||||
|
||||
torch.cuda.empty_cache()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
# ── end-to-end (dense bf16 vs everything-auto) ───────────────────────────────
|
||||
|
||||
_E2E_PROMPTS = [
|
||||
"A cozy reading nook by a rain-streaked window, warm lamplight, a cat asleep on a stack of books",
|
||||
"A lone lighthouse on a rocky cliff at sunset, dramatic clouds, crashing waves, highly detailed",
|
||||
"A bustling night market street in the rain, neon signs reflected in puddles, cinematic",
|
||||
"A close-up portrait of an elderly fisherman, weathered skin, soft window light, film grain",
|
||||
"A red fox trotting through a snowy pine forest at dawn, volumetric light",
|
||||
"A steaming bowl of ramen on a wooden table, chopsticks, shallow depth of field",
|
||||
]
|
||||
|
||||
|
||||
def _apply_auto(pipe: Any, family: str, components: list[str]) -> dict[str, Optional[str]]:
|
||||
"""Apply the shipped auto stack in place for the selected components (subset of
|
||||
{transformer, text_encoder, vae}), so the VAE's end-to-end contribution can be isolated."""
|
||||
import torch
|
||||
|
||||
from core.inference import diffusion_precision as dp
|
||||
from core.inference import diffusion_transformer_quant as tq
|
||||
from core.inference import diffusion_vae_quant as vq
|
||||
|
||||
class _Target:
|
||||
device = "cuda"
|
||||
dtype = torch.bfloat16
|
||||
|
||||
tgt = _Target()
|
||||
engaged: dict[str, Optional[str]] = {}
|
||||
if "transformer" in components:
|
||||
engaged["transformer"] = tq.quantize_transformer(pipe, tgt, mode="auto", family=family)
|
||||
if "text_encoder" in components:
|
||||
try:
|
||||
engaged["text_encoder"] = dp.quantize_text_encoders(pipe, tgt, mode="auto", family=family)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
engaged["text_encoder"] = f"err:{type(exc).__name__}"
|
||||
if "vae" in components:
|
||||
engaged["vae"] = vq.quantize_vae(pipe, tgt, mode="auto", family=family)
|
||||
return engaged
|
||||
|
||||
|
||||
def _sweep_e2e(args: argparse.Namespace, lp: "_Lpips", out_dir: Path) -> list[dict]:
|
||||
import torch
|
||||
|
||||
diffusers = _import_diffusers()
|
||||
rows: list[dict] = []
|
||||
for family in args.family:
|
||||
model = args.e2e_model or _VAE_FAMILIES.get(family, {}).get("repo")
|
||||
print(f"\n=== E2E {family} ({model}) ===", flush=True)
|
||||
prompts = args.prompts or _E2E_PROMPTS
|
||||
seeds = args.seeds
|
||||
|
||||
def _gen(pipe):
|
||||
imgs = []
|
||||
for pi, prompt in enumerate(prompts):
|
||||
for seed in seeds:
|
||||
g = torch.Generator(device="cuda").manual_seed(seed)
|
||||
kw = dict(
|
||||
prompt=prompt,
|
||||
num_inference_steps=args.steps,
|
||||
generator=g,
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
)
|
||||
if args.guidance is not None:
|
||||
kw["guidance_scale"] = args.guidance
|
||||
out = pipe(**kw)
|
||||
imgs.append((pi, seed, out.images[0]))
|
||||
return imgs
|
||||
|
||||
try:
|
||||
pipe = diffusers.AutoPipelineForText2Image.from_pretrained(
|
||||
model, torch_dtype=torch.bfloat16
|
||||
).to("cuda")
|
||||
except Exception as exc: # noqa: BLE001
|
||||
print(f" pipe load FAILED: {type(exc).__name__}: {str(exc)[:200]}", flush=True)
|
||||
rows.append({"family": family, "error": f"load: {exc}"})
|
||||
continue
|
||||
ref = _gen(pipe)
|
||||
del pipe
|
||||
_empty_cache()
|
||||
|
||||
pipe2 = diffusers.AutoPipelineForText2Image.from_pretrained(
|
||||
model, torch_dtype=torch.bfloat16
|
||||
).to("cuda")
|
||||
engaged = _apply_auto(pipe2, family, args.e2e_components)
|
||||
print(f" engaged: {engaged}", flush=True)
|
||||
q = _gen(pipe2)
|
||||
del pipe2
|
||||
_empty_cache()
|
||||
|
||||
import numpy as np
|
||||
|
||||
ls = []
|
||||
fam_dir = out_dir / f"e2e_{family}"
|
||||
fam_dir.mkdir(parents=True, exist_ok=True)
|
||||
for (pi, seed, a), (_, _, b) in zip(ref, q):
|
||||
aa, bb = np.asarray(a.convert("RGB")), np.asarray(b.convert("RGB"))
|
||||
ls.append(lp(aa, bb))
|
||||
a.save(fam_dir / f"dense_p{pi}_s{seed}.png")
|
||||
b.save(fam_dir / f"auto_p{pi}_s{seed}.png")
|
||||
mean_l = round(sum(ls) / len(ls), 4) if ls else None
|
||||
verdict = "PASS" if (mean_l is not None and mean_l <= E2E_LPIPS_BAR) else "FAIL"
|
||||
rows.append(
|
||||
{"family": family, "engaged": engaged, "mean_lpips": mean_l, "verdict": verdict}
|
||||
)
|
||||
print(f" mean LPIPS={mean_l} -> {verdict}", flush=True)
|
||||
return rows
|
||||
|
||||
|
||||
# ── output ───────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _write(out_dir: Path, mode: str, rows: list[dict]) -> None:
|
||||
(out_dir / f"{mode}_results.json").write_text(json.dumps(rows, indent=2))
|
||||
print(f"\nwrote {out_dir / f'{mode}_results.json'}", flush=True)
|
||||
print(f"\n=== {mode.upper()} RESULTS ===", flush=True)
|
||||
if mode == "vae":
|
||||
print(f" bars: LPIPS <= {LPIPS_BAR}, SSIM >= {SSIM_BAR}", flush=True)
|
||||
print(f" {'family':<20}{'scheme':<14}{'LPIPS':>9}{'PSNR':>9}{'SSIM':>9} verdict", flush=True)
|
||||
for r in rows:
|
||||
if "error" in r:
|
||||
print(f" {r['family']:<20}{'(error)':<14} {r['error'][:60]}", flush=True)
|
||||
continue
|
||||
print(
|
||||
f" {r['family']:<20}{r['scheme']:<14}{_f(r.get('lpips')):>9}"
|
||||
f"{_f(r.get('psnr')):>9}{_f(r.get('ssim')):>9} {r.get('verdict')}",
|
||||
flush=True,
|
||||
)
|
||||
else:
|
||||
print(f" bar: mean LPIPS <= {E2E_LPIPS_BAR}", flush=True)
|
||||
for r in rows:
|
||||
if "error" in r:
|
||||
print(f" {r['family']}: (error) {r['error'][:80]}", flush=True)
|
||||
continue
|
||||
print(f" {r['family']:<20} mean_lpips={r.get('mean_lpips')} {r.get('verdict')} {r.get('engaged')}", flush=True)
|
||||
|
||||
|
||||
def _f(v: Any) -> str:
|
||||
return f"{v:.4f}" if isinstance(v, (int, float)) else "-"
|
||||
|
||||
|
||||
# ── cli ──────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _build_parser() -> argparse.ArgumentParser:
|
||||
p = argparse.ArgumentParser(
|
||||
description="Decoded-image accuracy sweep for the auto VAE / end-to-end quantisation.",
|
||||
formatter_class=argparse.ArgumentDefaultsHelpFormatter,
|
||||
)
|
||||
p.add_argument("--mode", choices=["vae", "e2e"], default="vae")
|
||||
p.add_argument("--family", nargs="+", default=list(_VAE_FAMILIES.keys()))
|
||||
p.add_argument("--scheme", nargs="+", default=["fp8_dynamic", "fp8_dynamic_no1x1", "fp8"])
|
||||
p.add_argument("--num-samples", type=int, default=5, help="latents to average over")
|
||||
p.add_argument("--latents", choices=["encode", "random"], default="encode",
|
||||
help="encode natural photos (in-distribution) or seeded N(0,1) latents")
|
||||
p.add_argument("--ref-image-dir", default="outputs/quant_accuracy/_refs",
|
||||
help="natural photos to encode for the round-trip")
|
||||
p.add_argument("--enc-hw", type=int, default=512, help="2D encode pixel H=W")
|
||||
p.add_argument("--enc-hw-3d", type=int, default=256, help="3D encode pixel H=W")
|
||||
p.add_argument("--enc-frames", type=int, default=9, help="3D encode pixel frame count")
|
||||
p.add_argument("--latent-hw", type=int, default=64, help="2D random-latent H=W (x8 -> 512px)")
|
||||
p.add_argument("--latent-hw-3d", type=int, default=32, help="3D random-latent H=W")
|
||||
p.add_argument("--latent-t", type=int, default=3, help="3D random-latent temporal length")
|
||||
p.add_argument("--lpips-device", default="cpu", help="device for the LPIPS net (keep off the measured GPU)")
|
||||
p.add_argument("--out-dir", default="outputs/quant_accuracy")
|
||||
# e2e-only
|
||||
p.add_argument("--e2e-model", default=None, help="full model repo for --mode e2e")
|
||||
p.add_argument("--e2e-components", nargs="+", default=["transformer", "text_encoder", "vae"],
|
||||
choices=["transformer", "text_encoder", "vae"],
|
||||
help="which components to auto-quantise for the e2e (isolate the VAE with: --e2e-components vae)")
|
||||
p.add_argument("--prompts", nargs="*", default=None)
|
||||
p.add_argument("--seeds", nargs="*", type=int, default=[12345])
|
||||
p.add_argument("--steps", type=int, default=8)
|
||||
p.add_argument("--guidance", type=float, default=None)
|
||||
p.add_argument("--height", type=int, default=1024)
|
||||
p.add_argument("--width", type=int, default=1024)
|
||||
return p
|
||||
|
||||
|
||||
def main(argv: Optional[list[str]] = None) -> int:
|
||||
args = _build_parser().parse_args(argv)
|
||||
_check_deps()
|
||||
out_dir = Path(args.out_dir).resolve()
|
||||
out_dir.mkdir(parents=True, exist_ok=True)
|
||||
lp = _Lpips(args.lpips_device)
|
||||
if args.mode == "vae":
|
||||
rows = _sweep_vae(args, lp, out_dir)
|
||||
else:
|
||||
rows = _sweep_e2e(args, lp, out_dir)
|
||||
_write(out_dir, args.mode, rows)
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
|
|
@ -12,20 +12,25 @@ int8 path does not apply -- there is no int8 Conv3d kernel. Exactly two schemes
|
|||
Float8DynamicActivationFloat8WeightConfig quantises Conv2d (4D) and
|
||||
Conv3d (5D) weights, not just Linear, and torchao auto-skips any conv
|
||||
whose C_out / C_in is not a multiple of 16 (so the 3-channel RGB
|
||||
``conv_out`` head stays dense). Needs fp8-GEMM silicon (cc >= 8.9) and a
|
||||
resident (non-offloaded) VAE -- the fp8 tensor subclasses reject the
|
||||
Module.to() an offload hook uses.
|
||||
``conv_out`` head stays dense). torchao 0.17's fp8 conv kernel additionally
|
||||
rejects POINTWISE (1x1 / 1x1x1) convs ("Activation and filter channels must
|
||||
match" at decode), so those are kept dense too. Needs fp8-GEMM silicon
|
||||
(cc >= 8.9) and a resident (non-offloaded) VAE -- the fp8 tensor subclasses
|
||||
reject the Module.to() an offload hook uses.
|
||||
fp8 - diffusers layerwise casting: 8-bit (e4m3) STORAGE, upcast per layer to the
|
||||
compute dtype. Storage-only, so it runs on ANY conv (2D / 3D) on any
|
||||
fp8-capable card (cc >= 8.9) and survives group offload.
|
||||
|
||||
There is no int8 (no Conv3d int8 kernel) and no nvfp4 for the VAE. ``auto`` (the loader
|
||||
default) walks (fp8_dynamic, fp8): fp8_dynamic on resident data-center / Ada+ silicon that
|
||||
passes a live conv smoke probe, else layerwise fp8, else dense. ``none``/``off`` keeps the
|
||||
VAE dense bf16; an explicit scheme forces it (re-gated). A quantised VAE MUST NOT be
|
||||
``.to(dtype=...)``'d afterwards (the fp8 tensor subclasses mishandle it), so the loader
|
||||
skips the img2img/inpaint VAE re-align when a scheme engaged. torch / diffusers / torchao
|
||||
are imported lazily so the module stays importable in a no-torch runtime.
|
||||
There is no int8 (no Conv3d int8 kernel) and no nvfp4 for the VAE. A decoded-image
|
||||
LPIPS / SSIM sweep vs the dense bf16 VAE (B200) showed layerwise ``fp8`` holds across
|
||||
families (SSIM >= 0.977 on all but SDXL) while ``fp8_dynamic`` (PerTensor conv compute)
|
||||
only stays in-bar on a couple of VAEs (FLUX.2, Hunyuan) and is catastrophic on others
|
||||
(Qwen-Image SSIM 0.46). So ``auto`` (the loader default) engages layerwise ``fp8`` ONLY --
|
||||
the memory win with no measured quality loss -- and ``fp8_dynamic`` is an explicit opt-in,
|
||||
re-gated by a per-family deny list. ``none``/``off`` keeps the VAE dense bf16. A quantised
|
||||
VAE MUST NOT be ``.to(dtype=...)``'d afterwards (the fp8 tensor subclasses mishandle it), so
|
||||
the loader skips the img2img/inpaint VAE re-align when a scheme engaged. torch / diffusers /
|
||||
torchao are imported lazily so the module stays importable in a no-torch runtime.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
|
@ -46,15 +51,32 @@ VAE_QUANT_MODES = (VAE_QUANT_FP8, VAE_QUANT_FP8_DYNAMIC)
|
|||
# group-norm), whose low-magnitude outputs the coarse fp8 grid would band.
|
||||
_VAE_KEEP_DENSE_TOKENS = ("conv_out", "proj_out", "conv_norm_out", "norm_out")
|
||||
|
||||
# Best-first ``auto`` order: fp8_dynamic (compute fp8 on the conv tensor cores) leads; layerwise
|
||||
# ``fp8`` (storage-only) is the universal fallback and the sole scheme that survives group offload.
|
||||
_VAE_AUTO_LADDER = (VAE_QUANT_FP8_DYNAMIC, VAE_QUANT_FP8)
|
||||
# ``auto`` engages layerwise ``fp8`` ONLY. The decoded-image accuracy sweep found fp8_dynamic
|
||||
# (PerTensor conv compute) in-bar on just FLUX.2 / Hunyuan and out-of-bar or catastrophic
|
||||
# elsewhere, and for a VAE decode (a few % of end-to-end) fp8_dynamic's fp8-matmul speedup over
|
||||
# storage-only fp8 is negligible -- so auto never risks it. fp8_dynamic stays an EXPLICIT opt-in
|
||||
# (reachable via a direct request, re-gated by the deny list below). Layerwise fp8 is storage-only
|
||||
# and is the sole scheme that survives group offload. The loop in select_vae_quant_scheme keeps
|
||||
# its per-scheme gates generic so re-adding fp8_dynamic here later needs no extra plumbing.
|
||||
_VAE_AUTO_LADDER = (VAE_QUANT_FP8,)
|
||||
|
||||
# VAEs whose activation ranges break a scheme at the MODEL level (measured decoded-image
|
||||
# LPIPS / SSIM vs the dense bf16 VAE). Populated from the accuracy sweep; a denied scheme is
|
||||
# skipped by ``auto`` and refused when requested explicitly. Empty by default -- the
|
||||
# vae_force_fp32 video families are gated separately at the loader (they never quantise).
|
||||
_VAE_FAMILY_SCHEME_DENY: dict[str, frozenset[str]] = {}
|
||||
# VAEs whose activation ranges break a scheme at the MODEL level -- from the decoded-image
|
||||
# LPIPS / SSIM sweep vs the dense bf16 VAE (B200; bar LPIPS <= 0.05, SSIM >= 0.95). A denied
|
||||
# scheme is skipped by ``auto`` and refused when requested explicitly.
|
||||
# sdxl - layerwise fp8 marginal (SSIM 0.935) AND fp8_dynamic worse (0.894): deny BOTH,
|
||||
# so its small VAE just stays dense (negligible memory cost).
|
||||
# fp8_dynamic-only denials (layerwise fp8 is fine for these; auto still quantises them):
|
||||
# flux.1 / flux.1-kontext SSIM 0.935, qwen-image / qwen-image-edit SSIM 0.46 (catastrophic),
|
||||
# ltx-2 SSIM 0.942. FLUX.2 (klein/dev) and Hunyuan-1.5 pass fp8_dynamic, so they are not denied.
|
||||
# The vae_force_fp32 video families (Wan) are gated separately at the loader (never quantise).
|
||||
_VAE_FAMILY_SCHEME_DENY: dict[str, frozenset[str]] = {
|
||||
"sdxl": frozenset({VAE_QUANT_FP8, VAE_QUANT_FP8_DYNAMIC}),
|
||||
"flux.1": frozenset({VAE_QUANT_FP8_DYNAMIC}),
|
||||
"flux.1-kontext": frozenset({VAE_QUANT_FP8_DYNAMIC}),
|
||||
"qwen-image": frozenset({VAE_QUANT_FP8_DYNAMIC}),
|
||||
"qwen-image-edit": frozenset({VAE_QUANT_FP8_DYNAMIC}),
|
||||
"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] = {}
|
||||
|
|
@ -107,11 +129,13 @@ def _vae_family_denied(family: Optional[str], scheme: str) -> bool:
|
|||
|
||||
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, 1) (channels a multiple of 16 so torchao does not skip it) with the
|
||||
PerTensor fp8 config and run one forward. Cached per device. This makes ``auto`` 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 ladder
|
||||
falls to layerwise fp8 rather than crashing at the first decode."""
|
||||
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]
|
||||
ok = False
|
||||
|
|
@ -124,7 +148,7 @@ def _vae_fp8_dynamic_probe(device: str) -> bool:
|
|||
quantize_,
|
||||
)
|
||||
|
||||
conv = nn.Conv2d(16, 16, 1).to(device = device, dtype = torch.bfloat16)
|
||||
conv = nn.Conv2d(16, 16, 3, padding = 1).to(device = device, dtype = torch.bfloat16)
|
||||
quantize_(
|
||||
conv,
|
||||
Float8DynamicActivationFloat8WeightConfig(granularity = PerTensor()),
|
||||
|
|
@ -231,6 +255,13 @@ def quantize_vae(
|
|||
return None
|
||||
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
|
||||
vae = getattr(pipe, "vae", None)
|
||||
if vae is None:
|
||||
return None
|
||||
|
|
@ -266,6 +297,12 @@ def _cast_vae_fp8_dynamic(vae: Any, target: Any) -> None:
|
|||
return False
|
||||
if weight.shape[0] % 16 != 0 or weight.shape[1] % 16 != 0:
|
||||
return False
|
||||
# torchao 0.17's fp8 conv kernel (f8f8bf16_conv) rejects POINTWISE (1x1 / 1x1x1) convs
|
||||
# ("Activation and filter channels must match"): they cast fine but crash at the first
|
||||
# decode. Leave them dense (they are cheap 1x1 projections -- little to gain anyway).
|
||||
kernel_size = getattr(module, "kernel_size", None)
|
||||
if isinstance(kernel_size, tuple) and kernel_size and all(k == 1 for k in kernel_size):
|
||||
return False
|
||||
name = fqn.lower() if fqn else ""
|
||||
return not any(tok in name for tok in _VAE_KEEP_DENSE_TOKENS)
|
||||
|
||||
|
|
|
|||
|
|
@ -158,17 +158,21 @@ def test_vae_quant_supported_fp8_dynamic_requires_sm89(monkeypatch):
|
|||
# ── auto ladder (select_vae_quant_scheme) ───────────────────────────────────────
|
||||
|
||||
|
||||
def test_select_datacenter_prefers_fp8_dynamic(monkeypatch):
|
||||
# Data-center fp8-GEMM silicon with a passing conv probe: fp8_dynamic leads the ladder.
|
||||
def test_select_datacenter_uses_layerwise_fp8(monkeypatch):
|
||||
# ``auto`` engages layerwise fp8 ONLY -- the accuracy sweep keeps fp8_dynamic out of the
|
||||
# auto ladder (only in-bar on a couple of VAEs). Even on data-center fp8-GEMM silicon it
|
||||
# resolves to fp8, and the fp8_dynamic conv probe is never consulted for auto.
|
||||
_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)
|
||||
assert select_vae_quant_scheme(_target(), "auto", family = "flux.1") == VAE_QUANT_FP8_DYNAMIC
|
||||
monkeypatch.setattr(
|
||||
vq, "_vae_fp8_dynamic_probe", lambda device: pytest.fail("auto must not probe fp8_dynamic")
|
||||
)
|
||||
assert select_vae_quant_scheme(_target(), "auto", family = "flux.1") == VAE_QUANT_FP8
|
||||
|
||||
|
||||
def test_select_offload_uses_layerwise_fp8(monkeypatch):
|
||||
# Under offload the torchao fp8_dynamic mode (rejects Module.to()) is skipped BEFORE the
|
||||
# probe -> layerwise fp8. The probe must not even run.
|
||||
# Under offload ``auto`` still resolves to layerwise fp8 (the storage-only scheme that
|
||||
# survives Module.to()); the fp8_dynamic conv probe is never consulted.
|
||||
_stub_capability(monkeypatch, (10, 0))
|
||||
_allow_vae(monkeypatch, {VAE_QUANT_FP8_DYNAMIC, VAE_QUANT_FP8})
|
||||
monkeypatch.setattr(
|
||||
|
|
@ -186,14 +190,13 @@ def test_select_force_fp32_stays_dense(monkeypatch):
|
|||
|
||||
|
||||
def test_select_family_deny_skips_scheme(monkeypatch):
|
||||
# The only scheme ``auto`` walks is layerwise fp8, so denying fp8 for a family leaves it
|
||||
# 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_FAMILY_SCHEME_DENY", {"badfam": frozenset({VAE_QUANT_FP8_DYNAMIC})}
|
||||
)
|
||||
# fp8_dynamic denied for this family -> falls to layerwise fp8.
|
||||
assert select_vae_quant_scheme(_target(), "auto", family = "badfam") == VAE_QUANT_FP8
|
||||
monkeypatch.setattr(vq, "_VAE_FAMILY_SCHEME_DENY", {"badfam": frozenset({VAE_QUANT_FP8})})
|
||||
assert select_vae_quant_scheme(_target(), "auto", family = "badfam") is None
|
||||
|
||||
|
||||
def test_select_no_capability_is_none(monkeypatch):
|
||||
|
|
@ -202,19 +205,19 @@ def test_select_no_capability_is_none(monkeypatch):
|
|||
assert select_vae_quant_scheme(_target(), "auto") is None
|
||||
|
||||
|
||||
def test_select_support_gate_falls_to_fp8(monkeypatch):
|
||||
# fp8_dynamic hardware-unsupported (only fp8 allowed) -> layerwise fp8.
|
||||
def test_select_support_gate_uses_fp8(monkeypatch):
|
||||
# fp8-capable silicon with fp8 supported -> ``auto`` resolves to layerwise fp8.
|
||||
_stub_capability(monkeypatch, (8, 9))
|
||||
_allow_vae(monkeypatch, {VAE_QUANT_FP8})
|
||||
monkeypatch.setattr(vq, "_vae_fp8_dynamic_probe", lambda device: True)
|
||||
assert select_vae_quant_scheme(_target(), "auto") == VAE_QUANT_FP8
|
||||
|
||||
|
||||
def test_select_smoke_failure_falls_to_fp8(monkeypatch):
|
||||
# fp8_dynamic is hardware-supported but its conv probe fails -> layerwise fp8.
|
||||
def test_select_auto_never_uses_fp8_dynamic(monkeypatch):
|
||||
# Even with fp8_dynamic hardware-supported and its conv probe passing, ``auto`` stays on
|
||||
# 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: False)
|
||||
monkeypatch.setattr(vq, "_vae_fp8_dynamic_probe", lambda device: True)
|
||||
assert select_vae_quant_scheme(_target(), "auto") == VAE_QUANT_FP8
|
||||
|
||||
|
||||
|
|
@ -247,20 +250,26 @@ def test_fp8_dynamic_conv_filter(monkeypatch):
|
|||
ff = captured["filter_fn"]
|
||||
nn = torch.nn
|
||||
|
||||
def _mod(cls, shape):
|
||||
def _mod(cls, shape, kernel_size = None):
|
||||
m = cls()
|
||||
m.weight = _Weight(shape)
|
||||
if kernel_size is not None:
|
||||
m.kernel_size = kernel_size
|
||||
return m
|
||||
|
||||
# Conv2d (4D) / Conv3d (5D) with both channel dims a multiple of 16 quantise.
|
||||
assert ff(_mod(nn.Conv2d, (128, 128, 3, 3)), "decoder.up.0.resnets.0.conv1") is True
|
||||
assert ff(_mod(nn.Conv3d, (64, 64, 3, 3, 3)), "decoder.mid_block.conv3d") is True
|
||||
# nn.Linear (mid-block attention projections) also quantise.
|
||||
assert ff(_mod(nn.Conv2d, (128, 128, 3, 3), (3, 3)), "decoder.up.0.resnets.0.conv1") is True
|
||||
assert ff(_mod(nn.Conv3d, (64, 64, 3, 3, 3), (3, 3, 3)), "decoder.mid_block.conv3d") is True
|
||||
# nn.Linear (mid-block attention projections) also quantise (no kernel_size attr).
|
||||
assert ff(_mod(nn.Linear, (512, 512)), "decoder.mid_block.attentions.0.to_q") is True
|
||||
# POINTWISE (1x1 / 1x1x1) convs excluded even at %16 channels: torchao 0.17's fp8 conv
|
||||
# kernel rejects them ("Activation and filter channels must match") -> crash at decode.
|
||||
assert ff(_mod(nn.Conv2d, (128, 128, 1, 1), (1, 1)), "decoder.mid_block.attentions.0.proj_conv") is False
|
||||
assert ff(_mod(nn.Conv3d, (64, 64, 1, 1, 1), (1, 1, 1)), "decoder.time_mix.conv") is False
|
||||
# Channels not a multiple of 16 excluded (torchao would skip them regardless): the RGB
|
||||
# in/out head (C=3) and any off-16 dim.
|
||||
assert ff(_mod(nn.Conv2d, (128, 3, 3, 3)), "encoder.conv_in") is False
|
||||
assert ff(_mod(nn.Conv2d, (24, 128, 3, 3)), "decoder.up.1.upsamplers.0.conv") is False
|
||||
assert ff(_mod(nn.Conv2d, (128, 3, 3, 3), (3, 3)), "encoder.conv_in") is False
|
||||
assert ff(_mod(nn.Conv2d, (24, 128, 3, 3), (3, 3)), "decoder.up.1.upsamplers.0.conv") is False
|
||||
# conv_out / proj_out / norm_out excluded by NAME even with %16 channels.
|
||||
assert ff(_mod(nn.Conv2d, (16, 128, 3, 3)), "decoder.conv_out") is False
|
||||
assert ff(_mod(nn.Conv2d, (128, 128, 3, 3)), "decoder.conv_norm_out") is False
|
||||
|
|
@ -375,14 +384,49 @@ def test_quantize_vae_tolerates_caster_failure(monkeypatch):
|
|||
|
||||
|
||||
def test_quantize_vae_auto_resolves_and_applies(monkeypatch):
|
||||
# End-to-end: mode="auto" resolves via the ladder then applies the resolved caster.
|
||||
# End-to-end: mode="auto" resolves via the ladder (layerwise fp8) then applies that caster.
|
||||
_stub_torch(monkeypatch, cc = (10, 0))
|
||||
_stub_capability(monkeypatch, (10, 0))
|
||||
_allow_vae(monkeypatch, {VAE_QUANT_FP8_DYNAMIC, VAE_QUANT_FP8})
|
||||
monkeypatch.setattr(vq, "_cast_vae_fp8_dynamic", lambda v, t: pytest.fail("auto must not use fp8_dynamic"))
|
||||
calls: list = []
|
||||
monkeypatch.setattr(vq, "_cast_vae_fp8", lambda v, t: calls.append(v))
|
||||
vae = object()
|
||||
pipe = types.SimpleNamespace(vae = vae)
|
||||
assert quantize_vae(pipe, _target(), mode = "auto", family = "flux.1") == VAE_QUANT_FP8
|
||||
assert calls == [vae]
|
||||
|
||||
|
||||
def test_quantize_vae_explicit_fp8_dynamic_probe_gates(monkeypatch):
|
||||
# An explicit fp8_dynamic request runs the conv smoke probe: a build whose torchao lacks a
|
||||
# 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, "_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)
|
||||
calls: list = []
|
||||
monkeypatch.setattr(vq, "_cast_vae_fp8_dynamic", lambda v, t: calls.append(v))
|
||||
vae = object()
|
||||
pipe = types.SimpleNamespace(vae = vae)
|
||||
assert quantize_vae(pipe, _target(), mode = "auto", family = "flux.1") == VAE_QUANT_FP8_DYNAMIC
|
||||
assert calls == [vae]
|
||||
assert quantize_vae(pipe, _target(), mode = "fp8_dynamic", family = "flux.2-klein") == VAE_QUANT_FP8_DYNAMIC
|
||||
assert len(calls) == 1
|
||||
|
||||
|
||||
def test_real_family_deny_list_policy(monkeypatch):
|
||||
# The shipped _VAE_FAMILY_SCHEME_DENY (from the B200 sweep), exercised through the real
|
||||
# select path -- no deny-list monkeypatch. Confirms the per-family auto/explicit outcomes.
|
||||
_stub_capability(monkeypatch, (10, 0))
|
||||
_allow_vae(monkeypatch, {VAE_QUANT_FP8_DYNAMIC, VAE_QUANT_FP8})
|
||||
# SDXL denies BOTH schemes -> auto stays dense (its layerwise fp8 was marginal).
|
||||
assert select_vae_quant_scheme(_target(), "auto", family = "sdxl") is None
|
||||
assert select_vae_quant_scheme(_target(), "fp8", family = "sdxl") is None
|
||||
# Qwen-Image: auto -> layerwise fp8 (safe); explicit fp8_dynamic refused (catastrophic).
|
||||
assert select_vae_quant_scheme(_target(), "auto", family = "qwen-image") == VAE_QUANT_FP8
|
||||
assert select_vae_quant_scheme(_target(), "fp8_dynamic", family = "qwen-image") is None
|
||||
# FLUX.1 / LTX-2 keep layerwise fp8 on auto but deny explicit fp8_dynamic.
|
||||
assert select_vae_quant_scheme(_target(), "auto", family = "ltx-2") == VAE_QUANT_FP8
|
||||
assert select_vae_quant_scheme(_target(), "fp8_dynamic", family = "flux.1") is None
|
||||
# FLUX.2 / Hunyuan keep fp8_dynamic available as an explicit opt-in (measured in-bar).
|
||||
assert select_vae_quant_scheme(_target(), "fp8_dynamic", family = "flux.2-klein") == VAE_QUANT_FP8_DYNAMIC
|
||||
assert select_vae_quant_scheme(_target(), "fp8_dynamic", family = "hunyuanvideo-1.5") == VAE_QUANT_FP8_DYNAMIC
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue