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:
parent
3a21f12500
commit
b90f833469
11 changed files with 968 additions and 10 deletions
127
scripts/build_prequant_checkpoint.py
Normal file
127
scripts/build_prequant_checkpoint.py
Normal 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
170
scripts/prequant_probe.py
Normal 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)
|
||||
126
scripts/verify_prequant_backend.py
Normal file
126
scripts/verify_prequant_backend.py
Normal 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())
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
194
studio/backend/core/inference/diffusion_prequant.py
Normal file
194
studio/backend/core/inference/diffusion_prequant.py
Normal 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)
|
||||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
175
studio/backend/tests/test_diffusion_prequant.py
Normal file
175
studio/backend/tests/test_diffusion_prequant.py
Normal 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
|
||||
|
|
@ -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.
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue