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.
170 lines
6.6 KiB
Python
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)
|