unsloth/studio/backend/tests/test_diffusion_precision.py
2026-07-08 09:48:12 +00:00

513 lines
22 KiB
Python

# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
"""Unit tests for text-encoder quantisation (``diffusion_precision.py``).
Hermetic: torch + the diffusers / torchao casters are stubbed via ``sys.modules`` so
gating and the apply path run without a GPU, real diffusers, or real torchao.
"""
from __future__ import annotations
import sys
import types
import pytest
import core.inference.diffusion_precision as dp
from core.inference.diffusion_precision import (
TE_QUANT_AUTO,
TE_QUANT_FP8,
TE_QUANT_FP8_DYNAMIC,
TE_QUANT_INT8,
TE_QUANT_NVFP4,
_cast_int8_selective,
_cast_nvfp4,
_keep_bf16_block_fqns,
normalize_te_quant,
quantize_text_encoders,
select_te_quant_scheme,
te_quant_supported,
)
def _target(
*,
device = "cuda",
dtype = "bfloat16",
cc = (10, 0),
):
return types.SimpleNamespace(device = device, dtype = dtype, _cc = cc)
def _stub_torch(
monkeypatch,
*,
with_fp8 = True,
cc = (10, 0),
):
torch = types.ModuleType("torch")
torch.bfloat16 = "bfloat16"
torch.float16 = "float16"
if with_fp8:
torch.float8_e4m3fn = "float8_e4m3fn"
# _cast_fp8 skips nn.Embedding tables (skip_modules_classes) to keep prompt
# tokens full precision, and _keep_bf16_block_fqns walks for nn.ModuleList block
# stacks, so the stub torch must expose both.
torch.nn = types.SimpleNamespace(
Embedding = type("Embedding", (), {}),
ModuleList = type("ModuleList", (list,), {}),
)
torch.cuda = types.SimpleNamespace(get_device_capability = lambda *a: cc)
monkeypatch.setitem(sys.modules, "torch", torch)
return torch
def _stub_casters(monkeypatch, recorder):
# diffusers fp8 layerwise casting
hooks = types.ModuleType("diffusers.hooks")
casting = types.ModuleType("diffusers.hooks.layerwise_casting")
casting.DEFAULT_SKIP_MODULES_PATTERN = ("norm",)
hooks.apply_layerwise_casting = lambda module, **kw: recorder.append(("fp8", module))
monkeypatch.setitem(sys.modules, "diffusers.hooks", hooks)
monkeypatch.setitem(sys.modules, "diffusers.hooks.layerwise_casting", casting)
# torchao nvfp4 -- quantize_ now receives the vision-tower exclusion filter_fn; accept + ignore.
tq = types.ModuleType("torchao.quantization")
tq.quantize_ = lambda module, config, filter_fn = None: recorder.append(("nvfp4", module))
mx = types.ModuleType("torchao.prototype.mx_formats")
mx.NVFP4WeightOnlyConfig = lambda: "nvfp4cfg"
monkeypatch.setitem(sys.modules, "torchao.quantization", tq)
monkeypatch.setitem(sys.modules, "torchao.prototype.mx_formats", mx)
# _cast_nvfp4 / _cast_fp8_dynamic pull the shared linear filter from the transformer-quant module.
dtq = types.ModuleType("core.inference.diffusion_transformer_quant")
dtq.DEFAULT_MIN_LINEAR_FEATURES = 512
dtq.make_filter_fn = lambda min_features, exclude = (), *, require_bf16 = False: (
lambda module, fqn = "": True
)
monkeypatch.setitem(sys.modules, "core.inference.diffusion_transformer_quant", dtq)
# ── normalisation ─────────────────────────────────────────────────────────────
def test_normalize_te_quant():
assert normalize_te_quant(None) is None
assert normalize_te_quant("") is None
assert normalize_te_quant("none") is None
# "off" disables (like the transformer's normalize) -> dense.
assert normalize_te_quant("off") is None
# "auto" passes through for select_te_quant_scheme to resolve.
assert normalize_te_quant("AUTO") == TE_QUANT_AUTO
assert normalize_te_quant("FP8") == TE_QUANT_FP8
assert normalize_te_quant("NVFP4") == TE_QUANT_NVFP4
assert normalize_te_quant("int8") == TE_QUANT_INT8
# Hyphens fold to underscores so "fp8-dynamic" is accepted.
assert normalize_te_quant("FP8-Dynamic") == TE_QUANT_FP8_DYNAMIC
with pytest.raises(ValueError):
normalize_te_quant("int2")
# ── gating ────────────────────────────────────────────────────────────────────
def test_fp8_supported_requires_cuda_bf16_and_fp8(monkeypatch):
_stub_torch(monkeypatch, with_fp8 = True)
assert te_quant_supported(_target(), TE_QUANT_FP8) is True
assert te_quant_supported(_target(device = "cpu"), TE_QUANT_FP8) is False
assert te_quant_supported(_target(dtype = "float16"), TE_QUANT_FP8) is False
def test_nvfp4_supported_requires_blackwell(monkeypatch):
_stub_torch(monkeypatch, cc = (10, 0))
assert te_quant_supported(_target(), TE_QUANT_NVFP4) is True
# Hopper (cc 9.0) has no NVFP4 tensor cores.
_stub_torch(monkeypatch, cc = (9, 0))
assert te_quant_supported(_target(), TE_QUANT_NVFP4) is False
def test_int8_supported_requires_sm80(monkeypatch):
# int8 tensor cores (torch._int_mm) need Ampere sm_80+.
_stub_torch(monkeypatch, cc = (8, 0))
assert te_quant_supported(_target(), TE_QUANT_INT8) is True
_stub_torch(monkeypatch, cc = (7, 5))
assert te_quant_supported(_target(), TE_QUANT_INT8) is False
# Still needs CUDA + bf16 like every mode.
_stub_torch(monkeypatch, cc = (8, 0))
assert te_quant_supported(_target(device = "cpu"), TE_QUANT_INT8) is False
def test_fp8_dynamic_supported_requires_sm89_and_fp8(monkeypatch):
# Compute fp8 (torch._scaled_mm) needs fp8-GEMM silicon: Ada sm_89+ / Hopper / Blackwell.
_stub_torch(monkeypatch, cc = (8, 9))
assert te_quant_supported(_target(), TE_QUANT_FP8_DYNAMIC) is True
_stub_torch(monkeypatch, cc = (9, 0))
assert te_quant_supported(_target(), TE_QUANT_FP8_DYNAMIC) is True
# Ampere (8.0) has int8 but not fp8 GEMM.
_stub_torch(monkeypatch, cc = (8, 0))
assert te_quant_supported(_target(), TE_QUANT_FP8_DYNAMIC) is False
# No fp8 dtype at all -> unsupported regardless of arch.
_stub_torch(monkeypatch, with_fp8 = False, cc = (9, 0))
assert te_quant_supported(_target(), TE_QUANT_FP8_DYNAMIC) is False
# ── apply ─────────────────────────────────────────────────────────────────────
def test_quantize_disabled_returns_none(monkeypatch):
_stub_torch(monkeypatch)
pipe = types.SimpleNamespace(text_encoder = object())
assert quantize_text_encoders(pipe, _target(), mode = None) is None
assert quantize_text_encoders(pipe, _target(), mode = "none") is None
def test_quantize_fp8_casts_all_encoders(monkeypatch):
_stub_torch(monkeypatch)
recorder: list = []
_stub_casters(monkeypatch, recorder)
te1, te3 = object(), object()
pipe = types.SimpleNamespace(text_encoder = te1, text_encoder_2 = None, text_encoder_3 = te3)
mode = quantize_text_encoders(pipe, _target(), mode = "fp8")
assert mode == TE_QUANT_FP8
assert recorder == [("fp8", te1), ("fp8", te3)]
def test_quantize_nvfp4_uses_torchao(monkeypatch):
_stub_torch(monkeypatch, cc = (10, 0))
recorder: list = []
_stub_casters(monkeypatch, recorder)
te = object()
pipe = types.SimpleNamespace(text_encoder = te)
mode = quantize_text_encoders(pipe, _target(), mode = "nvfp4")
assert mode == TE_QUANT_NVFP4
assert recorder == [("nvfp4", te)]
def test_quantize_nvfp4_unsupported_on_hopper_is_noop(monkeypatch):
_stub_torch(monkeypatch, cc = (9, 0))
recorder: list = []
_stub_casters(monkeypatch, recorder)
pipe = types.SimpleNamespace(text_encoder = object())
assert quantize_text_encoders(pipe, _target(cc = (9, 0)), mode = "nvfp4") is None
assert recorder == []
def test_quantize_tolerates_caster_failure(monkeypatch):
_stub_torch(monkeypatch)
hooks = types.ModuleType("diffusers.hooks")
casting = types.ModuleType("diffusers.hooks.layerwise_casting")
casting.DEFAULT_SKIP_MODULES_PATTERN = ("norm",)
def _boom(module, **kwargs):
raise RuntimeError("fp8 unsupported for this layer")
hooks.apply_layerwise_casting = _boom
monkeypatch.setitem(sys.modules, "diffusers.hooks", hooks)
monkeypatch.setitem(sys.modules, "diffusers.hooks.layerwise_casting", casting)
pipe = types.SimpleNamespace(text_encoder = object())
# The only encoder fails to cast -> nothing applied -> None.
assert quantize_text_encoders(pipe, _target(), mode = "fp8") is None
# ── int8 (selective) + fp8_dynamic routing ─────────────────────────────────────
def test_quantize_int8_uses_family_keep_bf16_schedule(monkeypatch):
# int8 for a family with a measured schedule routes to the selective caster with
# that family's (skip_first, skip_last); qwen-image keeps first+last 6 blocks bf16.
_stub_torch(monkeypatch, cc = (10, 0))
calls: list = []
monkeypatch.setattr(
dp, "_cast_int8_selective", lambda enc, tgt, first, last: calls.append((enc, first, last))
)
te = object()
pipe = types.SimpleNamespace(text_encoder = te)
mode = quantize_text_encoders(pipe, _target(), mode = "int8", family = "qwen-image")
assert mode == TE_QUANT_INT8
assert calls == [(te, 6, 6)]
def test_quantize_int8_unknown_family_falls_back_to_fp8(monkeypatch):
# A family without an int8 keep-bf16 schedule falls back to layerwise fp8 (logged),
# never silently running full int8 that would degrade the encoder.
_stub_torch(monkeypatch, cc = (10, 0))
int8_calls: list = []
fp8_calls: list = []
monkeypatch.setattr(dp, "_cast_int8_selective", lambda *a: int8_calls.append(a))
monkeypatch.setattr(dp, "_cast_fp8", lambda enc, tgt: fp8_calls.append(enc))
te = object()
pipe = types.SimpleNamespace(text_encoder = te)
mode = quantize_text_encoders(pipe, _target(), mode = "int8", family = "wan-umt5")
assert mode == TE_QUANT_FP8
assert int8_calls == [] and fp8_calls == [te]
def test_quantize_fp8_dynamic_uses_compute_caster(monkeypatch):
# fp8_dynamic routes to the torchao per-row compute caster (not the layerwise one)
# and needs no per-family schedule.
_stub_torch(monkeypatch, cc = (9, 0))
calls: list = []
monkeypatch.setattr(dp, "_cast_fp8_dynamic", lambda enc, tgt: calls.append(enc))
te = object()
pipe = types.SimpleNamespace(text_encoder = te)
mode = quantize_text_encoders(pipe, _target(), mode = "fp8_dynamic")
assert mode == TE_QUANT_FP8_DYNAMIC
assert calls == [te]
def test_quantize_int8_unsupported_hw_is_noop(monkeypatch):
# int8 on pre-Ampere silicon (no int8 tensor cores) applies nothing.
_stub_torch(monkeypatch, cc = (7, 5))
monkeypatch.setattr(dp, "_cast_int8_selective", lambda *a: pytest.fail("must not cast"))
pipe = types.SimpleNamespace(text_encoder = object())
assert quantize_text_encoders(pipe, _target(), mode = "int8", family = "qwen-image") is None
def test_quantize_te_skips_torchao_modes_under_offload(monkeypatch):
# The torchao modes (int8-with-schedule / fp8_dynamic / nvfp4) produce tensor subclasses that
# reject Module.to(), which an offload hook uses, so they must be skipped under offload (the DiT
# path skips torchao quant for the same reason). Hardware supports every mode here, so a None
# result proves the offload skip, not a capability gate; the casters fail if wrongly invoked.
_stub_torch(monkeypatch, cc = (10, 0))
monkeypatch.setattr(
dp, "_cast_fp8_dynamic", lambda *a: pytest.fail("torchao caster must not run")
)
monkeypatch.setattr(dp, "_cast_nvfp4", lambda *a: pytest.fail("torchao caster must not run"))
monkeypatch.setattr(
dp, "_cast_int8_selective", lambda *a: pytest.fail("torchao caster must not run")
)
pipe = types.SimpleNamespace(text_encoder = object())
assert quantize_text_encoders(pipe, _target(), mode = "fp8_dynamic", offload_active = True) is None
assert quantize_text_encoders(pipe, _target(), mode = "nvfp4", offload_active = True) is None
assert (
quantize_text_encoders(
pipe, _target(), mode = "int8", family = "qwen-image", offload_active = True
)
is None
)
# Layerwise fp8 is not torchao and streams fine under offload, so it still engages.
fp8_calls: list = []
monkeypatch.setattr(dp, "_cast_fp8", lambda enc, tgt: fp8_calls.append(enc))
assert quantize_text_encoders(pipe, _target(), mode = "fp8", offload_active = True) == TE_QUANT_FP8
assert len(fp8_calls) == 1
# ── block selection + real int8 filter closure ─────────────────────────────────
def test_keep_bf16_block_fqns_selects_first_and_last(monkeypatch):
torch = _stub_torch(monkeypatch)
module_list = torch.nn.ModuleList
layers = module_list([object() for _ in range(10)])
# A short stack (<= skip_first + skip_last) contributes nothing (keeping it all would
# leave no interior to quantise).
short = module_list([object() for _ in range(4)])
enc = types.SimpleNamespace()
enc.named_modules = lambda: [("", enc), ("model.layers", layers), ("aux.blocks", short)]
keep = _keep_bf16_block_fqns(enc, 3, 2)
assert keep == {
"model.layers.0",
"model.layers.1",
"model.layers.2",
"model.layers.8",
"model.layers.9",
}
def _stub_transformer_quant(monkeypatch, captured):
# Reuse the committed factory's names but record what the int8 caster hands quantize_().
dtq = types.ModuleType("core.inference.diffusion_transformer_quant")
dtq.TQ_INT8 = "int8"
dtq.TQ_FP8 = "fp8"
dtq.DEFAULT_MIN_LINEAR_FEATURES = 512
dtq._make_quant_config = lambda scheme, *a, **k: f"cfg:{scheme}"
dtq.exclude_tokens_for_scheme = lambda scheme: ("modulation",)
def _make_filter_fn(
min_features,
exclude_name_tokens = (),
*,
require_bf16 = False,
):
def _f(module, fqn = ""):
return not any(tok in fqn for tok in exclude_name_tokens)
return _f
dtq.make_filter_fn = _make_filter_fn
monkeypatch.setitem(sys.modules, "core.inference.diffusion_transformer_quant", dtq)
tq = types.ModuleType("torchao.quantization")
def _quantize_(
module,
config,
filter_fn = None,
):
captured["config"] = config
captured["filter_fn"] = filter_fn
tq.quantize_ = _quantize_
monkeypatch.setitem(sys.modules, "torchao.quantization", tq)
# _cast_nvfp4 builds its config from here.
mx = types.ModuleType("torchao.prototype.mx_formats")
mx.NVFP4WeightOnlyConfig = lambda: "nvfp4cfg"
monkeypatch.setitem(sys.modules, "torchao.prototype.mx_formats", mx)
def test_int8_filter_keeps_blocks_and_towers_dense(monkeypatch):
# The real selective closure: interior Linears quantise, but the kept first blocks,
# the vision tower, lm_head, and the encoder's fp32-kept modules (T5 "wo") stay bf16.
torch = _stub_torch(monkeypatch)
captured: dict = {}
_stub_transformer_quant(monkeypatch, captured)
layers = torch.nn.ModuleList([object() for _ in range(8)])
enc = types.SimpleNamespace(_keep_in_fp32_modules = ["wo"])
enc.named_modules = lambda: [("model.layers", layers)]
_cast_int8_selective(enc, _target(), 3, 0)
assert captured["config"] == "cfg:int8"
ff = captured["filter_fn"]
# Kept first-3 decoder blocks stay bf16.
assert ff(object(), "model.layers.0.self_attn.q_proj") is False
assert ff(object(), "model.layers.2.mlp.gate_proj") is False
# An interior block is quantised.
assert ff(object(), "model.layers.5.self_attn.q_proj") is True
# Vision tower / lm_head / T5 wo are excluded by the shared token filter.
assert ff(object(), "visual.blocks.0.attn.qkv") is False
assert ff(object(), "lm_head") is False
assert ff(object(), "model.decoder.wo") is False
def test_nvfp4_filter_keeps_vision_tower_dense(monkeypatch):
# Weight-only NVFP4 on a text encoder must exclude the VLM vision tower / lm_head / T5 "wo"
# like the int8 / fp8 torchao TE modes -- 4-bit-ing a qwen-image(-edit) Qwen2.5-VL image tower
# degrades the edit/image conditioning the sibling schemes deliberately protect. Before the fix
# _cast_nvfp4 quantised every nn.Linear (no filter_fn), so the tower was silently 4-bit.
_stub_torch(monkeypatch)
captured: dict = {}
_stub_transformer_quant(monkeypatch, captured)
enc = types.SimpleNamespace(_keep_in_fp32_modules = ["wo"])
_cast_nvfp4(enc, _target())
assert captured["config"] == "nvfp4cfg"
ff = captured["filter_fn"]
assert ff is not None # a filter is passed now, not None (which quantised everything)
# Vision tower / lm_head / T5 wo stay bf16; an interior projection still quantises.
assert ff(object(), "visual.blocks.0.attn.qkv") is False
assert ff(object(), "vision_tower.encoder.layers.0.mlp.fc1") is False
assert ff(object(), "lm_head") is False
assert ff(object(), "model.decoder.wo") is False
assert ff(object(), "model.layers.5.self_attn.q_proj") is True
# ── auto ladder (select_te_quant_scheme) ────────────────────────────────────────
def _stub_tq_select(
monkeypatch,
*,
cc,
consumer = False,
smoke = True,
):
"""Stub the transformer module's shared helpers that select_te_quant_scheme imports:
capability, GPU class, and the kernel smoke probe (bool or a (tq, dev) predicate)."""
dtq = types.ModuleType("core.inference.diffusion_transformer_quant")
dtq._capability = lambda: cc
dtq._is_consumer_gpu = lambda device = None: consumer
dtq._smoke_probe = smoke if callable(smoke) else (lambda tq, dev: smoke)
monkeypatch.setitem(sys.modules, "core.inference.diffusion_transformer_quant", dtq)
return dtq
def _allow_te(monkeypatch, allowed):
"""Force te_quant_supported to accept only ``allowed`` (simulates the hardware gate)."""
monkeypatch.setattr(dp, "te_quant_supported", lambda target, mode: mode in allowed)
def test_select_te_auto_datacenter_prefers_fp8_dynamic(monkeypatch):
# Data-center fp8-GEMM silicon: fp8_dynamic (compute fp8) leads the ladder.
_stub_tq_select(monkeypatch, cc = (10, 0), consumer = False)
_allow_te(monkeypatch, {TE_QUANT_FP8_DYNAMIC, TE_QUANT_INT8, TE_QUANT_FP8})
assert select_te_quant_scheme(_target(), "auto", family = "qwen-image") == TE_QUANT_FP8_DYNAMIC
def test_select_te_auto_falls_through_to_int8_then_fp8(monkeypatch):
_stub_tq_select(monkeypatch, cc = (10, 0))
# fp8_dynamic unavailable -> int8 (family has a keep-bf16 schedule).
_allow_te(monkeypatch, {TE_QUANT_INT8, TE_QUANT_FP8})
assert select_te_quant_scheme(_target(), "auto", family = "qwen-image") == TE_QUANT_INT8
# A family with NO int8 schedule skips int8 -> layerwise fp8.
assert select_te_quant_scheme(_target(), "auto", family = "z-image") == TE_QUANT_FP8
def test_select_te_auto_consumer_prefers_int8(monkeypatch):
# Consumer GDDR halves fp8 FP32-accumulate but runs int8 full-rate -> int8 first.
_stub_tq_select(monkeypatch, cc = (10, 0), consumer = True)
_allow_te(monkeypatch, {TE_QUANT_FP8_DYNAMIC, TE_QUANT_INT8, TE_QUANT_FP8})
assert select_te_quant_scheme(_target(), "auto", family = "qwen-image") == TE_QUANT_INT8
def test_select_te_auto_offload_uses_layerwise_fp8(monkeypatch):
# Under offload the torchao modes (reject Module.to()) are skipped -> layerwise fp8.
_stub_tq_select(monkeypatch, cc = (10, 0))
_allow_te(monkeypatch, {TE_QUANT_FP8_DYNAMIC, TE_QUANT_INT8, TE_QUANT_FP8})
assert (
select_te_quant_scheme(_target(), "auto", family = "qwen-image", offload_active = True)
== TE_QUANT_FP8
)
def test_select_te_auto_ampere_uses_int8(monkeypatch):
# Ampere sm_80 has no fp8 GEMM; the tier is (int8, fp8).
_stub_tq_select(monkeypatch, cc = (8, 0))
_allow_te(monkeypatch, {TE_QUANT_INT8, TE_QUANT_FP8})
assert select_te_quant_scheme(_target(), "auto", family = "qwen-image") == TE_QUANT_INT8
def test_select_te_auto_family_deny_skips_scheme(monkeypatch):
_stub_tq_select(monkeypatch, cc = (10, 0))
_allow_te(monkeypatch, {TE_QUANT_FP8_DYNAMIC, TE_QUANT_INT8, TE_QUANT_FP8})
monkeypatch.setattr(
dp, "_TE_FAMILY_SCHEME_DENY", {"qwen-image": frozenset({TE_QUANT_FP8_DYNAMIC})}
)
# fp8_dynamic denied for this family -> falls to int8.
assert select_te_quant_scheme(_target(), "auto", family = "qwen-image") == TE_QUANT_INT8
def test_select_te_auto_smoke_failure_skips_scheme(monkeypatch):
# fp8_dynamic is hardware-supported but its kernel smoke-probe fails -> skip to int8.
_stub_tq_select(monkeypatch, cc = (10, 0), smoke = lambda tq, dev: tq != "fp8")
_allow_te(monkeypatch, {TE_QUANT_FP8_DYNAMIC, TE_QUANT_INT8, TE_QUANT_FP8})
assert select_te_quant_scheme(_target(), "auto", family = "qwen-image") == TE_QUANT_INT8
def test_select_te_auto_pre_ampere_and_no_cuda_are_none(monkeypatch):
_stub_tq_select(monkeypatch, cc = (7, 5))
_allow_te(monkeypatch, {TE_QUANT_FP8})
assert select_te_quant_scheme(_target(), "auto", family = "qwen-image") is None
_stub_tq_select(monkeypatch, cc = None)
assert select_te_quant_scheme(_target(), "auto", family = "qwen-image") is None
def test_select_te_explicit_scheme_passes_through(monkeypatch):
# An explicit request is returned as-is (quantize_text_encoders re-gates it); no ladder walk,
# so no transformer-module stub is needed.
assert select_te_quant_scheme(_target(), "fp8") == TE_QUANT_FP8
assert select_te_quant_scheme(_target(), "int8") == TE_QUANT_INT8
assert select_te_quant_scheme(_target(), None) is None
assert select_te_quant_scheme(_target(), "none") is None
def test_quantize_text_encoders_auto_resolves_and_applies(monkeypatch):
# End-to-end: mode="auto" resolves via the ladder then applies the resolved caster.
_stub_tq_select(monkeypatch, cc = (10, 0))
_allow_te(monkeypatch, {TE_QUANT_FP8_DYNAMIC, TE_QUANT_INT8, TE_QUANT_FP8})
calls: list = []
monkeypatch.setattr(dp, "_cast_fp8_dynamic", lambda enc, tgt: calls.append(enc))
te = object()
pipe = types.SimpleNamespace(text_encoder = te)
mode = quantize_text_encoders(pipe, _target(), mode = "auto", family = "qwen-image")
assert mode == TE_QUANT_FP8_DYNAMIC
assert calls == [te]