Merge diffusion-train-tab-2: qwen dense-quant family deny (black frames, measured)

This commit is contained in:
Daniel Han 2026-07-04 08:52:32 +00:00
commit ef03dd2780
4 changed files with 135 additions and 15 deletions

View file

@ -425,7 +425,9 @@ class DiffusionBackend:
target = self._resolve_device_target(fam)
if not dense_transformer_supported(target):
return False
scheme = select_transformer_quant_scheme(target, mode)
scheme = select_transformer_quant_scheme(
target, mode, family = getattr(fam, "name", None)
)
if scheme is None:
return False
source = resolve_prequant_source(
@ -1310,7 +1312,7 @@ class DiffusionBackend:
BEFORE the loader compiles the repeated block, so the order stays quantize ->
compile -> placement."""
# 1. Pre-quantized checkpoint, when one is configured for the resolved scheme.
scheme = select_transformer_quant_scheme(target, mode)
scheme = select_transformer_quant_scheme(target, mode, family = getattr(fam, "name", None))
if scheme is None:
# Bail BEFORE the (multi-GB) dense download: an explicit unsupported scheme
# (e.g. fp8 on Ampere, nvfp4 off Blackwell) would otherwise materialise the
@ -1349,7 +1351,14 @@ class DiffusionBackend:
base, subfolder = "transformer", torch_dtype = dtype, token = hf_token
)
pipe = self._assemble_pipe(pipeline_cls, base, transformer, dtype, hf_token, device)
scheme = quantize_transformer(pipe, target, mode = mode, fast_accum = fast_accum, logger = logger)
scheme = quantize_transformer(
pipe,
target,
mode = mode,
family = getattr(fam, "name", None),
fast_accum = fast_accum,
logger = logger,
)
if scheme is None:
raise RuntimeError("transformer quant unsupported for this device/scheme")
return pipe, scheme

View file

@ -105,6 +105,30 @@ _AUTO_LADDER: tuple[tuple[tuple[int, int], tuple[str, ...]], ...] = (
((8, 0), (TQ_INT8,)), # Ampere sm_80 / sm_86
)
# Families whose activation ranges break specific dense-quant schemes at the MODEL
# level. The kernel smoke probe below cannot see this (it only proves the GEMM runs);
# these were measured with the 28-pair prequant accuracy gate on a B200
# (scripts/prequant_accuracy_gate.py) and reproduced with on-the-fly quantisation:
# qwen-image + fp8 -> every frame black (mean luma 0.0000, SSIM 0.016 vs bf16). The
# same per-row fp8 that matches bf16 on Z-Image / FLUX: Qwen's
# activation outliers exceed even per-row fp8's dynamic range.
# qwen-image + mxfp8 -> real semantic damage at 1024px (CLIP delta mean 0.0146, worst
# cases 0.064 / 0.102 -- 2x the per-case bound).
# qwen-image + nvfp4 -> LPIPS mean 0.51 vs bf16: unusable.
# int8 dynamic (per-token) is excellent on Qwen (LPIPS mean 0.069 / SSIM 0.958), so the
# auto ladder falls through to it. The deny also applies to an EXPLICIT request: a
# scheme that renders black frames has no legitimate use, and returning None gives the
# caller the same fallback contract as an unsupported scheme (GGUF build).
_FAMILY_SCHEME_DENY: dict[str, frozenset[str]] = {
"qwen-image": frozenset({TQ_FP8, TQ_MXFP8, TQ_NVFP4}),
"qwen-image-edit": frozenset({TQ_FP8, TQ_MXFP8, TQ_NVFP4}), # same DiT + activations
}
def _family_denied(family, scheme: str) -> bool:
return scheme in _FAMILY_SCHEME_DENY.get(str(family or "").strip().lower(), ())
# Cache of (scheme, device) -> bool so the quantise+matmul smoke test runs once.
_SMOKE_CACHE: dict[tuple[str, str], bool] = {}
@ -205,18 +229,27 @@ def dense_transformer_supported(target: Any) -> bool:
return False
def select_transformer_quant_scheme(target: Any, requested: Optional[str]) -> Optional[str]:
def select_transformer_quant_scheme(
target: Any,
requested: Optional[str],
family: Optional[str] = None,
) -> Optional[str]:
"""The concrete scheme to apply, or None to fall back to GGUF.
``auto`` walks the per-arch ladder and returns the first scheme that passes a real
quantise+matmul smoke test, so on a box where the Blackwell fp4 / mx kernels are
unavailable it lands on fp8 / int8 with no error. An explicit scheme is honored only
if supported (else None -> GGUF), never silently swapped for a different one."""
if supported (else None -> GGUF), never silently swapped for a different one.
``family`` additionally applies the measured model-level deny list
(``_FAMILY_SCHEME_DENY``): schemes that produce black frames or out-of-bar drift on
that family are skipped by ``auto`` and refused when explicit."""
requested = normalize_transformer_quant(requested)
if requested is None or not dense_transformer_supported(target):
return None
device = str(getattr(target, "device", "cuda"))
if requested != TQ_AUTO:
if _family_denied(family, requested):
return None
return requested if _scheme_supported(requested, device) else None
cap = _capability()
if cap is None:
@ -224,6 +257,8 @@ def select_transformer_quant_scheme(target: Any, requested: Optional[str]) -> Op
for floor, schemes in _AUTO_LADDER:
if cap >= floor:
for scheme in _prefer_consumer_scheme(schemes, device):
if _family_denied(family, scheme):
continue
if _scheme_supported(scheme, device):
return scheme
return None
@ -389,6 +424,7 @@ def quantize_transformer(
target: Any,
*,
mode: Optional[str],
family: Optional[str] = None,
min_features: int = DEFAULT_MIN_LINEAR_FEATURES,
fast_accum: Optional[bool] = None,
logger: Any = None,
@ -400,7 +436,7 @@ def quantize_transformer(
``fast_accum`` (fp8 only) overrides the per-GPU-class accumulate choice: None
auto-detects (fast on consumer, precise on data-center), True/False force it."""
scheme = select_transformer_quant_scheme(target, mode)
scheme = select_transformer_quant_scheme(target, mode, family = family)
if scheme is None:
return None
transformer = getattr(pipe, "transformer", None)

View file

@ -1772,7 +1772,9 @@ def _stub_dense_quant(monkeypatch, *, scheme = "fp8"):
monkeypatch.setattr(dmod, "dense_transformer_supported", lambda target: True)
# Resolve the scheme without the real GPU smoke probe, and configure no pre-quant
# checkpoint so the dense materialise+quantise branch is the one exercised.
monkeypatch.setattr(dmod, "select_transformer_quant_scheme", lambda target, mode: scheme)
monkeypatch.setattr(
dmod, "select_transformer_quant_scheme", lambda target, mode, family = None: scheme
)
monkeypatch.setattr(dmod, "resolve_prequant_source", lambda fam, scheme, **kw: None)
def _quantize(pipe, target, *, mode, **kw):
@ -1836,7 +1838,9 @@ def test_transformer_quant_prequant_path_engaged(fake_runtime, tmp_path, monkeyp
backend = DiffusionBackend()
_force_cuda_target(backend, monkeypatch)
monkeypatch.setattr(dmod, "dense_transformer_supported", lambda target: True)
monkeypatch.setattr(dmod, "select_transformer_quant_scheme", lambda target, mode: "fp8")
monkeypatch.setattr(
dmod, "select_transformer_quant_scheme", lambda target, mode, family = None: "fp8"
)
monkeypatch.setattr(dmod, "resolve_prequant_source", lambda fam, scheme, **kw: object())
prequant_obj = object()
loaded: dict = {"n": 0}
@ -1966,7 +1970,9 @@ def test_transformer_quant_unsupported_scheme_skips_dense_download(
backend = DiffusionBackend()
_force_cuda_target(backend, monkeypatch)
monkeypatch.setattr(dmod, "dense_transformer_supported", lambda target: True)
monkeypatch.setattr(dmod, "select_transformer_quant_scheme", lambda target, mode: None)
monkeypatch.setattr(
dmod, "select_transformer_quant_scheme", lambda target, mode, family = None: None
)
monkeypatch.setattr(dmod, "resolve_prequant_source", lambda fam, scheme, **kw: None)
@classmethod
@ -2013,7 +2019,9 @@ def test_dense_quant_prefetch_needed_gates(fake_runtime, monkeypatch):
_force_cuda_target(backend, monkeypatch)
fam = detect_family("unsloth/Z-Image-Turbo-GGUF")
monkeypatch.setattr(dmod, "dense_transformer_supported", lambda target: True)
monkeypatch.setattr(dmod, "select_transformer_quant_scheme", lambda target, mode: "fp8")
monkeypatch.setattr(
dmod, "select_transformer_quant_scheme", lambda target, mode, family = None: "fp8"
)
monkeypatch.setattr(dmod, "resolve_prequant_source", lambda fam, scheme, **kw: None)
assert backend._dense_quant_prefetch_needed(fam, {"transformer_quant": "fp8"}) is True
@ -2024,10 +2032,14 @@ def test_dense_quant_prefetch_needed_gates(fake_runtime, monkeypatch):
assert backend._dense_quant_prefetch_needed(fam, {"transformer_quant": "fp8"}) is False
# Unsupported scheme bails before the dense path (and so must the prefetch).
monkeypatch.setattr(dmod, "resolve_prequant_source", lambda fam, scheme, **kw: None)
monkeypatch.setattr(dmod, "select_transformer_quant_scheme", lambda target, mode: None)
monkeypatch.setattr(
dmod, "select_transformer_quant_scheme", lambda target, mode, family = None: None
)
assert backend._dense_quant_prefetch_needed(fam, {"transformer_quant": "fp8"}) is False
# Device without dense support (e.g. non-CUDA) never widens.
monkeypatch.setattr(dmod, "select_transformer_quant_scheme", lambda target, mode: "fp8")
monkeypatch.setattr(
dmod, "select_transformer_quant_scheme", lambda target, mode, family = None: "fp8"
)
monkeypatch.setattr(dmod, "dense_transformer_supported", lambda target: False)
assert backend._dense_quant_prefetch_needed(fam, {"transformer_quant": "fp8"}) is False

View file

@ -433,7 +433,9 @@ def test_fp8_config_uses_per_row_granularity():
def test_quantize_transformer_applies_and_marks(monkeypatch):
monkeypatch.setattr(tq, "select_transformer_quant_scheme", lambda target, mode: TQ_FP8)
monkeypatch.setattr(
tq, "select_transformer_quant_scheme", lambda target, mode, family = None: TQ_FP8
)
seen: dict = {}
def _mk(scheme, fast_accum = None):
@ -458,13 +460,17 @@ def test_quantize_transformer_applies_and_marks(monkeypatch):
def test_quantize_transformer_none_when_unsupported(monkeypatch):
monkeypatch.setattr(tq, "select_transformer_quant_scheme", lambda target, mode: None)
monkeypatch.setattr(
tq, "select_transformer_quant_scheme", lambda target, mode, family = None: None
)
pipe = types.SimpleNamespace(transformer = types.SimpleNamespace())
assert quantize_transformer(pipe, _target(), mode = "auto") is None
def test_quantize_transformer_tolerates_failure(monkeypatch):
monkeypatch.setattr(tq, "select_transformer_quant_scheme", lambda target, mode: TQ_INT8)
monkeypatch.setattr(
tq, "select_transformer_quant_scheme", lambda target, mode, family = None: TQ_INT8
)
monkeypatch.setattr(tq, "_make_quant_config", lambda scheme: "cfg")
tqz = types.ModuleType("torchao.quantization")
@ -480,3 +486,60 @@ def test_quantize_transformer_tolerates_failure(monkeypatch):
pipe = types.SimpleNamespace(transformer = types.SimpleNamespace())
# A quantise failure returns None (caller falls back to GGUF), never raises.
assert quantize_transformer(pipe, _target(), mode = "int8") is None
# ── family scheme deny (measured model-level breakage) ────────────────────────
def test_family_deny_auto_skips_fp8_for_qwen(monkeypatch):
# B200 with every scheme available: auto must NOT pick fp8 / nvfp4 / mxfp8 for the
# Qwen DiT (per-row fp8 renders black frames on it; see _FAMILY_SCHEME_DENY) and
# falls through the ladder to int8, which measures excellent on Qwen.
_stub_torch(monkeypatch, cc = (10, 0))
_allow(monkeypatch, {TQ_FP8, TQ_NVFP4, TQ_MXFP8, TQ_INT8})
assert select_transformer_quant_scheme(_target(), "auto", family = "qwen-image") == TQ_INT8
assert select_transformer_quant_scheme(_target(), "auto", family = "qwen-image-edit") == TQ_INT8
def test_family_deny_refuses_explicit_fp8_for_qwen(monkeypatch):
# An explicit fp8 request on qwen-image returns None (same contract as an
# unsupported scheme: the caller builds the GGUF pipeline instead). int8 stays
# honored on qwen, and fp8 stays honored on families outside the deny table.
_stub_torch(monkeypatch, cc = (10, 0))
_allow(monkeypatch, {TQ_FP8, TQ_INT8})
assert select_transformer_quant_scheme(_target(), "fp8", family = "qwen-image") is None
assert select_transformer_quant_scheme(_target(), "int8", family = "qwen-image") == TQ_INT8
assert select_transformer_quant_scheme(_target(), "fp8", family = "z-image") == TQ_FP8
def test_family_deny_no_family_keeps_ladder(monkeypatch):
# Without a family (or an unknown one) the ladder is unchanged: fp8 first on B200.
_stub_torch(monkeypatch, cc = (10, 0))
_allow(monkeypatch, {TQ_FP8, TQ_INT8})
assert select_transformer_quant_scheme(_target(), "auto") == TQ_FP8
assert select_transformer_quant_scheme(_target(), "auto", family = "sdxl") == TQ_FP8
def test_quantize_transformer_threads_family(monkeypatch):
# quantize_transformer passes the family down to the selector, so a denied
# (family, scheme) pair never reaches torchao.
_stub_torch(monkeypatch, cc = (10, 0))
_allow(monkeypatch, {TQ_FP8, TQ_INT8})
pipe = types.SimpleNamespace(transformer = types.SimpleNamespace())
called = {}
tqz = types.ModuleType("torchao.quantization")
def _quantize(
module,
config,
filter_fn = None,
):
called["scheme"] = True
tqz.quantize_ = _quantize
tqz.Int8DynamicActivationInt8WeightConfig = lambda: "int8-cfg"
tqz.Float8DynamicActivationFloat8WeightConfig = lambda **kw: "fp8-cfg"
tqz.PerRow = lambda: "per-row"
monkeypatch.setitem(sys.modules, "torchao.quantization", tqz)
assert quantize_transformer(pipe, _target(), mode = "fp8", family = "qwen-image") is None
assert called == {}