From 7bf80f6a4ec7573e25772a25fc8f4bc79695ad47 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Sat, 4 Jul 2026 08:51:10 +0000 Subject: [PATCH 1/2] Deny fp8/mxfp8/nvfp4 dense quant for the Qwen DiT (black frames, measured) A 28-pair accuracy gate on a B200 (same-seed vs the dense bf16 reference) found per-row fp8 dynamic quant renders EVERY qwen-image frame black (mean luma 0.0000, SSIM 0.016), reproduced identically with on-the-fly quantize_ on the dense transformer, so it is the model's activation range, not a checkpoint artifact. mxfp8 shows real semantic damage at 1024px (CLIP delta mean 0.0146, worst cases 0.064/0.102) and nvfp4 measures LPIPS mean 0.51. int8 dynamic (per-token scales) is excellent on Qwen: LPIPS mean 0.069, SSIM 0.958. The per-scheme smoke probe only proves the GEMM kernel runs, so it cannot catch model-level breakage. Add _FAMILY_SCHEME_DENY consulted by select_transformer_quant_scheme: auto skips denied schemes (Qwen lands on int8) and an explicit denied request returns None, the same GGUF-fallback contract as an unsupported scheme. Family is threaded from the three diffusion.py call sites; existing behavior is unchanged for every other family. 4 new tests; 529 diffusion tests green; CI-sim green. --- studio/backend/core/inference/diffusion.py | 13 +++- .../inference/diffusion_transformer_quant.py | 40 +++++++++++- .../backend/tests/test_diffusion_backend.py | 12 ++-- .../tests/test_diffusion_transformer_quant.py | 62 ++++++++++++++++++- 4 files changed, 112 insertions(+), 15 deletions(-) diff --git a/studio/backend/core/inference/diffusion.py b/studio/backend/core/inference/diffusion.py index 28eafe6829..2d1a547166 100644 --- a/studio/backend/core/inference/diffusion.py +++ b/studio/backend/core/inference/diffusion.py @@ -418,7 +418,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( @@ -1292,7 +1294,9 @@ 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 @@ -1331,7 +1335,10 @@ 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 39aa751778..363098bcdb 100644 --- a/studio/backend/core/inference/diffusion_transformer_quant.py +++ b/studio/backend/core/inference/diffusion_transformer_quant.py @@ -101,6 +101,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] = {} @@ -201,18 +225,25 @@ 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: @@ -220,6 +251,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 @@ -385,6 +418,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, @@ -396,7 +430,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 d36e62f3cf..8a337557d4 100644 --- a/studio/backend/tests/test_diffusion_backend.py +++ b/studio/backend/tests/test_diffusion_backend.py @@ -1764,7 +1764,7 @@ 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): @@ -1828,7 +1828,7 @@ 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} @@ -1958,7 +1958,7 @@ 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 @@ -2005,7 +2005,7 @@ 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 @@ -2016,10 +2016,10 @@ 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..8fa4c71e7d 100644 --- a/studio/backend/tests/test_diffusion_transformer_quant.py +++ b/studio/backend/tests/test_diffusion_transformer_quant.py @@ -433,7 +433,7 @@ 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 +458,13 @@ 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 +480,59 @@ 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 == {} From ab56d819356021d809553c372abba2348476b54d Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Sat, 4 Jul 2026 08:51:43 +0000 Subject: [PATCH 2/2] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- studio/backend/core/inference/diffusion.py | 12 ++++---- .../inference/diffusion_transformer_quant.py | 4 ++- .../backend/tests/test_diffusion_backend.py | 24 +++++++++++---- .../tests/test_diffusion_transformer_quant.py | 29 ++++++++++++------- 4 files changed, 46 insertions(+), 23 deletions(-) diff --git a/studio/backend/core/inference/diffusion.py b/studio/backend/core/inference/diffusion.py index 2d1a547166..fda2a27a4b 100644 --- a/studio/backend/core/inference/diffusion.py +++ b/studio/backend/core/inference/diffusion.py @@ -1294,9 +1294,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, family = getattr(fam, "name", None) - ) + 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 @@ -1336,8 +1334,12 @@ class DiffusionBackend: ) pipe = self._assemble_pipe(pipeline_cls, base, transformer, dtype, hf_token, device) scheme = quantize_transformer( - pipe, target, mode = mode, family = getattr(fam, "name", None), - fast_accum = fast_accum, logger = logger, + 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") diff --git a/studio/backend/core/inference/diffusion_transformer_quant.py b/studio/backend/core/inference/diffusion_transformer_quant.py index 363098bcdb..4d31ae6f25 100644 --- a/studio/backend/core/inference/diffusion_transformer_quant.py +++ b/studio/backend/core/inference/diffusion_transformer_quant.py @@ -226,7 +226,9 @@ def dense_transformer_supported(target: Any) -> bool: def select_transformer_quant_scheme( - target: Any, requested: Optional[str], family: Optional[str] = None + target: Any, + requested: Optional[str], + family: Optional[str] = None, ) -> Optional[str]: """The concrete scheme to apply, or None to fall back to GGUF. diff --git a/studio/backend/tests/test_diffusion_backend.py b/studio/backend/tests/test_diffusion_backend.py index 8a337557d4..cd25103a24 100644 --- a/studio/backend/tests/test_diffusion_backend.py +++ b/studio/backend/tests/test_diffusion_backend.py @@ -1764,7 +1764,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, family = None: 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): @@ -1828,7 +1830,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, family = None: "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} @@ -1958,7 +1962,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, family = None: 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 @@ -2005,7 +2011,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, family = None: "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 @@ -2016,10 +2024,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, family = None: 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, family = None: "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 8fa4c71e7d..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, family = None: 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, family = None: 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, family = None: 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") @@ -492,10 +498,7 @@ def test_family_deny_auto_skips_fp8_for_qwen(monkeypatch): _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 - ) + assert select_transformer_quant_scheme(_target(), "auto", family = "qwen-image-edit") == TQ_INT8 def test_family_deny_refuses_explicit_fp8_for_qwen(monkeypatch): @@ -525,14 +528,18 @@ def test_quantize_transformer_threads_family(monkeypatch): pipe = types.SimpleNamespace(transformer = types.SimpleNamespace()) called = {} tqz = types.ModuleType("torchao.quantization") - def _quantize(module, config, filter_fn = None): + + 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 quantize_transformer(pipe, _target(), mode = "fp8", family = "qwen-image") is None assert called == {}