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.
126 lines
5.1 KiB
Python
126 lines
5.1 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
|
|
|
|
"""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())
|