Studio diffusion (Phase 8): opt-in fast transformer (torchao int8/fp8/fp4 on a dense source)

Add an opt-in transformer_quant mode that loads the dense bf16 transformer and
torchao-quantises it onto the low-precision tensor cores, instead of the GGUF
transformer (which dequantises to bf16 per matmul and so runs at bf16 rate). On a
B200 (Z-Image-Turbo, 1024px/8 steps): auto picks fp8 at 0.614s vs GGUF+compile's
0.823s (1.34x), int8 0.626s (1.32x), both at lower LPIPS than GGUF's own 4-bit floor.

GGUF+compile stays the low-memory default and the fallback. The mode is gated on
CUDA + bf16 + resident VRAM headroom (the dense load peaks ~21GB vs GGUF's 13GB);
any unsupported arch/scheme, OOM, or quant failure falls back to GGUF with a logged
reason. auto picks the best scheme per GPU via a real quantise+matmul smoke probe
(Blackwell nvfp4/fp8/mxfp8, Ada/Hopper fp8, Ampere int8); a min-features filter skips
the tiny projections that crash int8's torch._int_mm. New module mirrors
diffusion_precision.py; quant runs before compile before placement.

184 -> tests pass; new test_diffusion_transformer_quant.py plus backend/route
coverage. scripts/diffusion_bench.py gains --transformer-quant; scripts/quant_probe.py
is the standalone torchao lever probe.
This commit is contained in:
Daniel Han 2026-06-26 06:15:18 +00:00
commit 983f2c14a0
9 changed files with 1040 additions and 28 deletions

View file

@ -216,6 +216,7 @@ def _run(args: argparse.Namespace) -> dict[str, Any]:
memory_mode = args.memory_mode,
speed_mode = args.speed_mode,
text_encoder_quant = args.text_encoder_quant,
transformer_quant = args.transformer_quant,
)
_wait_for_load(backend)
_cuda_sync()
@ -297,6 +298,7 @@ def _run(args: argparse.Namespace) -> dict[str, Any]:
"speed_mode": args.speed_mode,
"cpu_offload": args.cpu_offload,
"text_encoder_quant": args.text_encoder_quant,
"transformer_quant": args.transformer_quant,
},
}
@ -456,6 +458,14 @@ def _build_parser() -> argparse.ArgumentParser:
choices = ["fp8", "nvfp4"],
help = "quantise the companion text encoder (fp8 or nvfp4)",
)
p.add_argument(
"--transformer-quant",
default = None,
choices = ["auto", "int8", "fp8", "nvfp4", "mxfp8"],
help = "opt-in fast transformer: load the DENSE bf16 transformer and torchao-"
"quantise it onto the low-precision tensor cores (faster than GGUF, higher "
"VRAM). auto picks per GPU; falls back to GGUF if unsupported / no VRAM",
)
p.add_argument(
"--cpu-offload", action = "store_true", help = "legacy: force whole-module CPU offload"
)

273
scripts/quant_probe.py Normal file
View file

@ -0,0 +1,273 @@
# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
"""Empirical quant probe: torchao int8/fp8/fp4 dynamic quant vs GGUF+compile.
Question this answers: GGUF stores the Z-Image DiT at 4-bit but dequantizes to bf16
per matmul, so it runs at bf16 tensor-core rate. Can a low-precision *tensor-core*
path (int8dq on any Ampere+, fp8dq on Ada+, NVFP4/MXFP8 on Blackwell), loaded from
the dense bf16 transformer, beat GGUF+compile on speed while staying inside the
quality bar -- and how does its quality compare to GGUF's own 4-bit loss?
Reference for all quality numbers is the DENSE bf16 EAGER image (the best this model
can do). Each config is a fresh pipeline (no compile/quant cross-contamination).
Reports median latency, PSNR + LPIPS vs reference, and peak VRAM. Run on one CUDA GPU.
"""
from __future__ import annotations
import argparse
import sys
import time
from pathlib import Path
import numpy as np
REPO = "unsloth/Z-Image-Turbo-GGUF"
GGUF = "z-image-turbo-Q4_K_M.gguf"
BASE = "Tongyi-MAI/Z-Image-Turbo"
PROMPT = "A cinematic photograph of a red fox in a snowy forest at dawn, highly detailed"
OUT = Path("/mnt/disks/unslothai/ubuntu/workspace_81/outputs/quant_research/probe_images")
def _psnr(a, b):
mse = float(np.mean((a.astype(np.float64) - b.astype(np.float64)) ** 2))
return float("inf") if mse == 0 else float(10 * np.log10(255.0**2 / mse))
_LPIPS = {"fn": None}
def _lpips(ref_arr, arr):
"""Perceptual LPIPS (alexnet) vs reference; lower is closer. None if unavailable."""
try:
import torch
import lpips
if _LPIPS["fn"] is None:
_LPIPS["fn"] = lpips.LPIPS(net="alex", verbose=False).cuda().eval()
def t(x):
t = torch.from_numpy(x).float().permute(2, 0, 1).unsqueeze(0) / 127.5 - 1.0
return t.cuda()
with torch.no_grad():
return float(_LPIPS["fn"](t(ref_arr), t(arr)).item())
except Exception as exc: # noqa: BLE001
print(f" (lpips unavailable: {type(exc).__name__}: {str(exc)[:80]})", flush=True)
return None
def _load_dense():
import torch
import diffusers
t = diffusers.ZImageTransformer2DModel.from_pretrained(
BASE, subfolder="transformer", torch_dtype=torch.bfloat16
)
pipe = diffusers.ZImagePipeline.from_pretrained(BASE, torch_dtype=torch.bfloat16, transformer=t)
pipe.to("cuda")
return pipe
def _load_gguf():
import torch
import diffusers
from huggingface_hub import hf_hub_download
t = diffusers.ZImageTransformer2DModel.from_single_file(
hf_hub_download(REPO, GGUF),
quantization_config=diffusers.GGUFQuantizationConfig(compute_dtype=torch.bfloat16),
torch_dtype=torch.bfloat16, config=BASE, subfolder="transformer",
)
pipe = diffusers.ZImagePipeline.from_pretrained(BASE, torch_dtype=torch.bfloat16, transformer=t)
pipe.to("cuda")
return pipe
def _quant_config(name):
"""Return a torchao config instance for `name`, or raise to mark FAILED."""
from torchao.quantization import (
Int8WeightOnlyConfig,
Int8DynamicActivationInt8WeightConfig,
Float8DynamicActivationFloat8WeightConfig,
)
if name == "int8wo":
return Int8WeightOnlyConfig()
if name == "int8dq":
return Int8DynamicActivationInt8WeightConfig()
if name == "fp8dq":
return Float8DynamicActivationFloat8WeightConfig()
if name == "nvfp4":
from torchao.prototype.mx_formats import NVFP4DynamicActivationNVFP4WeightConfig
return NVFP4DynamicActivationNVFP4WeightConfig()
if name == "mxfp8":
from torchao.prototype.mx_formats import MXDynamicActivationMXWeightConfig
try:
import torch
return MXDynamicActivationMXWeightConfig(
activation_dtype=torch.float8_e4m3fn, weight_dtype=torch.float8_e4m3fn
)
except TypeError:
return MXDynamicActivationMXWeightConfig()
raise ValueError(name)
def _make_filter_fn(min_features):
"""Keep only the FLOP-heavy linears: nn.Linear with both in/out >= min_features.
The int8 dynamic path uses torch._int_mm (needs activation M>16), and the tiny
timestep/pooled projections (in_features=256) run at M=1 and crash it -- skip them."""
import torch.nn as nn
def filter_fn(module, fqn=""):
return (
isinstance(module, nn.Linear)
and getattr(module, "in_features", 0) >= min_features
and getattr(module, "out_features", 0) >= min_features
)
return filter_fn
def _apply_quant(pipe, name, log, min_features):
import torch.nn as nn
from torchao.quantization import quantize_
cfg = _quant_config(name)
total = sum(1 for m in pipe.transformer.modules() if isinstance(m, nn.Linear))
filt = _make_filter_fn(min_features)
q = sum(1 for n, m in pipe.transformer.named_modules() if filt(m, n))
quantize_(pipe.transformer, cfg, filter_fn=filt)
log(f" quantized transformer with {name} ({q}/{total} linears >= {min_features} feat)")
def _compile(pipe, log):
fn = getattr(pipe.transformer, "compile_repeated_blocks", None)
if not callable(fn):
return False
for kw in ({"fullgraph": True, "dynamic": True}, {"dynamic": True}, {}):
try:
fn(**kw)
log(f" compiled repeated blocks {kw}")
return True
except Exception as exc: # noqa: BLE001
log(f" compile {kw} failed: {type(exc).__name__}: {str(exc)[:90]}")
return False
def _gen(pipe, steps, seed, res):
import torch
g = torch.Generator(device="cuda").manual_seed(seed)
torch.cuda.synchronize()
t0 = time.time()
img = pipe(prompt=PROMPT, width=res, height=res, num_inference_steps=steps,
guidance_scale=0.0, generator=g).images[0]
torch.cuda.synchronize()
return img, time.time() - t0
def _median(xs):
return sorted(xs)[len(xs) // 2]
def main(argv=None) -> int:
p = argparse.ArgumentParser()
p.add_argument("--steps", type=int, default=8)
p.add_argument("--res", type=int, default=1024)
p.add_argument("--seed", type=int, default=42)
p.add_argument("--iters", type=int, default=3)
p.add_argument("--min-feat", type=int, default=512,
help="only quantize Linear with in&out features >= this (int8 _int_mm needs M>16)")
p.add_argument("--configs", default="bf16,bf16_c,gguf_c,int8dq_c,fp8dq_c,nvfp4_c,mxfp8_c,int8wo_c",
help="comma list; suffix _c = +compile")
args = p.parse_args(argv)
steps, res, seed, iters = args.steps, args.res, args.seed, args.iters
import torch
OUT.mkdir(parents=True, exist_ok=True)
def run(tag, *, source, quant=None, compile=False):
torch.compiler.reset()
torch.cuda.empty_cache()
torch.cuda.reset_peak_memory_stats()
pipe = _load_dense() if source == "dense" else _load_gguf()
load_peak = torch.cuda.max_memory_allocated() / 1e9
if quant is not None:
_apply_quant(pipe, quant, print_, args.min_feat)
if compile:
_compile(pipe, print_)
_gen(pipe, steps, seed, res) # warmup / compilation
else:
_gen(pipe, steps, seed, res) # allocator warmup
torch.cuda.reset_peak_memory_stats()
dts, img = [], None
for _ in range(iters):
img, dt = _gen(pipe, steps, seed, res)
dts.append(dt)
gen_peak = torch.cuda.max_memory_allocated() / 1e9
arr = np.array(img)
img.save(OUT / f"{tag}.png")
del pipe
torch.cuda.empty_cache()
return tag, _median(dts), arr, load_peak, gen_peak
print_ = lambda s: print(s, flush=True) # noqa: E731
# config table: tag -> (source, quant, compile)
table = {
"bf16": ("dense", None, False),
"bf16_c": ("dense", None, True),
"gguf_c": ("gguf", None, True),
"int8wo_c": ("dense", "int8wo", True),
"int8dq_c": ("dense", "int8dq", True),
"fp8dq_c": ("dense", "fp8dq", True),
"nvfp4_c": ("dense", "nvfp4", True),
"mxfp8_c": ("dense", "mxfp8", True),
}
want = [c.strip() for c in args.configs.split(",") if c.strip()]
print(f"== quant probe (Z-Image-Turbo, {res}px, {steps} steps, seed {seed}) ==", flush=True)
ref_arr = None
rows = []
for tag in want:
if tag not in table:
print(f" {tag}: unknown config, skipping", flush=True)
continue
source, quant, compile = table[tag]
print(f"-- {tag} (source={source} quant={quant} compile={compile}) --", flush=True)
try:
_, med, arr, lp, gp = run(tag, source=source, quant=quant, compile=compile)
except Exception as exc: # noqa: BLE001
import traceback
print(f" {tag:10s} FAILED: {type(exc).__name__}: {str(exc)[:160]}", flush=True)
traceback.print_exc()
rows.append((tag, None, None, None, None, None))
continue
if ref_arr is None and tag == "bf16":
ref_arr = arr
psnr = _psnr(ref_arr, arr) if ref_arr is not None else None
lpips_v = _lpips(ref_arr, arr) if (ref_arr is not None and tag != "bf16") else (0.0 if tag == "bf16" else None)
rows.append((tag, med, psnr, lpips_v, lp, gp))
ps = f"{psnr:.1f}dB" if psnr is not None else "n/a"
lps = f"{lpips_v:.3f}" if lpips_v is not None else "n/a"
print(f" {tag:10s} {med:.3f}s PSNR={ps:>7s} LPIPS={lps:>6s} loadVRAM={lp:.1f}G genVRAM={gp:.1f}G", flush=True)
base = next((r[1] for r in rows if r[0] == "bf16" and r[1]), None)
gguf = next((r[1] for r in rows if r[0] == "gguf_c" and r[1]), None)
print("\n==== SUMMARY (ref = bf16 dense eager) ====", flush=True)
print(f"{'config':10s} {'sec':>7s} {'vs_bf16':>8s} {'vs_gguf':>8s} {'PSNR':>8s} {'LPIPS':>7s} {'loadG':>6s} {'genG':>6s}", flush=True)
for tag, med, psnr, lpips_v, lp, gp in rows:
if med is None:
print(f"{tag:10s} {'FAILED':>7s}", flush=True)
continue
vb = f"{base/med:.2f}x" if base else "-"
vg = f"{gguf/med:.2f}x" if gguf else "-"
ps = f"{psnr:.1f}" if psnr is not None and psnr != float('inf') else ("inf" if psnr == float('inf') else "n/a")
lps = f"{lpips_v:.3f}" if lpips_v is not None else "n/a"
print(f"{tag:10s} {med:>7.3f} {vb:>8s} {vg:>8s} {ps:>8s} {lps:>7s} {lp:>6.1f} {gp:>6.1f}", flush=True)
print("QUANT-PROBE-DONE", flush=True)
return 0
if __name__ == "__main__":
sys.path.insert(0, str(Path(__file__).resolve().parent.parent / "studio" / "backend"))
sys.exit(main())

View file

@ -52,6 +52,11 @@ from .diffusion_speed import (
snapshot_backend_flags,
)
from .diffusion_precision import quantize_text_encoders
from .diffusion_transformer_quant import (
dense_transformer_supported,
normalize_transformer_quant,
quantize_transformer,
)
logger = get_logger(__name__)
@ -81,6 +86,9 @@ class _LoadState:
backend_flags_before: Optional[dict] = None
# Text-encoder quantisation actually engaged: "fp8" | "nvfp4" | None (Phase 2B/2C).
text_encoder_quant: Optional[str] = None
# Transformer quant actually engaged on the opt-in dense fast path: "int8" | "fp8"
# | "nvfp4" | "mxfp8" | None. None means the default GGUF transformer was loaded.
transformer_quant: Optional[str] = None
@dataclass
@ -267,6 +275,7 @@ class DiffusionBackend:
memory_mode: Optional[str] = None,
speed_mode: Optional[str] = None,
text_encoder_quant: Optional[str] = None,
transformer_quant: Optional[str] = None,
) -> dict[str, Any]:
"""Validate, then run the (slow) load on a daemon thread. Returns at once."""
fam = self.validate_load_request(
@ -298,6 +307,7 @@ class DiffusionBackend:
memory_mode = memory_mode,
speed_mode = speed_mode,
text_encoder_quant = text_encoder_quant,
transformer_quant = transformer_quant,
_load_token = token,
),
daemon = True,
@ -419,6 +429,7 @@ class DiffusionBackend:
memory_mode: Optional[str] = None,
speed_mode: Optional[str] = None,
text_encoder_quant: Optional[str] = None,
transformer_quant: Optional[str] = None,
_load_token: Optional[int] = None,
) -> dict[str, Any]:
# Validate first (cheap, no torch/diffusers) so a direct call with a bad
@ -451,26 +462,62 @@ class DiffusionBackend:
# checkpoints never sit in VRAM at once.
self._unload_locked()
# Dequantise the GGUF transformer on-device; the VAE / text-encoder /
# scheduler come from the base diffusers repo (GGUF is transformer-only).
gguf_path = self._resolve_gguf_path(repo_id, gguf_filename, hf_token)
transformer_cls = getattr(diffusers, fam.transformer_class)
transformer = transformer_cls.from_single_file(
gguf_path,
quantization_config = diffusers.GGUFQuantizationConfig(compute_dtype = dtype),
torch_dtype = dtype,
config = base,
subfolder = "transformer",
# Forward the token: the config is fetched from the (possibly gated)
# base repo before from_pretrained gets a chance to authenticate.
token = hf_token,
pipeline_cls = getattr(diffusers, fam.pipeline_class)
# Decide placement up front (the weights are still on CPU, so free VRAM is
# the real budget) -- this also doubles as the dense-quant preflight: the
# dense bf16 transformer must fit resident, so the fast path is offered only
# when the plan is `none`.
plan = self._plan_memory(
target, gguf_path, gguf_filename, base, fam, memory_mode, cpu_offload
)
pipe_kwargs: dict[str, Any] = {"torch_dtype": dtype, "transformer": transformer}
if hf_token:
pipe_kwargs["token"] = hf_token
pipeline_cls = getattr(diffusers, fam.pipeline_class)
pipe = pipeline_cls.from_pretrained(base, **pipe_kwargs)
# Opt-in fast path: load the DENSE bf16 transformer and torchao-quantise it
# (int8 / fp8 / fp4 tensor cores), which beats GGUF's bf16-rate per-matmul
# dequant on both speed and quality, at the cost of a higher-memory dense
# load. Gated on CUDA + bf16 + a resident fit; ANY failure (unsupported arch
# / scheme, OOM, partial quant) falls back to the GGUF build below.
pipe = None
transformer_quant_engaged = None
if (
normalize_transformer_quant(transformer_quant) is not None
and dense_transformer_supported(target)
and plan.offload_policy == OFFLOAD_NONE
):
try:
pipe, transformer_quant_engaged = self._load_dense_quant_pipeline(
transformer_cls, pipeline_cls, base, device, dtype, hf_token,
target, transformer_quant,
)
except Exception as exc: # noqa: BLE001 — fall back to the GGUF build
logger.warning(
"diffusion.transformer_quant_fallback: %s (loading GGUF)", exc
)
pipe = None
transformer_quant_engaged = None
clear_gpu_cache()
if pipe is None:
# Default: dequantise the single-file GGUF transformer on-device; the
# VAE / text-encoder / scheduler come from the base diffusers repo
# (GGUF is transformer-only).
transformer = transformer_cls.from_single_file(
gguf_path,
quantization_config = diffusers.GGUFQuantizationConfig(compute_dtype = dtype),
torch_dtype = dtype,
config = base,
subfolder = "transformer",
# Forward the token: the config is fetched from the (possibly gated)
# base repo before from_pretrained gets a chance to authenticate.
token = hf_token,
)
pipe_kwargs: dict[str, Any] = {"torch_dtype": dtype, "transformer": transformer}
if hf_token:
pipe_kwargs["token"] = hf_token
pipe = pipeline_cls.from_pretrained(base, **pipe_kwargs)
# Resolve the effective speed mode: GGUF models default to the
# near-lossless `default` profile (compile is ~2.2x and sits below
@ -499,18 +546,12 @@ class DiffusionBackend:
logger = logger,
)
# Decide placement from MEASURED free device memory vs the model's
# estimated resident size (transformer GGUF dequantised + the
# companion text-encoder / VAE already cached for `base`), then
# apply it. Computed here, after the build but before placement,
# because the weights are still on CPU so free VRAM is the real
# budget. `cpu_offload=True` stays an explicit override.
plan = self._plan_memory(
target, gguf_path, gguf_filename, base, fam, memory_mode, cpu_offload
)
# apply_memory_plan returns the (policy, tiling) ACTUALLY engaged (it
# may fall back to whole-module offload, and tiling is a no-op on a
# pipeline with no tiling control), so status stays honest.
# Apply the placement planned above (from MEASURED free device memory vs
# the model's estimated resident size). apply_memory_plan returns the
# (policy, tiling) ACTUALLY engaged (it may fall back to whole-module
# offload, and tiling is a no-op on a pipeline with no tiling control), so
# status stays honest. The dense fast path already placed the pipe resident;
# for the `none` policy this is an idempotent re-placement.
effective_policy, effective_tiling = apply_memory_plan(
pipe, plan, device = device, logger = logger
)
@ -530,6 +571,7 @@ class DiffusionBackend:
speed_optims = tuple(k for k, v in speed_applied.items() if v),
backend_flags_before = backend_flags_before,
text_encoder_quant = te_quant,
transformer_quant = transformer_quant_engaged,
)
logger.info(
@ -543,6 +585,38 @@ class DiffusionBackend:
)
return self.status()
def _load_dense_quant_pipeline(
self,
transformer_cls: Any,
pipeline_cls: Any,
base: str,
device: str,
dtype: Any,
hf_token: Optional[str],
target: DiffusionDeviceTarget,
mode: Optional[str],
) -> tuple[Any, str]:
"""Build the opt-in fast pipeline: load the DENSE bf16 transformer from the base
repo (``subfolder="transformer"``), assemble the pipeline, place it on the device,
and torchao-quantise the transformer in place. Returns ``(pipe, engaged_scheme)``.
Raises if the scheme is unsupported or quantisation fails, so ``load_pipeline``
catches it and falls back to the GGUF build. Quantisation runs ON the device (the
dynamic int8 / fp8 / fp4 kernels need the weights on CUDA) and BEFORE the loader
compiles the repeated block, so the order is quantize -> compile -> placement."""
transformer = transformer_cls.from_pretrained(
base, subfolder = "transformer", torch_dtype = dtype, token = hf_token
)
pipe_kwargs: dict[str, Any] = {"torch_dtype": dtype, "transformer": transformer}
if hf_token:
pipe_kwargs["token"] = hf_token
pipe = pipeline_cls.from_pretrained(base, **pipe_kwargs)
pipe.to(device)
scheme = quantize_transformer(pipe, target, mode = mode, logger = logger)
if scheme is None:
raise RuntimeError("transformer quant unsupported for this device/scheme")
return pipe, scheme
def _plan_memory(
self,
target: DiffusionDeviceTarget,
@ -751,6 +825,7 @@ class DiffusionBackend:
"speed_mode": None,
"speed_optims": [],
"text_encoder_quant": None,
"transformer_quant": None,
}
return {
"loaded": True,
@ -766,6 +841,7 @@ class DiffusionBackend:
"speed_mode": state.speed_mode,
"speed_optims": list(state.speed_optims),
"text_encoder_quant": state.text_encoder_quant,
"transformer_quant": state.transformer_quant,
}

View file

@ -0,0 +1,254 @@
# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
"""Opt-in low-precision quantisation of the diffusion DiT transformer.
The default path loads the transformer as a single-file GGUF, which stores weights
4-bit but DEQUANTISES to bf16 on every matmul -- so it runs at bf16 tensor-core rate
and never touches the int8 / fp8 / fp4 tensor cores. It is a memory win that costs
speed. This module is the opt-in alternative: load the DENSE bf16 transformer from the
base repo and torchao-quantise it with a DYNAMIC-ACTIVATION scheme so the matmul runs
on the low-precision tensor cores. Measured on a B200 (Z-Image-Turbo, 1024px / 8 steps)
vs the GGUF+compile default (0.802s, LPIPS 0.083 vs dense bf16): fp8 dynamic 0.585s
(1.37x), int8 dynamic 0.603s (1.33x), both at LOWER LPIPS than GGUF -- faster AND a hair
more accurate, at the cost of a higher-memory dense load. So it is strictly opt-in; the
loader keeps GGUF as the low-memory default and the fallback.
Scheme by architecture (``auto`` picks the best supported, best first):
nvfp4 / mxfp8 - Blackwell sm_100+ FP4 / MX tensor cores (biggest win; prototype).
fp8 - Ada / Hopper / Blackwell (sm_89+) fp8 tensor cores.
int8 - Ampere+ (sm_80+) int8 tensor cores -- the broadest-hardware lever.
Every scheme needs ``torch.compile`` to realise the speedup (dynamic quant is ~30x
slower eager); the loader already compiles the repeated block AFTER this runs. torch /
torchao are imported lazily so the module stays importable in a no-torch runtime, and
every probe is best-effort: an unsupported scheme yields None and the caller loads GGUF.
"""
from __future__ import annotations
from typing import Any, Optional
TQ_INT8 = "int8"
TQ_FP8 = "fp8"
TQ_NVFP4 = "nvfp4"
TQ_MXFP8 = "mxfp8"
TQ_AUTO = "auto"
TQ_SCHEMES = (TQ_INT8, TQ_FP8, TQ_NVFP4, TQ_MXFP8)
TQ_MODES = (TQ_AUTO,) + TQ_SCHEMES
# Skip linears whose in/out features are below this. The int8 dynamic path uses
# torch._int_mm, which requires the activation row count M > 16, and the DiT's tiny
# timestep / pooled / modulation projections run at M=1 and crash it. They are a
# negligible share of the FLOPs, so leaving them bf16 costs ~nothing (measured:
# 239/276 Z-Image linears quantised, full speedup) and keeps quality a touch higher.
DEFAULT_MIN_LINEAR_FEATURES = 512
# Per-architecture preference order for ``auto`` -- best (fastest, in-bar) first, with
# the lower-precision schemes listed as fallbacks for that arch tier. On Blackwell, fp8
# is preferred over mxfp8: measured on a B200, plain fp8 dynamic is both faster AND a
# touch more accurate than mxfp8 (block scaling adds overhead without a speed win here),
# so mxfp8 sits below fp8 as a fallback. nvfp4 stays first for its larger speedup when
# its kernels are available.
_AUTO_LADDER: tuple[tuple[tuple[int, int], tuple[str, ...]], ...] = (
((10, 0), (TQ_NVFP4, TQ_FP8, TQ_MXFP8, TQ_INT8)), # Blackwell sm_100+
((8, 9), (TQ_FP8, TQ_INT8)), # Ada sm_89 / Hopper sm_90
((8, 0), (TQ_INT8,)), # Ampere sm_80 / sm_86
)
# Cache of (scheme, device) -> bool so the quantise+matmul smoke test runs once.
_SMOKE_CACHE: dict[tuple[str, str], bool] = {}
def normalize_transformer_quant(value: Optional[str]) -> Optional[str]:
"""Lower/strip a requested transformer quant; None / "" / "none" / "off" -> None.
Raises ValueError for an unsupported value so a bad request is rejected cheaply."""
if value is None:
return None
normalized = str(value).strip().lower().replace("-", "_")
if not normalized or normalized in ("none", "off"):
return None
if normalized not in TQ_MODES:
raise ValueError(
f"Unsupported transformer_quant '{value}'. Use one of: {', '.join(TQ_MODES)}."
)
return normalized
def dense_transformer_supported(target: Any) -> bool:
"""Whether the dense-source quant path is usable for ``target``: a CUDA device with
a bf16 compute dtype (the only configuration any torchao dynamic scheme accelerates).
A cheap pre-check the loader runs before loading the (large) dense transformer."""
if getattr(target, "device", None) != "cuda":
return False
try:
import torch
return getattr(target, "dtype", None) is torch.bfloat16
except Exception:
return False
def select_transformer_quant_scheme(target: Any, requested: Optional[str]) -> 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."""
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:
return requested if _scheme_supported(requested, device) else None
cap = _capability()
if cap is None:
return None
for floor, schemes in _AUTO_LADDER:
if cap >= floor:
for scheme in schemes:
if _scheme_supported(scheme, device):
return scheme
return None
return None
def _capability() -> Optional[tuple[int, int]]:
try:
import torch
major, minor = torch.cuda.get_device_capability()
return (int(major), int(minor))
except Exception:
return None
def _scheme_supported(scheme: str, device: str) -> bool:
"""CUDA + (for fp8) the fp8 dtype + a cached quantise+matmul smoke test for ``scheme``."""
try:
import torch
if not torch.cuda.is_available():
return False
if scheme == TQ_FP8 and not hasattr(torch, "float8_e4m3fn"):
return False
except Exception:
return False
return _smoke_probe(scheme, device)
def _smoke_probe(scheme: str, device: str) -> bool:
"""True iff a tiny Linear quantised with ``scheme`` runs one M=32 forward without
error. Cached per (scheme, device). This is what makes ``auto`` robust to a torch /
torchao build where a prototype (nvfp4 / mxfp8) kernel is unavailable: it fails here
and the ladder moves on, rather than crashing at the first real denoise step."""
key = (scheme, device)
if key in _SMOKE_CACHE:
return _SMOKE_CACHE[key]
ok = False
try:
import torch
from torchao.quantization import quantize_
lin = torch.nn.Linear(512, 512, bias = False).to(device = device, dtype = torch.bfloat16)
quantize_(lin, _make_quant_config(scheme), filter_fn = make_filter_fn(0))
x = torch.randn(32, 512, device = device, dtype = torch.bfloat16)
with torch.no_grad():
lin(x)
torch.cuda.synchronize()
ok = True
except Exception:
ok = False
_SMOKE_CACHE[key] = ok
return ok
def _make_quant_config(scheme: str) -> Any:
"""The torchao dynamic-activation config for ``scheme`` (lazy import; prototype
import for the Blackwell fp4 / mx schemes is inside the branch that needs it)."""
from torchao.quantization import (
Float8DynamicActivationFloat8WeightConfig,
Int8DynamicActivationInt8WeightConfig,
)
if scheme == TQ_INT8:
return Int8DynamicActivationInt8WeightConfig()
if scheme == TQ_FP8:
return Float8DynamicActivationFloat8WeightConfig()
if scheme == TQ_NVFP4:
from torchao.prototype.mx_formats import NVFP4DynamicActivationNVFP4WeightConfig
return NVFP4DynamicActivationNVFP4WeightConfig()
if scheme == TQ_MXFP8:
import torch
from torchao.prototype.mx_formats import MXDynamicActivationMXWeightConfig
try:
return MXDynamicActivationMXWeightConfig(
activation_dtype = torch.float8_e4m3fn, weight_dtype = torch.float8_e4m3fn
)
except TypeError:
return MXDynamicActivationMXWeightConfig()
raise ValueError(f"unknown transformer quant scheme '{scheme}'")
def make_filter_fn(min_features: int):
"""A torchao ``quantize_`` filter keeping only the FLOP-heavy linears: nn.Linear
with both in/out features >= ``min_features``. Hides the (module, fqn) callback arity."""
def filter_fn(module: Any, fqn: str = "") -> bool:
try:
import torch
if not isinstance(module, torch.nn.Linear):
return False
except Exception:
return False
in_features = getattr(module, "in_features", None)
out_features = getattr(module, "out_features", None)
if in_features is None or out_features is None:
return False
return in_features >= min_features and out_features >= min_features
return filter_fn
def quantize_transformer(
pipe: Any,
target: Any,
*,
mode: Optional[str],
min_features: int = DEFAULT_MIN_LINEAR_FEATURES,
logger: Any = None,
) -> Optional[str]:
"""Quantise ``pipe.transformer``'s FLOP-heavy linears in place with the arch-chosen
dynamic scheme. Returns the scheme actually engaged, or None when disabled /
unsupported / failed -- the caller then loads GGUF instead. Best-effort: it never
raises for an ordinary unsupported environment (a failure leaves the module dense)."""
scheme = select_transformer_quant_scheme(target, mode)
if scheme is None:
return None
transformer = getattr(pipe, "transformer", None)
if transformer is None:
return None
try:
from torchao.quantization import quantize_
quantize_(transformer, _make_quant_config(scheme), filter_fn = make_filter_fn(min_features))
# Runtime-only marker (torchao tensors are not safetensors-serializable; this
# backend is inference-only, so this is purely diagnostic).
try:
transformer._unsloth_runtime_quant = scheme
except Exception: # noqa: BLE001 — marker is best-effort
pass
return scheme
except Exception as exc: # noqa: BLE001 — leave the transformer dense -> GGUF fallback
_warn(logger, scheme, exc)
return None
def _warn(logger: Any, what: str, exc: Exception) -> None:
if logger is not None:
logger.warning("diffusion.transformer_quant: %s failed: %s", what, exc)

View file

@ -1721,6 +1721,15 @@ class DiffusionLoadRequest(BaseModel):
"memory-vs-quality tradeoff (shifts fine detail), not free; "
"pairs well with balanced mode.",
)
transformer_quant: Optional[Literal["auto", "int8", "fp8", "nvfp4", "mxfp8"]] = Field(
None,
description = "Opt-in fast transformer: load the DENSE bf16 transformer instead "
"of the GGUF and torchao-quantise it onto the low-precision tensor "
"cores (faster than GGUF's bf16-rate dequant, at higher VRAM). auto "
"picks the best for the GPU (Blackwell nvfp4/mxfp8, Ada/Hopper fp8, "
"Ampere int8); an explicit scheme forces it. Needs CUDA + bf16 + room "
"for the dense load; falls back to GGUF otherwise.",
)
class DiffusionGenerateRequest(BaseModel):
@ -1828,3 +1837,8 @@ class DiffusionStatusResponse(BaseModel):
text_encoder_quant: Optional[str] = Field(
None, description = "Text-encoder quantisation engaged: fp8 | nvfp4 | null"
)
transformer_quant: Optional[str] = Field(
None,
description = "Transformer quant engaged on the dense fast path: int8 | fp8 | "
"nvfp4 | mxfp8 | null (null = the GGUF transformer was loaded)",
)

View file

@ -10085,6 +10085,7 @@ async def load_diffusion_model(
memory_mode = request.memory_mode,
speed_mode = request.speed_mode,
text_encoder_quant = request.text_encoder_quant,
transformer_quant = request.transformer_quant,
)
return DiffusionStatusResponse(**status_dict)
except (ValueError, FileNotFoundError) as exc:

View file

@ -885,3 +885,120 @@ def test_load_fast_mode_stays_resident_on_cuda(fake_runtime, tmp_path, monkeypat
)
assert status["offload_policy"] == "none" and status["cpu_offload"] is False
assert backend._state.pipe.moved_to == "cuda"
# ── transformer quant (opt-in dense fast path) ────────────────────────────────
def _stub_dense_quant(monkeypatch, *, scheme = "fp8"):
"""Force the dense+quant branch hermetically: a supported dense source, a
from_pretrained on the fake transformer, and a quantizer that engages `scheme`.
Returns a dict recording the dense-loader / quantizer calls."""
from core.inference import diffusion as dmod
calls: dict = {"from_pretrained": 0, "quantize": 0, "quant_mode": None}
@classmethod
def _from_pretrained(cls, base, **kwargs):
calls["from_pretrained"] += 1
calls["fp_kwargs"] = {"base": base, **kwargs}
return object()
monkeypatch.setattr(_FakeTransformer, "from_pretrained", _from_pretrained, raising = False)
monkeypatch.setattr(dmod, "dense_transformer_supported", lambda target: True)
def _quantize(pipe, target, *, mode, **kw):
calls["quantize"] += 1
calls["quant_mode"] = mode
return scheme
monkeypatch.setattr(dmod, "quantize_transformer", _quantize)
return calls
def test_default_load_skips_dense_quant_path(fake_runtime, tmp_path, monkeypatch):
# With no transformer_quant flag the GGUF path is taken and the dense gate is
# never even consulted (short-circuit), so the default cannot regress.
from core.inference import diffusion as dmod
monkeypatch.setattr(
dmod, "dense_transformer_supported",
lambda *a, **k: pytest.fail("dense path must not run without the flag"),
)
(tmp_path / "m.gguf").write_bytes(b"x")
backend = DiffusionBackend()
status = backend.load_pipeline(str(tmp_path), gguf_filename = "m.gguf", family_override = "z-image")
assert status["transformer_quant"] is None
assert _FakeTransformer.last["path"] # GGUF from_single_file was used
def test_transformer_quant_dense_path_engaged(fake_runtime, tmp_path, monkeypatch):
# transformer_quant + a CUDA resident plan -> load the DENSE transformer from the
# base repo, place it on the device, quantise it, and report the engaged scheme.
backend = DiffusionBackend()
_force_cuda_target(backend, monkeypatch)
calls = _stub_dense_quant(monkeypatch, scheme = "fp8")
(tmp_path / "m.gguf").write_bytes(b"x")
status = backend.load_pipeline(
str(tmp_path), gguf_filename = "m.gguf", family_override = "z-image",
transformer_quant = "fp8",
)
assert status["transformer_quant"] == "fp8"
assert calls["from_pretrained"] == 1 and calls["quantize"] == 1
assert calls["quant_mode"] == "fp8"
assert calls["fp_kwargs"]["subfolder"] == "transformer" # dense transformer subfolder
# The GGUF single-file path was NOT used for the transformer.
assert _FakeTransformer.last == {}
# quantize ran on-device: the dense pipe was placed on cuda (before compile).
assert backend._state.pipe.moved_to == "cuda"
assert status["offload_policy"] == "none"
def test_transformer_quant_falls_back_to_gguf_on_failure(fake_runtime, tmp_path, monkeypatch):
# A dense/quant failure (here: quantize returns None -> unsupported) must fall back
# to the GGUF build, not error -- status reports no transformer_quant engaged.
from core.inference import diffusion as dmod
backend = DiffusionBackend()
_force_cuda_target(backend, monkeypatch)
monkeypatch.setattr(dmod, "dense_transformer_supported", lambda target: True)
@classmethod
def _from_pretrained(cls, base, **kwargs):
return object()
monkeypatch.setattr(_FakeTransformer, "from_pretrained", _from_pretrained, raising = False)
monkeypatch.setattr(dmod, "quantize_transformer", lambda pipe, target, **kw: None)
(tmp_path / "m.gguf").write_bytes(b"x")
status = backend.load_pipeline(
str(tmp_path), gguf_filename = "m.gguf", family_override = "z-image",
transformer_quant = "fp8",
)
assert status["loaded"] is True
assert status["transformer_quant"] is None # fell back
assert _FakeTransformer.last["path"] # GGUF from_single_file used
def test_transformer_quant_skipped_when_plan_offloads(fake_runtime, tmp_path, monkeypatch):
# The dense bf16 transformer only fits resident, so when the memory plan would
# offload (here low_vram) the fast path is skipped and GGUF loads instead -- the
# dense transformer is never even loaded.
from core.inference import diffusion as dmod
backend = DiffusionBackend()
_force_cuda_target(backend, monkeypatch)
monkeypatch.setattr(dmod, "dense_transformer_supported", lambda target: True)
@classmethod
def _fp_fail(cls, *a, **k):
pytest.fail("dense transformer must not load when the plan offloads")
monkeypatch.setattr(_FakeTransformer, "from_pretrained", _fp_fail, raising = False)
(tmp_path / "m.gguf").write_bytes(b"x")
status = backend.load_pipeline(
str(tmp_path), gguf_filename = "m.gguf", family_override = "z-image",
transformer_quant = "fp8", memory_mode = "low_vram",
)
assert status["transformer_quant"] is None
assert status["offload_policy"] == "model"
assert _FakeTransformer.last["path"] # GGUF path used

View file

@ -316,6 +316,28 @@ def test_memory_mode_threads_through_to_backend(client, monkeypatch):
assert backend.last_load_kwargs.get("memory_mode") == "low_vram"
def test_transformer_quant_threads_through_to_backend(client, monkeypatch):
backend = _FakeBackend()
monkeypatch.setattr(diffusion_module, "get_diffusion_backend", lambda: backend)
resp = client.post(
"/api/inference/images/load",
json = {"model_path": "x/z-image", "gguf_filename": "q.gguf", "transformer_quant": "auto"},
)
assert resp.status_code == 200
assert backend.last_load_kwargs.get("transformer_quant") == "auto"
def test_invalid_transformer_quant_returns_422_without_eviction(client):
# An unsupported transformer_quant is rejected by the request schema (Literal), so
# the GPU is never acquired and no chat model is evicted.
resp = client.post(
"/api/inference/images/load",
json = {"model_path": "x/z-image", "gguf_filename": "q.gguf", "transformer_quant": "int2"},
)
assert resp.status_code == 422
assert gpu_arbiter._owner is None
def test_invalid_memory_mode_returns_422_without_eviction(client):
# An unsupported memory_mode is rejected by the request schema (Literal), so the
# GPU is never acquired and no chat model is evicted.

View file

@ -0,0 +1,245 @@
# 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 transformer quantisation (``diffusion_transformer_quant.py``).
Hermetic: torch + torchao are stubbed via ``sys.modules``, and the per-scheme smoke
probe (``_scheme_supported`` / ``_smoke_probe``) is monkeypatched where the test cares
about the selection ladder rather than the GPU probe, so everything runs CPU-only.
"""
from __future__ import annotations
import sys
import types
import pytest
import core.inference.diffusion_transformer_quant as tq
from core.inference.diffusion_transformer_quant import (
TQ_FP8,
TQ_INT8,
TQ_MXFP8,
TQ_NVFP4,
dense_transformer_supported,
make_filter_fn,
normalize_transformer_quant,
quantize_transformer,
select_transformer_quant_scheme,
)
def _target(*, device = "cuda", dtype = "bfloat16"):
return types.SimpleNamespace(device = device, dtype = dtype)
def _stub_torch(monkeypatch, *, cc = (10, 0), with_fp8 = True, cuda_available = True):
torch = types.ModuleType("torch")
torch.bfloat16 = "bfloat16"
torch.float16 = "float16"
if with_fp8:
torch.float8_e4m3fn = "float8_e4m3fn"
torch.cuda = types.SimpleNamespace(
is_available = lambda: cuda_available,
get_device_capability = lambda *a: cc,
)
monkeypatch.setitem(sys.modules, "torch", torch)
return torch
# ── normalisation ─────────────────────────────────────────────────────────────
def test_normalize_transformer_quant():
assert normalize_transformer_quant(None) is None
assert normalize_transformer_quant("") is None
assert normalize_transformer_quant("none") is None
assert normalize_transformer_quant("off") is None
assert normalize_transformer_quant("AUTO") == "auto"
assert normalize_transformer_quant("INT8") == TQ_INT8
assert normalize_transformer_quant("fp8") == TQ_FP8
with pytest.raises(ValueError):
normalize_transformer_quant("int2")
# ── dense-source gate ───────────────────────────────────────────────────────────
def test_dense_transformer_supported_requires_cuda_bf16(monkeypatch):
_stub_torch(monkeypatch)
assert dense_transformer_supported(_target()) is True
assert dense_transformer_supported(_target(device = "cpu")) is False
assert dense_transformer_supported(_target(dtype = "float16")) is False
# ── scheme selection ladder ─────────────────────────────────────────────────────
def _allow(monkeypatch, allowed):
"""Force ``_scheme_supported`` to accept only ``allowed`` (simulates smoke results)."""
monkeypatch.setattr(tq, "_scheme_supported", lambda scheme, device: scheme in allowed)
def test_auto_blackwell_prefers_nvfp4_then_falls_back(monkeypatch):
_stub_torch(monkeypatch, cc = (10, 0))
_allow(monkeypatch, {TQ_NVFP4, TQ_MXFP8, TQ_FP8, TQ_INT8})
assert select_transformer_quant_scheme(_target(), "auto") == TQ_NVFP4
# nvfp4 unavailable: fp8 is preferred over mxfp8 (measured faster + a touch more
# accurate on B200), even though mxfp8 is also supported.
_allow(monkeypatch, {TQ_MXFP8, TQ_FP8, TQ_INT8})
assert select_transformer_quant_scheme(_target(), "auto") == TQ_FP8
# Only mxfp8 + int8 left -> mxfp8 (still above int8).
_allow(monkeypatch, {TQ_MXFP8, TQ_INT8})
assert select_transformer_quant_scheme(_target(), "auto") == TQ_MXFP8
# Only int8 usable -> int8.
_allow(monkeypatch, {TQ_INT8})
assert select_transformer_quant_scheme(_target(), "auto") == TQ_INT8
def test_auto_ada_hopper_prefers_fp8(monkeypatch):
_stub_torch(monkeypatch, cc = (8, 9))
_allow(monkeypatch, {TQ_NVFP4, TQ_MXFP8, TQ_FP8, TQ_INT8})
assert select_transformer_quant_scheme(_target(), "auto") == TQ_FP8
_stub_torch(monkeypatch, cc = (9, 0)) # Hopper
assert select_transformer_quant_scheme(_target(), "auto") == TQ_FP8
def test_auto_ampere_prefers_int8(monkeypatch):
_stub_torch(monkeypatch, cc = (8, 0))
_allow(monkeypatch, {TQ_FP8, TQ_INT8}) # fp8 cores absent on Ampere -> int8 only in ladder
assert select_transformer_quant_scheme(_target(), "auto") == TQ_INT8
_stub_torch(monkeypatch, cc = (8, 6))
assert select_transformer_quant_scheme(_target(), "auto") == TQ_INT8
def test_auto_pre_ampere_unsupported(monkeypatch):
_stub_torch(monkeypatch, cc = (7, 5)) # Turing: below the int8-dynamic floor
_allow(monkeypatch, {TQ_INT8, TQ_FP8})
assert select_transformer_quant_scheme(_target(), "auto") is None
def test_explicit_scheme_honored_or_none(monkeypatch):
_stub_torch(monkeypatch, cc = (8, 0))
_allow(monkeypatch, {TQ_INT8})
assert select_transformer_quant_scheme(_target(), "int8") == TQ_INT8
# Explicit unsupported scheme is NOT silently downgraded -> None (-> GGUF fallback).
assert select_transformer_quant_scheme(_target(), "fp8") is None
assert select_transformer_quant_scheme(_target(), "nvfp4") is None
def test_select_none_when_disabled_or_non_cuda(monkeypatch):
_stub_torch(monkeypatch)
_allow(monkeypatch, {TQ_INT8, TQ_FP8, TQ_NVFP4})
assert select_transformer_quant_scheme(_target(), None) is None
assert select_transformer_quant_scheme(_target(device = "cpu"), "auto") is None
# ── _scheme_supported / _smoke_probe ────────────────────────────────────────────
def test_scheme_supported_shortcircuits(monkeypatch):
# No CUDA -> False without running the smoke probe.
_stub_torch(monkeypatch, cuda_available = False)
monkeypatch.setattr(tq, "_smoke_probe", lambda *a: pytest.fail("probe should not run"))
assert tq._scheme_supported(TQ_INT8, "cuda") is False
# fp8 requested but the fp8 dtype is missing -> False before the probe.
_stub_torch(monkeypatch, with_fp8 = False)
monkeypatch.setattr(tq, "_smoke_probe", lambda *a: pytest.fail("probe should not run"))
assert tq._scheme_supported(TQ_FP8, "cuda") is False
def test_smoke_probe_caches_and_tolerates_failure(monkeypatch):
tq._SMOKE_CACHE.clear()
calls = {"n": 0}
class _Lin:
def __init__(self, *a, **k):
pass
def to(self, **k):
return self
torch = types.ModuleType("torch")
torch.bfloat16 = "bfloat16"
torch.nn = types.SimpleNamespace(Linear = _Lin)
torch.randn = lambda *a, **k: object()
torch.no_grad = lambda: __import__("contextlib").nullcontext()
torch.cuda = types.SimpleNamespace(is_available = lambda: True, synchronize = lambda: None)
monkeypatch.setitem(sys.modules, "torch", torch)
tqz = types.ModuleType("torchao.quantization")
def _quantize_ok(module, config, filter_fn = None):
calls["n"] += 1
tqz.quantize_ = _quantize_ok
tqz.Int8DynamicActivationInt8WeightConfig = lambda: "int8cfg"
tqz.Float8DynamicActivationFloat8WeightConfig = lambda: "fp8cfg"
monkeypatch.setitem(sys.modules, "torchao.quantization", tqz)
# _Lin is callable? No -> the forward lin(x) would fail. Make instances callable.
_Lin.__call__ = lambda self, x: x
assert tq._smoke_probe(TQ_INT8, "cuda") is True
assert tq._smoke_probe(TQ_INT8, "cuda") is True # cached, no second quantize_
assert calls["n"] == 1
# A scheme whose quantize_ raises -> probe False (and cached).
tq._SMOKE_CACHE.clear()
def _quantize_boom(module, config, filter_fn = None):
raise RuntimeError("kernel unavailable")
tqz.quantize_ = _quantize_boom
assert tq._smoke_probe(TQ_FP8, "cuda") is False
# ── filter ──────────────────────────────────────────────────────────────────────
def test_make_filter_fn(monkeypatch):
class _Lin:
def __init__(self, i, o):
self.in_features, self.out_features = i, o
torch = types.ModuleType("torch")
torch.nn = types.SimpleNamespace(Linear = _Lin)
monkeypatch.setitem(sys.modules, "torch", torch)
keep = make_filter_fn(512)
assert keep(_Lin(1024, 4096), "blocks.0.attn.to_q") is True
assert keep(_Lin(256, 4096), "time_proj") is False # small in_features -> skip
assert keep(_Lin(4096, 256), "out_proj") is False # small out_features -> skip
assert keep(object(), "not_linear") is False # non-Linear -> skip
assert keep(types.SimpleNamespace(), "no_attrs") is False
# ── apply ───────────────────────────────────────────────────────────────────────
def test_quantize_transformer_applies_and_marks(monkeypatch):
monkeypatch.setattr(tq, "select_transformer_quant_scheme", lambda target, mode: TQ_FP8)
monkeypatch.setattr(tq, "_make_quant_config", lambda scheme: f"{scheme}cfg")
recorder: list = []
tqz = types.ModuleType("torchao.quantization")
tqz.quantize_ = lambda module, config, filter_fn = None: recorder.append((module, config, filter_fn))
monkeypatch.setitem(sys.modules, "torchao.quantization", tqz)
transformer = types.SimpleNamespace()
pipe = types.SimpleNamespace(transformer = transformer)
assert quantize_transformer(pipe, _target(), mode = "fp8") == TQ_FP8
assert len(recorder) == 1 and recorder[0][0] is transformer and recorder[0][1] == "fp8cfg"
assert callable(recorder[0][2]) # a filter_fn was passed
assert transformer._unsloth_runtime_quant == TQ_FP8 # diagnostic marker set
def test_quantize_transformer_none_when_unsupported(monkeypatch):
monkeypatch.setattr(tq, "select_transformer_quant_scheme", lambda target, mode: 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, "_make_quant_config", lambda scheme: "cfg")
tqz = types.ModuleType("torchao.quantization")
def _boom(module, config, filter_fn = None):
raise RuntimeError("partial quant failure")
tqz.quantize_ = _boom
monkeypatch.setitem(sys.modules, "torchao.quantization", tqz)
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