From b90f833469c853bdceba09cec413fc7fde21cde7 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Fri, 26 Jun 2026 11:23:20 +0000 Subject: [PATCH] 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. --- scripts/build_prequant_checkpoint.py | 127 ++++++++++++ scripts/prequant_probe.py | 170 +++++++++++++++ scripts/verify_prequant_backend.py | 126 ++++++++++++ studio/backend/core/inference/diffusion.py | 74 ++++++- .../core/inference/diffusion_families.py | 14 ++ .../core/inference/diffusion_prequant.py | 194 ++++++++++++++++++ studio/backend/models/inference.py | 8 + studio/backend/routes/inference.py | 1 + .../backend/tests/test_diffusion_backend.py | 73 +++++++ .../backend/tests/test_diffusion_prequant.py | 175 ++++++++++++++++ studio/backend/tests/test_diffusion_routes.py | 16 ++ 11 files changed, 968 insertions(+), 10 deletions(-) create mode 100644 scripts/build_prequant_checkpoint.py create mode 100644 scripts/prequant_probe.py create mode 100644 scripts/verify_prequant_backend.py create mode 100644 studio/backend/core/inference/diffusion_prequant.py create mode 100644 studio/backend/tests/test_diffusion_prequant.py diff --git a/scripts/build_prequant_checkpoint.py b/scripts/build_prequant_checkpoint.py new file mode 100644 index 0000000000..4270bbca1c --- /dev/null +++ b/scripts/build_prequant_checkpoint.py @@ -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()) diff --git a/scripts/prequant_probe.py b/scripts/prequant_probe.py new file mode 100644 index 0000000000..6508bb7bd6 --- /dev/null +++ b/scripts/prequant_probe.py @@ -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) diff --git a/scripts/verify_prequant_backend.py b/scripts/verify_prequant_backend.py new file mode 100644 index 0000000000..dd662fb259 --- /dev/null +++ b/scripts/verify_prequant_backend.py @@ -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()) diff --git a/studio/backend/core/inference/diffusion.py b/studio/backend/core/inference/diffusion.py index 65e7ce8faa..4302e06ec1 100644 --- a/studio/backend/core/inference/diffusion.py +++ b/studio/backend/core/inference/diffusion.py @@ -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, diff --git a/studio/backend/core/inference/diffusion_families.py b/studio/backend/core/inference/diffusion_families.py index 4e0d938178..f9812ffc2c 100644 --- a/studio/backend/core/inference/diffusion_families.py +++ b/studio/backend/core/inference/diffusion_families.py @@ -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. diff --git a/studio/backend/core/inference/diffusion_prequant.py b/studio/backend/core/inference/diffusion_prequant.py new file mode 100644 index 0000000000..f85782bdda --- /dev/null +++ b/studio/backend/core/inference/diffusion_prequant.py @@ -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) diff --git a/studio/backend/models/inference.py b/studio/backend/models/inference.py index d0d62393dc..928df30358 100644 --- a/studio/backend/models/inference.py +++ b/studio/backend/models/inference.py @@ -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): diff --git a/studio/backend/routes/inference.py b/studio/backend/routes/inference.py index ec84b4ee7f..06df9eb154 100644 --- a/studio/backend/routes/inference.py +++ b/studio/backend/routes/inference.py @@ -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: diff --git a/studio/backend/tests/test_diffusion_backend.py b/studio/backend/tests/test_diffusion_backend.py index 634f322639..ef217ecb48 100644 --- a/studio/backend/tests/test_diffusion_backend.py +++ b/studio/backend/tests/test_diffusion_backend.py @@ -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. diff --git a/studio/backend/tests/test_diffusion_prequant.py b/studio/backend/tests/test_diffusion_prequant.py new file mode 100644 index 0000000000..e601979a2e --- /dev/null +++ b/studio/backend/tests/test_diffusion_prequant.py @@ -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 diff --git a/studio/backend/tests/test_diffusion_routes.py b/studio/backend/tests/test_diffusion_routes.py index 9d66edd750..cd41397901 100644 --- a/studio/backend/tests/test_diffusion_routes.py +++ b/studio/backend/tests/test_diffusion_routes.py @@ -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.