unsloth/scripts/prequant_probe.py
Daniel Han b90f833469 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.
2026-06-26 11:23:20 +00:00

170 lines
6.6 KiB
Python

# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
"""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)