Studio diffusion (Phase 9): pre-quantized transformer loading

The Phase 8 fast transformer_quant path materialises the dense bf16 transformer on
the GPU and torchao-quantises it in place, so its load peak is ~2x GGUF's (~21 vs
13.4 GB) plus a ~12 GB download. Add a pre-quantized branch: quantise once offline
(scripts/build_prequant_checkpoint.py) and at runtime build the transformer skeleton
on the meta device (accelerate.init_empty_weights) and load_state_dict(assign=True)
the quantized weights, so the dense bf16 never touches the GPU.

Measured (B200, Z-Image fp8): full-pipeline GPU load peak 21.2 -> 14.6 GB (matching
GGUF's 13.4), on-disk 12 -> 6.28 GB, output bit-identical (LPIPS 0.0). It is the same
torchao config + min_features filter the runtime path uses, applied ahead of time.

New core/inference/diffusion_prequant.py (resolve_prequant_source +
load_prequantized_transformer, best-effort, lazy imports). diffusion.py
_load_dense_quant_pipeline tries the pre-quant source first and falls back to the
dense materialise+quantise path, then to GGUF, so the default is unchanged.
DiffusionLoadRequest gains transformer_prequant_path; DiffusionFamily gains an empty
prequant_repos map for hosted checkpoints (hosting deferred). Hermetic CPU tests for
the resolver, the meta-init+assign loader, and the backend branch selection +
fallbacks; GPU verification via scripts/verify_prequant_backend.py.
This commit is contained in:
Daniel Han 2026-06-26 11:23:20 +00:00
commit b90f833469
11 changed files with 968 additions and 10 deletions

View file

@ -0,0 +1,127 @@
# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
"""Build a pre-quantized transformer checkpoint for the Studio diffusion fast path.
Quantise a model's dense bf16 DiT transformer ONCE and save the quantized state dict, so
the backend can load the already-quantized weights at runtime (meta-init +
load_state_dict(assign=True)) instead of materialising the dense bf16 on the GPU. That
drops the transformer GPU load peak ~2x and the download ~2x for fp8 (measured on Z-Image:
12.9 -> 6.3 GB peak, 12 -> 6.28 GB on disk), with bit-identical output -- it is the exact
same torchao config + min_features filter the runtime path uses, applied ahead of time.
Run on one CUDA (Blackwell / Ada / Hopper) GPU. fp8 works on torch 2.9+; the FP4/MX schemes
need the newer kernels (see scripts/nvfp4_t211_probe.py).
python scripts/build_prequant_checkpoint.py \
--base Tongyi-MAI/Z-Image-Turbo --family z-image --scheme fp8 \
--out outputs/quant_research/prequant_fp8/transformer_fp8.pt [--upload-repo ORG/REPO]
"""
from __future__ import annotations
import argparse
import sys
import time
from pathlib import Path
BACKEND = Path(__file__).resolve().parent.parent / "studio" / "backend"
def main(argv=None) -> int:
p = argparse.ArgumentParser()
p.add_argument("--base", required=True, help="diffusers base repo (carries the transformer subfolder)")
p.add_argument("--family", required=True, help="diffusion family name/alias (e.g. z-image)")
p.add_argument("--scheme", required=True, help="quant scheme: int8 | fp8 | nvfp4 | mxfp8")
p.add_argument("--out", required=True, help="output .pt path for the checkpoint")
p.add_argument("--min-features", type=int, default=512)
p.add_argument("--dtype", default="bfloat16", choices=["bfloat16"])
p.add_argument("--hf-token", default=None)
p.add_argument("--upload-repo", default=None, help="optional HF repo id to upload the checkpoint to")
p.add_argument("--upload-revision", default=None)
args = p.parse_args(argv)
sys.path.insert(0, str(BACKEND))
import torch
import torchao
import diffusers
from core.inference.diffusion_families import detect_family
from core.inference.diffusion_prequant import PREQUANT_FORMAT, prequant_filename
# Reuse the runtime quant factory + filter so offline == runtime (the LPIPS-0 invariant).
from core.inference.diffusion_transformer_quant import (
TQ_SCHEMES,
_make_quant_config,
make_filter_fn,
)
from torchao.quantization import quantize_
scheme = args.scheme.strip().lower()
if scheme not in TQ_SCHEMES:
print(f"error: --scheme must be one of {TQ_SCHEMES} (not 'auto')", flush=True)
return 2
fam = detect_family(args.base, override=args.family)
if fam is None:
print(f"error: unknown family '{args.family}'", flush=True)
return 2
transformer_cls = getattr(diffusers, fam.transformer_class)
print(f"== build prequant ({fam.name}/{scheme}, min_feat={args.min_features}) ==", flush=True)
print(f" loading dense transformer from {args.base} (subfolder=transformer) ...", flush=True)
t0 = time.time()
transformer = transformer_cls.from_pretrained(
args.base, subfolder="transformer", torch_dtype=torch.bfloat16, token=args.hf_token
).to("cuda")
print(f" quantising in place ({scheme}) ...", flush=True)
quantize_(transformer, _make_quant_config(scheme), filter_fn=make_filter_fn(args.min_features))
# Move the state dict to CPU for a portable, GPU-free artifact.
state_dict = {
k: (v.detach().to("cpu") if hasattr(v, "detach") else v)
for k, v in transformer.state_dict().items()
}
ckpt = {
"format": PREQUANT_FORMAT,
"metadata": {
"base_model_id": args.base,
"family": fam.name,
"scheme": scheme,
"min_features": args.min_features,
"torch_dtype": args.dtype,
"quant_backend": "torchao",
"transformer_class": fam.transformer_class,
"torch_version": torch.__version__,
"torchao_version": getattr(torchao, "__version__", "?"),
"diffusers_version": diffusers.__version__,
},
"state_dict": state_dict,
}
out = Path(args.out)
out.parent.mkdir(parents=True, exist_ok=True)
torch.save(ckpt, out)
size_gb = out.stat().st_size / 1e9
print(f" saved {out} ({size_gb:.2f} GB) in {time.time() - t0:.0f}s", flush=True)
print(f" metadata: {ckpt['metadata']}", flush=True)
if args.upload_repo:
from huggingface_hub import HfApi
dest = prequant_filename(scheme)
print(f" uploading -> {args.upload_repo}:{dest} ...", flush=True)
api = HfApi(token=args.hf_token)
api.create_repo(args.upload_repo, exist_ok=True)
api.upload_file(
path_or_fileobj=str(out),
path_in_repo=dest,
repo_id=args.upload_repo,
revision=args.upload_revision,
)
print(f" uploaded {dest} to {args.upload_repo}", flush=True)
print("BUILD-PREQUANT-DONE", flush=True)
return 0
if __name__ == "__main__":
sys.exit(main())

170
scripts/prequant_probe.py Normal file
View file

@ -0,0 +1,170 @@
# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
"""Does a *pre-quantized* checkpoint fix the dense-quant load-VRAM spike?
The current fast-transformer path materialises the dense bf16 transformer on the GPU
and quantises it in place -> ~2x the GGUF load peak. This probe checks the fix: quantise
once, ``torch.save`` the quantized state dict, then load it onto an empty (meta) model
with ``load_state_dict(assign=True)`` so the bf16 never touches the GPU.
Modes (run each in its own process so peak VRAM is clean):
build -- load dense bf16, quantize_ fp8, torch.save the state dict + on-disk size.
baseline -- current path: from_pretrained bf16 -> quantize_ on GPU. Report load peak + gen.
prequant -- meta-init -> load_state_dict(saved, assign=True) -> cuda. Report load peak + gen.
Run on one CUDA (Blackwell) GPU. Reference image for LPIPS is the baseline path."""
from __future__ import annotations
import argparse
import sys
import time
from pathlib import Path
import numpy as np
BASE = "Tongyi-MAI/Z-Image-Turbo"
PROMPT = "A cinematic photograph of a red fox in a snowy forest at dawn, highly detailed"
ROOT = Path("/mnt/disks/unslothai/ubuntu/workspace_81/outputs/quant_research")
CKPT = ROOT / "prequant_fp8" / "transformer_fp8_state.pt"
OUT = ROOT / "prequant_images"
MIN_FEAT = 512
def _filt(mod, fqn=""):
import torch.nn as nn
return isinstance(mod, nn.Linear) and mod.in_features >= MIN_FEAT and mod.out_features >= MIN_FEAT
def _fp8_cfg():
from torchao.quantization import Float8DynamicActivationFloat8WeightConfig
return Float8DynamicActivationFloat8WeightConfig()
def _build():
import torch
import diffusers
from torchao.quantization import quantize_
torch.cuda.reset_peak_memory_stats()
t = diffusers.ZImageTransformer2DModel.from_pretrained(
BASE, subfolder="transformer", torch_dtype=torch.bfloat16).to("cuda")
quantize_(t, _fp8_cfg(), filter_fn=_filt)
CKPT.parent.mkdir(parents=True, exist_ok=True)
sd = t.state_dict()
# move to cpu for a portable, gpu-free checkpoint
sd = {k: (v.detach().to("cpu") if hasattr(v, "detach") else v) for k, v in sd.items()}
torch.save(sd, CKPT)
sz = CKPT.stat().st_size / 1e9
peak = torch.cuda.max_memory_allocated() / 1e9
print(f"[build] saved {CKPT.name} on-disk={sz:.2f} GB build_gpu_peak={peak:.1f} GB", flush=True)
return 0
def _make_pipe_from_transformer(t):
import diffusers
import torch
pipe = diffusers.ZImagePipeline.from_pretrained(BASE, torch_dtype=torch.bfloat16, transformer=t)
pipe.to("cuda")
return pipe
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 _baseline(steps, seed, res):
import torch
import diffusers
from torchao.quantization import quantize_
torch.cuda.reset_peak_memory_stats()
t = diffusers.ZImageTransformer2DModel.from_pretrained(
BASE, subfolder="transformer", torch_dtype=torch.bfloat16).to("cuda")
quantize_(t, _fp8_cfg(), filter_fn=_filt)
load_peak = torch.cuda.max_memory_allocated() / 1e9
pipe = _make_pipe_from_transformer(t)
img, dt = _gen(pipe, steps, seed, res) # warmup
img, dt = _gen(pipe, steps, seed, res)
OUT.mkdir(parents=True, exist_ok=True)
img.save(OUT / "baseline.png")
print(f"[baseline] transformer_load_gpu_peak={load_peak:.1f} GB gen={dt:.3f}s", flush=True)
return 0
def _prequant(steps, seed, res):
import torch
import diffusers
from accelerate import init_empty_weights
if not CKPT.exists():
print(f"[prequant] missing checkpoint {CKPT}; run --mode build first", flush=True)
return 1
torch.cuda.reset_peak_memory_stats()
cfg = diffusers.ZImageTransformer2DModel.load_config(BASE, subfolder="transformer")
with init_empty_weights():
t = diffusers.ZImageTransformer2DModel.from_config(cfg)
sd = torch.load(CKPT, weights_only=False, map_location="cpu")
missing, unexpected = t.load_state_dict(sd, strict=False, assign=True)
# any param/buffer still on meta (e.g. non-persistent buffers) -> materialise on cuda
leftover = [n for n, p in t.named_parameters() if p.is_meta] + [n for n, b in t.named_buffers() if b.is_meta]
if leftover:
print(f"[prequant] {len(leftover)} meta leftovers (non-persistent buffers): {leftover[:4]}", flush=True)
t = t.to_empty(device="cuda") # fallback path; re-loads sd below
t.load_state_dict(sd, strict=False, assign=True)
t = t.to(torch.bfloat16).to("cuda")
load_peak = torch.cuda.max_memory_allocated() / 1e9
print(f"[prequant] missing={len(missing)} unexpected={len(unexpected)} "
f"transformer_load_gpu_peak={load_peak:.1f} GB", flush=True)
pipe = _make_pipe_from_transformer(t)
img, dt = _gen(pipe, steps, seed, res) # warmup
img, dt = _gen(pipe, steps, seed, res)
OUT.mkdir(parents=True, exist_ok=True)
img.save(OUT / "prequant.png")
# LPIPS vs baseline if present
bpath = OUT / "baseline.png"
lp = None
if bpath.exists():
try:
import lpips
from PIL import Image
fn = lpips.LPIPS(net="alex", verbose=False).cuda().eval()
def tt(p):
a = np.array(Image.open(p).convert("RGB"))
return (torch.from_numpy(a).float().permute(2, 0, 1).unsqueeze(0) / 127.5 - 1.0).cuda()
with torch.no_grad():
lp = float(fn(tt(bpath), tt(OUT / "prequant.png")).item())
except Exception as exc: # noqa: BLE001
print(f" (lpips: {type(exc).__name__})", flush=True)
print(f"[prequant] gen={dt:.3f}s LPIPS_vs_baseline={lp}", flush=True)
return 0
def main(argv=None) -> int:
p = argparse.ArgumentParser()
p.add_argument("--mode", choices=["build", "baseline", "prequant"], required=True)
p.add_argument("--steps", type=int, default=8)
p.add_argument("--res", type=int, default=1024)
p.add_argument("--seed", type=int, default=42)
args = p.parse_args(argv)
if args.mode == "build":
return _build()
if args.mode == "baseline":
return _baseline(args.steps, args.seed, args.res)
return _prequant(args.steps, args.seed, args.res)
if __name__ == "__main__":
sys.path.insert(0, str(Path(__file__).resolve().parent.parent / "studio" / "backend"))
rc = main()
print("PREQUANT-PROBE-DONE", flush=True)
sys.exit(rc)

View file

@ -0,0 +1,126 @@
# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
"""GPU verification of the Phase 9 pre-quantized load path through the real backend code.
Exercises the actual product functions (``load_prequantized_transformer`` and the runtime
``quantize_transformer``), not a reimplementation:
prequant -- load the checkpoint built by build_prequant_checkpoint.py via the real
``load_prequantized_transformer`` (meta-init + assign), measure GPU load peak,
generate.
runtime -- the existing path: from_pretrained dense bf16 -> ``quantize_transformer`` on
device, measure GPU load peak, generate (the LPIPS reference).
Asserts the prequant load peak is far below the dense one and the images match (LPIPS ~0).
Run each mode in its own process for a clean peak. One CUDA GPU."""
from __future__ import annotations
import argparse
import logging
import sys
import time
from pathlib import Path
import numpy as np
BACKEND = Path(__file__).resolve().parent.parent / "studio" / "backend"
BASE = "Tongyi-MAI/Z-Image-Turbo"
CKPT = "/mnt/disks/unslothai/ubuntu/workspace_81/outputs/quant_research/prequant_fp8/transformer_fp8.pt"
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/prequant_verify_images")
logging.basicConfig(level=logging.INFO, format="%(message)s")
LOGGER = logging.getLogger("verify_prequant")
def _target(dtype):
import types
return types.SimpleNamespace(device="cuda", dtype=dtype)
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 _lpips(ref, arr):
try:
import lpips, torch
fn = lpips.LPIPS(net="alex", verbose=False).cuda().eval()
def t(x):
return (torch.from_numpy(x).float().permute(2, 0, 1).unsqueeze(0) / 127.5 - 1.0).cuda()
with torch.no_grad():
return float(fn(t(ref), t(arr)).item())
except Exception as exc: # noqa: BLE001
print(f" (lpips: {type(exc).__name__})", flush=True)
return None
def run(mode, steps, seed, res):
sys.path.insert(0, str(BACKEND))
import torch
import diffusers
from core.inference.diffusion_prequant import PrequantSource, load_prequantized_transformer
from core.inference.diffusion_transformer_quant import quantize_transformer
OUT.mkdir(parents=True, exist_ok=True)
transformer_cls = diffusers.ZImageTransformer2DModel
torch.cuda.reset_peak_memory_stats(); torch.cuda.empty_cache()
if mode == "prequant":
source = PrequantSource(kind="path", location=CKPT, filename=None)
transformer = load_prequantized_transformer(
transformer_cls, BASE, source, device="cuda", dtype=torch.bfloat16,
hf_token=None, scheme="fp8", logger=LOGGER)
if transformer is None:
print("prequant load FAILED (returned None)", flush=True)
return 1
pipe = diffusers.ZImagePipeline.from_pretrained(BASE, torch_dtype=torch.bfloat16, transformer=transformer)
pipe.to("cuda")
load_peak = torch.cuda.max_memory_allocated() / 1e9
marker = getattr(transformer, "_unsloth_runtime_quant", None)
print(f"[prequant] load_gpu_peak={load_peak:.1f} GB marker={marker}", flush=True)
else: # runtime
transformer = transformer_cls.from_pretrained(BASE, subfolder="transformer", torch_dtype=torch.bfloat16).to("cuda")
pipe = diffusers.ZImagePipeline.from_pretrained(BASE, torch_dtype=torch.bfloat16, transformer=transformer)
pipe.to("cuda")
scheme = quantize_transformer(pipe, _target(torch.bfloat16), mode="fp8", logger=LOGGER)
load_peak = torch.cuda.max_memory_allocated() / 1e9
print(f"[runtime] engaged={scheme} load_gpu_peak={load_peak:.1f} GB", flush=True)
img, dt = _gen(pipe, steps, seed, res) # warmup
img, dt = _gen(pipe, steps, seed, res)
img.save(OUT / f"{mode}.png")
print(f"[{mode}] gen={dt:.3f}s saved {mode}.png", flush=True)
ref_path = OUT / "runtime.png"
if mode == "prequant" and ref_path.exists():
from PIL import Image
lp = _lpips(np.array(Image.open(ref_path).convert("RGB")), np.array(img))
print(f"[prequant] LPIPS_vs_runtime={lp}", flush=True)
return 0
def main(argv=None) -> int:
p = argparse.ArgumentParser()
p.add_argument("--mode", choices=["prequant", "runtime"], required=True)
p.add_argument("--steps", type=int, default=8)
p.add_argument("--res", type=int, default=1024)
p.add_argument("--seed", type=int, default=42)
args = p.parse_args(argv)
rc = run(args.mode, args.steps, args.seed, args.res)
print("VERIFY-PREQUANT-DONE", flush=True)
return rc
if __name__ == "__main__":
sys.exit(main())

View file

@ -52,10 +52,15 @@ from .diffusion_speed import (
snapshot_backend_flags,
)
from .diffusion_precision import quantize_text_encoders
from .diffusion_prequant import (
load_prequantized_transformer,
resolve_prequant_source,
)
from .diffusion_transformer_quant import (
dense_transformer_supported,
normalize_transformer_quant,
quantize_transformer,
select_transformer_quant_scheme,
)
logger = get_logger(__name__)
@ -277,6 +282,7 @@ class DiffusionBackend:
text_encoder_quant: Optional[str] = None,
transformer_quant: Optional[str] = None,
transformer_quant_fast_accum: Optional[bool] = None,
transformer_prequant_path: Optional[str] = None,
) -> dict[str, Any]:
"""Validate, then run the (slow) load on a daemon thread. Returns at once."""
fam = self.validate_load_request(
@ -310,6 +316,7 @@ class DiffusionBackend:
text_encoder_quant = text_encoder_quant,
transformer_quant = transformer_quant,
transformer_quant_fast_accum = transformer_quant_fast_accum,
transformer_prequant_path = transformer_prequant_path,
_load_token = token,
),
daemon = True,
@ -433,6 +440,7 @@ class DiffusionBackend:
text_encoder_quant: Optional[str] = None,
transformer_quant: Optional[str] = None,
transformer_quant_fast_accum: Optional[bool] = None,
transformer_prequant_path: Optional[str] = None,
_load_token: Optional[int] = None,
) -> dict[str, Any]:
# Validate first (cheap, no torch/diffusers) so a direct call with a bad
@ -500,6 +508,8 @@ class DiffusionBackend:
target,
transformer_quant,
transformer_quant_fast_accum,
fam = fam,
prequant_path = transformer_prequant_path,
)
except Exception as exc: # noqa: BLE001 — fall back to the GGUF build
logger.warning(
@ -606,27 +616,71 @@ class DiffusionBackend:
target: DiffusionDeviceTarget,
mode: Optional[str],
fast_accum: Optional[bool] = None,
*,
fam: Optional[DiffusionFamily] = None,
prequant_path: Optional[str] = None,
) -> 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)``.
"""Build the opt-in fast pipeline and return ``(pipe, engaged_scheme)``.
Two ways to get the quantized transformer, in order:
1. Pre-quantized: if a checkpoint is configured for the chosen scheme (an explicit
``prequant_path`` or the family's hosted repo), load the already-quantized
weights onto the meta device and assign them in -- the dense bf16 never lands on
the GPU, so the load peak is ~half and the download is smaller.
2. Dense + quantise (fallback): load the DENSE bf16 transformer from the base repo,
place it on the device, and torchao-quantise it in place.
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."""
catches it and falls back to the GGUF build. Quantisation runs ON the device and
BEFORE the loader compiles the repeated block, so the order stays quantize ->
compile -> placement."""
# 1. Pre-quantized checkpoint, when one is configured for the resolved scheme.
scheme = select_transformer_quant_scheme(target, mode)
if scheme is not None and fam is not None:
source = resolve_prequant_source(fam, scheme, path_override = prequant_path)
if source is not None:
transformer = load_prequantized_transformer(
transformer_cls,
base,
source,
device = device,
dtype = dtype,
hf_token = hf_token,
scheme = scheme,
logger = logger,
)
if transformer is not None:
pipe = self._assemble_pipe(pipeline_cls, base, transformer, dtype, hf_token, device)
return pipe, scheme
# 2. Fallback: materialise the dense bf16 transformer and quantise it on-device.
transformer = transformer_cls.from_pretrained(
base, subfolder = "transformer", torch_dtype = dtype, token = hf_token
)
pipe = self._assemble_pipe(pipeline_cls, base, transformer, dtype, hf_token, device)
scheme = quantize_transformer(pipe, target, mode = mode, fast_accum = fast_accum, logger = logger)
if scheme is None:
raise RuntimeError("transformer quant unsupported for this device/scheme")
return pipe, scheme
@staticmethod
def _assemble_pipe(
pipeline_cls: Any,
base: str,
transformer: Any,
dtype: Any,
hf_token: Optional[str],
device: str,
) -> Any:
"""Assemble the diffusers pipeline around ``transformer`` and place it on ``device``
(a no-op for an already-placed pre-quantized transformer; it moves the companions)."""
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, fast_accum = fast_accum, logger = logger)
if scheme is None:
raise RuntimeError("transformer quant unsupported for this device/scheme")
return pipe, scheme
return pipe
def _plan_memory(
self,

View file

@ -38,6 +38,12 @@ class DiffusionFamily:
# regional torch.compile. Now consulted on the GGUF path too (compile runs on the
# GGUF transformer); all current families compile, so this stays True.
supports_torch_compile: bool = True
# Optional pre-quantized transformer checkpoints, as (scheme, repo_id) pairs (a
# hashable mapping). When the fast transformer_quant path resolves a scheme with a
# hosted checkpoint, the loader fetches the already-quantized weights instead of
# materialising the dense bf16 transformer on the GPU (much lower load VRAM + a
# smaller download). Empty until checkpoints are hosted -> behaviour is unchanged.
prequant_repos: tuple[tuple[str, str], ...] = field(default_factory = tuple)
# Keyed by architecture, not per model variant: a checkpoint's specific base repo
@ -116,6 +122,14 @@ def resolve_base_repo(fam: DiffusionFamily, base_repo: Optional[str]) -> str:
return base or fam.base_repo
def family_prequant_repo(fam: DiffusionFamily, scheme: str) -> Optional[str]:
"""The hosted pre-quantized transformer repo for ``scheme`` in this family, or None."""
for entry_scheme, repo_id in fam.prequant_repos:
if entry_scheme == scheme:
return repo_id
return None
def resolve_local_gguf_child(repo_root: Path, gguf_filename: str) -> Path:
"""Resolve ``gguf_filename`` to a file under ``repo_root``, rejecting escapes.

View file

@ -0,0 +1,194 @@
# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
"""Load a *pre-quantized* transformer instead of quantising a dense one on the GPU.
The opt-in fast transformer_quant path (see ``diffusion_transformer_quant.py``) loads
the dense bf16 transformer and torchao-``quantize_``s it in place. That materialises the
full bf16 weights on the GPU before quantising, so the load peak is ~2x the GGUF's and it
pulls the full bf16 download. When a transformer has already been quantised once and saved
(``scripts/build_prequant_checkpoint.py``), this module loads those weights directly:
1. build the transformer skeleton on the ``meta`` device (no storage) via
``accelerate.init_empty_weights`` + ``from_config``;
2. ``load_state_dict(assign=True)`` the quantized state dict (the torchao weight subclass
tensors are assigned in, not copied), so the dense bf16 never touches the GPU;
3. move to the device.
Measured (B200, Z-Image fp8): transformer GPU load peak 12.9 -> 6.3 GB, download 12 ->
6.28 GB, output bit-identical (LPIPS 0.0). The checkpoint carries the exact same scheme +
``min_features`` as the runtime path, so the result is identical to quantising on the fly.
Best-effort and lazily imported throughout: a missing / mismatched / unreadable checkpoint
returns None and the caller falls back to the dense-quantise path (and then to GGUF). All
behaviour is gated on a configured source -- with nothing configured this module is inert.
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import Any, Optional
# torch.save dict layout this module reads (and the build script writes). Bumped if the
# on-disk structure changes so an old/foreign artifact is rejected rather than mis-loaded.
PREQUANT_FORMAT = "unsloth_prequant_transformer_state_dict_v1"
@dataclass(frozen = True)
class PrequantSource:
"""Where a pre-quantized transformer checkpoint lives. ``kind`` is "path" (a local
file) or "repo" (a Hub repo id in ``location`` + ``filename`` inside it)."""
kind: str
location: str
filename: Optional[str] = None
def prequant_filename(scheme: str) -> str:
"""The conventional checkpoint filename for ``scheme`` inside a Hub repo."""
return f"transformer_{scheme}.pt"
def resolve_prequant_source(
fam: Any,
scheme: str,
*,
path_override: Optional[str] = None,
) -> Optional[PrequantSource]:
"""Resolve where the pre-quantized checkpoint for ``(fam, scheme)`` should come from.
Priority: (1) an explicit local ``path_override`` (testing / power users); (2) the
family's hosted repo for ``scheme``; (3) None -> no pre-quant, caller quantises dense.
Pure: no IO, no torch -- it only decides the source, the loader fetches it.
"""
override = (path_override or "").strip()
if override:
return PrequantSource(kind = "path", location = override, filename = None)
try:
from .diffusion_families import family_prequant_repo
repo_id = family_prequant_repo(fam, scheme)
except Exception: # noqa: BLE001 — a bad family object must not break the load
repo_id = None
if repo_id:
return PrequantSource(kind = "repo", location = repo_id, filename = prequant_filename(scheme))
return None
def load_prequantized_transformer(
transformer_cls: Any,
base: str,
source: PrequantSource,
*,
device: str,
dtype: Any,
hf_token: Optional[str] = None,
scheme: str,
logger: Any = None,
) -> Optional[Any]:
"""Load the pre-quantized transformer described by ``source`` onto ``device``.
Returns the placed, already-quantized transformer, or None on any problem (missing /
mismatched / unreadable checkpoint, or a meta-init the class does not support) so the
caller falls back to the dense-quantise path. Best-effort: never raises for an
ordinary unavailable artifact.
"""
try:
path = _resolve_checkpoint_path(source, hf_token)
if path is None:
return None
import torch
# torchao weight subclasses are not safetensors-serializable, so the checkpoint is
# a torch.save pickle. weights_only=False is required to rebuild those subclasses;
# only a configured family repo (first-party) or an explicit local path reaches
# here, which is the trust signal -- this never loads an arbitrary remote pickle.
ckpt = torch.load(path, weights_only = False, map_location = "cpu")
if not _validate_checkpoint(ckpt, scheme, base, logger):
return None
state_dict = ckpt["state_dict"]
config = transformer_cls.load_config(base, subfolder = "transformer", token = hf_token)
from accelerate import init_empty_weights
with init_empty_weights():
transformer = transformer_cls.from_config(config)
# assign=True swaps in the loaded (quantized) tensors rather than copying into the
# meta tensors (a copy into meta is a no-op); strict=True since the saved state
# dict is the full state dict of the same class (non-persistent buffers excluded).
transformer.load_state_dict(state_dict, strict = True, assign = True)
if _has_meta_tensors(transformer):
# A class with non-persistent buffers (computed in __init__, absent from the
# state dict) leaves those on meta. Rebuild on CPU so the buffers hold their
# real values, then re-assign the quantized weights. The dense bf16 lives in
# CPU RAM only -- the GPU still receives just the quantized footprint.
transformer = transformer_cls.from_config(config)
transformer.load_state_dict(state_dict, strict = True, assign = True)
transformer = transformer.to(device)
try: # diagnostic marker, mirrors the runtime-quant path
transformer._unsloth_runtime_quant = scheme
except Exception: # noqa: BLE001 — marker is best-effort
pass
if logger is not None:
logger.info(
"diffusion.prequant: loaded %s checkpoint (%s) onto %s",
scheme,
source.kind,
device,
)
return transformer
except Exception as exc: # noqa: BLE001 — fall back to the dense-quantise path
_warn(logger, f"{scheme}:{source.kind}", exc)
return None
def _resolve_checkpoint_path(source: PrequantSource, hf_token: Optional[str]) -> Optional[str]:
"""The local file path for ``source``, downloading from the Hub if needed; None if absent."""
if source.kind == "path":
import os
return source.location if os.path.isfile(source.location) else None
if source.kind == "repo":
from huggingface_hub import hf_hub_download
return hf_hub_download(
repo_id = source.location, filename = source.filename, token = hf_token
)
return None
def _validate_checkpoint(ckpt: Any, scheme: str, base: str, logger: Any) -> bool:
"""Reject a checkpoint that is the wrong format / scheme / base model."""
if not isinstance(ckpt, dict) or ckpt.get("format") != PREQUANT_FORMAT:
_warn(logger, scheme, ValueError("unrecognised pre-quant checkpoint format"))
return False
if "state_dict" not in ckpt:
_warn(logger, scheme, ValueError("pre-quant checkpoint has no state_dict"))
return False
meta = ckpt.get("metadata") or {}
if meta.get("scheme") != scheme:
_warn(logger, scheme, ValueError(f"checkpoint scheme {meta.get('scheme')!r} != {scheme!r}"))
return False
ckpt_base = meta.get("base_model_id")
if ckpt_base and base and ckpt_base != base:
_warn(logger, scheme, ValueError(f"checkpoint base {ckpt_base!r} != {base!r}"))
return False
return True
def _has_meta_tensors(module: Any) -> bool:
"""True if any parameter or buffer is still on the meta device after loading."""
try:
for tensor in list(module.parameters()) + list(module.buffers()):
if getattr(tensor, "is_meta", False):
return True
except Exception: # noqa: BLE001
return False
return False
def _warn(logger: Any, what: str, exc: Exception) -> None:
if logger is not None:
logger.warning("diffusion.prequant: %s failed: %s", what, exc)

View file

@ -1738,6 +1738,14 @@ class DiffusionLoadRequest(BaseModel):
"HBM cards, which are not nerfed). true/false force it. Negligible "
"quality effect (below the fp8 quant noise floor); no overflow risk.",
)
transformer_prequant_path: Optional[str] = Field(
None,
description = "Local path to a pre-quantized transformer checkpoint (built by "
"scripts/build_prequant_checkpoint.py) for the requested transformer_quant "
"scheme. Loads the already-quantized weights with the dense bf16 never on the "
"GPU (~half the load VRAM and a smaller download). null uses the family's hosted "
"checkpoint if configured, else quantises the dense transformer at load time.",
)
class DiffusionGenerateRequest(BaseModel):

View file

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

View file

@ -906,6 +906,10 @@ def _stub_dense_quant(monkeypatch, *, scheme = "fp8"):
monkeypatch.setattr(_FakeTransformer, "from_pretrained", _from_pretrained, raising = False)
monkeypatch.setattr(dmod, "dense_transformer_supported", lambda target: True)
# Resolve the scheme without the real GPU smoke probe, and configure no pre-quant
# checkpoint so the dense materialise+quantise branch is the one exercised.
monkeypatch.setattr(dmod, "select_transformer_quant_scheme", lambda target, mode: scheme)
monkeypatch.setattr(dmod, "resolve_prequant_source", lambda fam, scheme, **kw: None)
def _quantize(pipe, target, *, mode, **kw):
calls["quantize"] += 1
@ -957,6 +961,75 @@ def test_transformer_quant_dense_path_engaged(fake_runtime, tmp_path, monkeypatc
assert status["offload_policy"] == "none"
def test_transformer_quant_prequant_path_engaged(fake_runtime, tmp_path, monkeypatch):
# A configured pre-quant checkpoint -> load the already-quantized transformer directly;
# the dense from_pretrained and the on-device quantize_transformer are NOT used.
from core.inference import diffusion as dmod
backend = DiffusionBackend()
_force_cuda_target(backend, monkeypatch)
monkeypatch.setattr(dmod, "dense_transformer_supported", lambda target: True)
monkeypatch.setattr(dmod, "select_transformer_quant_scheme", lambda target, mode: "fp8")
monkeypatch.setattr(dmod, "resolve_prequant_source", lambda fam, scheme, **kw: object())
prequant_obj = object()
loaded: dict = {"n": 0}
def _load_prequant(transformer_cls, base, source, **kw):
loaded["n"] += 1
loaded["scheme"] = kw.get("scheme")
return prequant_obj
monkeypatch.setattr(dmod, "load_prequantized_transformer", _load_prequant)
@classmethod
def _fp_fail(cls, *a, **k):
pytest.fail("dense from_pretrained must not run when a prequant checkpoint loads")
monkeypatch.setattr(_FakeTransformer, "from_pretrained", _fp_fail, raising = False)
monkeypatch.setattr(
dmod,
"quantize_transformer",
lambda *a, **k: pytest.fail("quantize_transformer must not run on the prequant path"),
)
(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",
transformer_prequant_path = str(tmp_path / "zimage_fp8.pt"),
)
assert status["transformer_quant"] == "fp8"
assert loaded["n"] == 1 and loaded["scheme"] == "fp8"
# The pre-quantized transformer object was assembled into the pipeline...
assert _FakePipeline.last.get("transformer") is prequant_obj
# ...and the GGUF single-file path was not used.
assert _FakeTransformer.last == {}
def test_transformer_quant_prequant_load_fails_falls_back_to_dense(fake_runtime, tmp_path, monkeypatch):
# A configured prequant source whose load returns None must fall back to the dense
# materialise+quantise path (not straight to GGUF), preserving the fast mode.
from core.inference import diffusion as dmod
backend = DiffusionBackend()
_force_cuda_target(backend, monkeypatch)
calls = _stub_dense_quant(monkeypatch, scheme = "fp8")
# Override the no-prequant default: a source resolves, but its load fails.
monkeypatch.setattr(dmod, "resolve_prequant_source", lambda fam, scheme, **kw: object())
monkeypatch.setattr(dmod, "load_prequantized_transformer", lambda *a, **k: 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["transformer_quant"] == "fp8"
assert calls["from_pretrained"] == 1 and calls["quantize"] == 1 # dense path ran
assert _FakeTransformer.last == {} # GGUF not used
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.

View file

@ -0,0 +1,175 @@
# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
"""Hermetic CPU tests for the pre-quantized transformer load path.
torch / accelerate are stubbed via ``sys.modules`` (the module under test imports them
lazily), and ``transformer_cls`` is a fake that records calls -- so the resolver, the
meta-init + ``load_state_dict(assign=True)`` flow, and the validation/fallback behaviour
are all exercised without CUDA, torchao, or a real diffusers model.
"""
from __future__ import annotations
import contextlib
import sys
import types
import core.inference.diffusion_prequant as pq
from core.inference.diffusion_families import DiffusionFamily
from core.inference.diffusion_prequant import (
PREQUANT_FORMAT,
PrequantSource,
load_prequantized_transformer,
resolve_prequant_source,
)
# ── resolve_prequant_source ──────────────────────────────────────────────────────
def _fam(prequant_repos=()):
return DiffusionFamily(
name = "z-image",
pipeline_class = "ZImagePipeline",
transformer_class = "ZImageTransformer2DModel",
base_repo = "Tongyi-MAI/Z-Image-Turbo",
prequant_repos = prequant_repos,
)
def test_resolve_path_override_wins():
fam = _fam(prequant_repos = (("fp8", "org/hosted-fp8"),))
src = resolve_prequant_source(fam, "fp8", path_override = "/tmp/local.pt")
assert src == PrequantSource(kind = "path", location = "/tmp/local.pt", filename = None)
def test_resolve_family_repo_by_scheme():
fam = _fam(prequant_repos = (("fp8", "org/hosted-fp8"), ("int8", "org/hosted-int8")))
src = resolve_prequant_source(fam, "int8")
assert src.kind == "repo" and src.location == "org/hosted-int8"
assert src.filename == "transformer_int8.pt"
def test_resolve_wrong_scheme_is_none():
fam = _fam(prequant_repos = (("fp8", "org/hosted-fp8"),))
assert resolve_prequant_source(fam, "int8") is None
def test_resolve_nothing_configured_is_none():
assert resolve_prequant_source(_fam(), "fp8") is None
assert resolve_prequant_source(_fam(), "fp8", path_override = "") is None
# ── load_prequantized_transformer ────────────────────────────────────────────────
class _FakeTransformer:
calls: dict = {}
def __init__(self):
self.assigned = None
self.moved = None
@classmethod
def load_config(cls, base, **kw):
cls.calls["load_config"] = {"base": base, **kw}
return {"cfg": True}
@classmethod
def from_config(cls, config):
cls.calls["from_config"] = config
return cls()
@classmethod
def from_pretrained(cls, *a, **k): # the dense path -- must never run here
cls.calls["from_pretrained"] = True
raise AssertionError("from_pretrained must not be called on the prequant path")
def load_state_dict(self, sd, strict = True, assign = False):
_FakeTransformer.calls["load_state_dict"] = {"strict": strict, "assign": assign}
self.assigned = sd
def parameters(self):
return []
def buffers(self):
return []
def to(self, device):
self.moved = device
return self
def _stub_torch_accelerate(monkeypatch, ckpt, *, load_raises=False):
torch = types.ModuleType("torch")
def _load(path, weights_only = False, map_location = None):
if load_raises:
raise RuntimeError("corrupt checkpoint")
return ckpt
torch.load = _load
monkeypatch.setitem(sys.modules, "torch", torch)
accelerate = types.ModuleType("accelerate")
accelerate.init_empty_weights = lambda: contextlib.nullcontext()
monkeypatch.setitem(sys.modules, "accelerate", accelerate)
def _good_ckpt(scheme="fp8", base="Tongyi-MAI/Z-Image-Turbo"):
return {
"format": PREQUANT_FORMAT,
"metadata": {"scheme": scheme, "base_model_id": base},
"state_dict": {"weight": object()},
}
def _load(monkeypatch, tmp_path, ckpt, *, scheme="fp8", load_raises=False, exists=True):
_FakeTransformer.calls = {}
_stub_torch_accelerate(monkeypatch, ckpt, load_raises = load_raises)
path = tmp_path / "ckpt.pt"
if exists:
path.write_bytes(b"x")
source = PrequantSource(kind = "path", location = str(path), filename = None)
return load_prequantized_transformer(
_FakeTransformer,
"Tongyi-MAI/Z-Image-Turbo",
source,
device = "cuda",
dtype = "bfloat16",
hf_token = None,
scheme = scheme,
logger = None,
)
def test_load_meta_init_and_assign(monkeypatch, tmp_path):
t = _load(monkeypatch, tmp_path, _good_ckpt())
assert t is not None
# meta-init path was used, not the dense from_pretrained.
assert "from_config" in _FakeTransformer.calls
assert "from_pretrained" not in _FakeTransformer.calls
# assign=True is the whole point (copy into meta is a no-op).
assert _FakeTransformer.calls["load_state_dict"] == {"strict": True, "assign": True}
assert t.moved == "cuda"
assert t._unsloth_runtime_quant == "fp8"
def test_load_missing_file_is_none(monkeypatch, tmp_path):
assert _load(monkeypatch, tmp_path, _good_ckpt(), exists = False) is None
def test_load_torch_load_raises_is_none(monkeypatch, tmp_path):
assert _load(monkeypatch, tmp_path, _good_ckpt(), load_raises = True) is None
def test_load_format_mismatch_is_none(monkeypatch, tmp_path):
bad = _good_ckpt()
bad["format"] = "something_else"
assert _load(monkeypatch, tmp_path, bad) is None
def test_load_scheme_mismatch_is_none(monkeypatch, tmp_path):
# checkpoint built for int8, but fp8 was requested.
assert _load(monkeypatch, tmp_path, _good_ckpt(scheme = "int8"), scheme = "fp8") is None
def test_load_base_mismatch_is_none(monkeypatch, tmp_path):
assert _load(monkeypatch, tmp_path, _good_ckpt(base = "other/model")) is None

View file

@ -343,6 +343,22 @@ def test_transformer_quant_fast_accum_threads_through(client, monkeypatch):
assert backend.last_load_kwargs.get("transformer_quant_fast_accum") is False
def test_transformer_prequant_path_threads_through(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": "fp8",
"transformer_prequant_path": "/data/zimage_fp8.pt",
},
)
assert resp.status_code == 200
assert backend.last_load_kwargs.get("transformer_prequant_path") == "/data/zimage_fp8.pt"
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.