# Conflicts: # studio/backend/core/inference/diffusion.py # studio/frontend/src/features/images/images-page.tsx
58 lines
5.7 KiB
Text
58 lines
5.7 KiB
Text
## Summary
|
|
|
|
An opt-in **fast transformer** mode for the Studio diffusion backend: load the dense bf16 transformer and torchao-quantise it onto the low-precision tensor cores, instead of the GGUF transformer. Stacked on the Phase 7 perf pass (#6690).
|
|
|
|
The motivation, measured on a B200 (Z-Image-Turbo, 1024px / 8 steps, LPIPS vs the dense bf16 reference): GGUF stores the transformer 4-bit but **dequantises to bf16 on every matmul**, so it runs at bf16 tensor-core rate and never touches the int8 / fp8 / fp4 cores. It is a memory win that costs speed. Loading the dense bf16 transformer and quantising it dynamically with torchao runs the matmul on the actual low-precision cores:
|
|
|
|
| config | sec | vs GGUF+compile | LPIPS vs bf16 |
|
|
| --- | --- | --- | --- |
|
|
| GGUF + compile (today's default) | 0.802 | 1.00x | 0.083 |
|
|
| dense bf16 + compile (no quant) | 0.671 | 1.20x | 0.004 |
|
|
| int8 dynamic + compile | 0.603 | **1.33x** | 0.069 |
|
|
| fp8 dynamic + compile | 0.585 | **1.37x** | 0.058 |
|
|
|
|
Every working scheme beats GGUF on **both** speed and quality (LPIPS lands *below* GGUF's own 4-bit floor of 0.083). The only cost is memory: the dense bf16 transformer must be loaded (~2x the 4-bit GGUF), so the mode is strictly opt-in and gated on resident VRAM headroom. GGUF + compile stays the low-memory default and the fallback.
|
|
|
|
## What it does
|
|
|
|
New flag `transformer_quant` on the load request (`auto | int8 | fp8 | nvfp4 | mxfp8`, default off). When set and the device qualifies (CUDA + bf16 + the dense weights fit resident), the loader:
|
|
|
|
1. loads the **dense bf16 transformer** from the family `base_repo` (`from_pretrained(subfolder="transformer")`) instead of the GGUF,
|
|
2. places it on the device and torchao-quantises the FLOP-heavy linears,
|
|
3. compiles the repeated block (existing Phase 7 regional compile), then applies placement — so the order is **quantize -> compile -> offload**.
|
|
|
|
`auto` picks the best scheme for the GPU via a real quantise+matmul **smoke probe** (Blackwell nvfp4/mxfp8 -> fp8 -> int8; Ada/Hopper fp8 -> int8; Ampere int8). An explicit scheme is honored only if supported, never silently swapped. **Any** failure (unsupported arch/scheme, OOM, partial quant, or the dense weights not fitting resident) falls back to the GGUF build with a logged reason — the default path cannot regress.
|
|
|
|
A `min_features=512` filter skips the tiny timestep/pooled/modulation projections: the int8 path uses `torch._int_mm`, which requires activation rows M>16, and those projections run at M=1 and crash it (measured: 239/276 Z-Image linears quantised, full speedup, no crash).
|
|
|
|
## Design
|
|
|
|
- New module `core/inference/diffusion_transformer_quant.py` mirrors `diffusion_precision.py` (the text-encoder quant module): pure functions, torch/torchao imported lazily, best-effort, hermetic CPU tests. It owns scheme selection + the arch/capability/smoke probe + the `quantize_` call.
|
|
- `diffusion.py` `load_pipeline` gains the source-branch (dense+quant or GGUF), the VRAM preflight (reuses `plan_diffusion_memory` / `estimate_gguf_dense_mib`), and the fallback. `transformer_quant` threads through `begin_load` / `_LoadState` / `status()` exactly like `text_encoder_quant`.
|
|
- Flag surface mirrors `text_encoder_quant`: request + status models in `models/inference.py`, forwarded in `routes/inference.py`.
|
|
- torchao tensors are not safetensors-serializable; this backend is inference-only, so the engaged transformer carries a diagnostic runtime marker but there is no save path to guard.
|
|
|
|
Blackwell nvfp4/mxfp8 are wired but currently smoke-fail on this box's torch 2.9 (the FP4/MX kernels need torch>=2.11 / the MX build); `auto` lands on fp8 there with no error, and they will activate automatically once the tooling supports them.
|
|
|
|
## Measured through the backend (single B200, Z-Image-Turbo, 1024x1024, 8 steps, seed 12345)
|
|
|
|
End to end through `DiffusionBackend` via `scripts/diffusion_bench.py`. LPIPS vs the dense bf16 reference is from `scripts/quant_probe.py` (the standalone lever probe).
|
|
|
|
| `transformer_quant` | engaged scheme | median/gen | vs GGUF | load VRAM | gen VRAM | LPIPS |
|
|
| --- | --- | --- | --- | --- | --- | --- |
|
|
| (unset) GGUF + compile | none | 0.823 s | reference | 13.4 GB | 15.3 GB | 0.083 |
|
|
| `auto` | fp8 | **0.614 s** | **1.34x** | 20.9 GB | 16.5 GB | 0.058 |
|
|
| `int8` | int8 | 0.626 s | 1.32x | 20.9 GB | 16.5 GB | 0.069 |
|
|
| `mxfp8` | mxfp8 | 0.651 s | 1.26x | 21.3 GB | 16.7 GB | n/m |
|
|
|
|
`auto` selects fp8 on this B200 (nvfp4 smoke-fails on torch 2.9 -> the ladder prefers the measured-faster fp8 over mxfp8). The speed-for-memory trade is explicit: the dense bf16 load peaks ~21 GB vs GGUF's 13.4 GB; resident generation VRAM is close (16.5 vs 15.3 GB). Every engaged scheme is faster than GGUF *and* lands below its 0.083 LPIPS floor.
|
|
|
|
## Tests
|
|
|
|
CPU-only, hermetic (torch / torchao stubbed via `sys.modules`):
|
|
- new `tests/test_diffusion_transformer_quant.py` — normalisation, the arch-selection ladder (Ampere int8 / Ada-Hopper fp8 / Blackwell nvfp4->mxfp8->fp8 fallback / pre-Ampere none / explicit-unsupported none), the smoke-probe cache + tolerance, the feature filter, and the apply path (calls `quantize_` with a filter_fn, sets the marker, tolerates failure).
|
|
- extended `tests/test_diffusion_backend.py` — default load skips the dense path; the dense path engages and reports the scheme; a quant failure falls back to GGUF; the path is skipped when the plan would offload.
|
|
- extended `tests/test_diffusion_routes.py` — the flag threads through to `begin_load` and an invalid enum is a 422.
|
|
|
|
`scripts/diffusion_bench.py` gains `--transformer-quant` (through-the-backend benchmark + regression guard); `scripts/quant_probe.py` is the standalone torchao lever probe (latency + PSNR + LPIPS + VRAM vs the dense reference, with the `--min-feat` filter).
|
|
|