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:
parent
ede94176f6
commit
983f2c14a0
9 changed files with 1040 additions and 28 deletions
|
|
@ -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
273
scripts/quant_probe.py
Normal 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())
|
||||
|
|
@ -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,
|
||||
}
|
||||
|
||||
|
||||
|
|
|
|||
254
studio/backend/core/inference/diffusion_transformer_quant.py
Normal file
254
studio/backend/core/inference/diffusion_transformer_quant.py
Normal 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)
|
||||
|
|
@ -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)",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
245
studio/backend/tests/test_diffusion_transformer_quant.py
Normal file
245
studio/backend/tests/test_diffusion_transformer_quant.py
Normal 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
|
||||
Loading…
Add table
Add a link
Reference in a new issue