diff --git a/studio/backend/core/inference/diffusion.py b/studio/backend/core/inference/diffusion.py index 72398e47ce..ca26bfea08 100644 --- a/studio/backend/core/inference/diffusion.py +++ b/studio/backend/core/inference/diffusion.py @@ -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 diff --git a/studio/backend/core/inference/diffusion_transformer_quant.py b/studio/backend/core/inference/diffusion_transformer_quant.py index 035234b034..b990efc563 100644 --- a/studio/backend/core/inference/diffusion_transformer_quant.py +++ b/studio/backend/core/inference/diffusion_transformer_quant.py @@ -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) diff --git a/studio/backend/tests/test_diffusion_backend.py b/studio/backend/tests/test_diffusion_backend.py index d2f0ea63fd..953980718f 100644 --- a/studio/backend/tests/test_diffusion_backend.py +++ b/studio/backend/tests/test_diffusion_backend.py @@ -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 diff --git a/studio/backend/tests/test_diffusion_transformer_quant.py b/studio/backend/tests/test_diffusion_transformer_quant.py index 60296bdfc5..a621bd3874 100644 --- a/studio/backend/tests/test_diffusion_transformer_quant.py +++ b/studio/backend/tests/test_diffusion_transformer_quant.py @@ -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 == {}