# SPDX-License-Identifier: AGPL-3.0-only # Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0 """Speed + memory benchmark for the text-encoder auto-quant and VAE quant casters. The accuracy half is ``quant_accuracy_sweep.py``; this measures memory saved + latency impact. Drives the REAL casters (``quantize_text_encoders`` / ``quantize_vae``) exactly as the loader does, on the components in isolation so the numbers aren't drowned by the DiT (held constant). Three modes: --mode te load ONLY the text_encoder subfolder(s), measure weight memory + encode latency, dense vs ``quantize_text_encoders(mode="auto")``. --mode vae load ONLY the vae subfolder, measure weight memory + decode latency, dense vs ``quantize_vae(mode="auto")``; for flux.2 also explicit fp8_dynamic. --mode e2e load the full pipeline (DiT dense) and measure load-peak / gen-peak memory + generation latency, dense TE+VAE vs auto TE+VAE. torch / torchao / diffusers imported lazily so ``--help`` works without them. Example: CUDA_VISIBLE_DEVICES=4 python scripts/quant_speedmem_bench.py --family qwen-image --mode te \\ --out outputs/quant_speedmem """ from __future__ import annotations import argparse import gc import json import os import sys import time from pathlib import Path from typing import Any, Optional os.environ.setdefault("BITSANDBYTES_NOWELCOME", "1") _REPO_ROOT = Path(__file__).resolve().parent.parent _BACKEND_ROOT = _REPO_ROOT / "studio" / "backend" for _p in (str(_BACKEND_ROOT), str(_REPO_ROOT / "scripts")): if _p not in sys.path: sys.path.insert(0, _p) PROMPT = "A photograph of an astronaut riding a horse on the surface of the moon, detailed, 8k" # base diffusers repos (cached). flux.2 exercises the explicit VAE fp8_dynamic opt-in path (the # only image family where fp8_dynamic is measured in-bar). _FAMILIES: dict[str, dict[str, Any]] = { "qwen-image": {"repo": "Qwen/Qwen-Image"}, "flux.1": {"repo": "black-forest-labs/FLUX.1-dev"}, "sdxl": {"repo": "stabilityai/stable-diffusion-xl-base-1.0"}, "flux.2": {"repo": "black-forest-labs/FLUX.2-dev"}, # video families -- Conv3d VAEs (where the VAE holds real memory). Here for the ISOLATED # te/vae modes only; the dit/e2e modes drive the image __call__ interface and reject them, # pointing at scripts/video_speedmem_bench.py. "ltx-2": {"repo": "Lightricks/LTX-2", "video": True}, "hunyuanvideo-1.5": { "repo": "hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v", "video": True, }, "wan2.2-ti2v-5b": { "repo": "Wan-AI/Wan2.2-TI2V-5B-Diffusers", "vae_force_fp32": True, "video": True, }, } _TE_ATTRS = ("text_encoder", "text_encoder_2", "text_encoder_3") # ── cuda memory / timing helpers (lifted from diffusion_bench.py) ────────────── def _sync() -> None: import torch if torch.cuda.is_available(): torch.cuda.synchronize() def _reset_peak() -> None: import torch if torch.cuda.is_available(): torch.cuda.reset_peak_memory_stats() def _alloc_gb() -> float: import torch return torch.cuda.memory_allocated() / 1e9 if torch.cuda.is_available() else 0.0 def _peak_gb() -> float: import torch return torch.cuda.max_memory_allocated() / 1e9 if torch.cuda.is_available() else 0.0 def _empty() -> None: import torch gc.collect() if torch.cuda.is_available(): torch.cuda.empty_cache() def _median(xs: list[float]) -> float: return sorted(xs)[len(xs) // 2] if xs else 0.0 _LP: dict = {} def _lpips_alex(ref_arr, arr): """LPIPS(AlexNet) between two HxWx3 uint8 images (net kept on CPU). None if lpips missing.""" try: import lpips import torch fn = _LP.get("fn") if fn is None: fn = lpips.LPIPS(net = "alex", verbose = False).eval() _LP["fn"] = fn def _t(a): import torch as _torch return _torch.from_numpy(a).float().permute(2, 0, 1).unsqueeze(0) / 127.5 - 1.0 with torch.no_grad(): return round(float(fn(_t(ref_arr), _t(arr)).item()), 4) except Exception: return None def _timed_generate(pipe, *, steps, res, seed): """One generation returning (PIL image, total seconds, [per-step ms]). Per-step wall-clock via callback_on_step_end (each step synchronised).""" import time as _time import torch g = torch.Generator(device = "cuda").manual_seed(seed) step_ts: list[float] = [] last = [0.0] def _cb(pp, i, t, kw): torch.cuda.synchronize() now = _time.perf_counter() if last[0]: step_ts.append((now - last[0]) * 1000.0) last[0] = now return kw _sync() t0 = _time.perf_counter() img = pipe( prompt = PROMPT, width = res, height = res, num_inference_steps = steps, generator = g, callback_on_step_end = _cb, ).images[0] _sync() return img, (_time.perf_counter() - t0), step_ts def _import_diffusers(): """diffusers with the bnb quantiser disabled (we quant via torchao / layerwise only).""" import torch # noqa: F401 (torch/torchao first so their extensions register) import torchao # noqa: F401 import diffusers.utils.import_utils as iu iu._bitsandbytes_available = False import diffusers return diffusers def _target(): """The object the casters read (.device == "cuda" string, .dtype is torch.bfloat16).""" import types import torch return types.SimpleNamespace(device = "cuda", dtype = torch.bfloat16) # ── model_index component class resolution ──────────────────────────────────── def _model_index(repo: str) -> dict: from huggingface_hub import hf_hub_download path = hf_hub_download(repo, "model_index.json") with open(path, "r", encoding = "utf-8") as fh: return json.load(fh) def _load_named(repo: str, name: str): """Load one pipeline sub-component (text_encoder[_2/_3]) with the exact class the pipeline uses, resolved from model_index.json -> [library, class]. Same modules the real loader casts.""" import importlib import torch idx = _model_index(repo) spec = idx.get(name) if not spec or not isinstance(spec, list) or len(spec) != 2 or spec[0] is None: return None lib, cls_name = spec module = importlib.import_module(lib) klass = getattr(module, cls_name) return klass.from_pretrained(repo, subfolder = name, torch_dtype = torch.bfloat16) def _load_text_encoders(repo: str, device: str): """A bag exposing .text_encoder[_2/_3] (the attrs quantize_text_encoders iterates).""" import types bag = types.SimpleNamespace() n = 0 for attr in _TE_ATTRS: mod = _load_named(repo, attr) if mod is not None: mod = mod.to(device).eval() n += 1 setattr(bag, attr, mod) if n == 0: raise RuntimeError(f"{repo}: no text encoders found in model_index.json") return bag _PROMPT_SUITE = ( "A photograph of an astronaut riding a horse on the moon", "a serene mountain lake at sunrise, mist over the water, ultra detailed", "portrait of an old fisherman with a weathered face, dramatic lighting", "a bustling futuristic city street at night, neon signs, rain reflections", "a bowl of ripe strawberries on a wooden table, macro, soft light", "an oil painting of a stormy sea with a lighthouse", ) def _load_tokenizers(repo: str): """One tokenizer per present text-encoder slot (tokenizer / tokenizer_2 / tokenizer_3).""" from transformers import AutoTokenizer toks = {} for te_attr, tk_sub in ( ("text_encoder", "tokenizer"), ("text_encoder_2", "tokenizer_2"), ("text_encoder_3", "tokenizer_3"), ): try: toks[te_attr] = AutoTokenizer.from_pretrained(repo, subfolder = tk_sub) except Exception: toks[te_attr] = None return toks def _encoder_hidden(te, ids, mask): """Mean-pooled last hidden state (float vector) -- the encoder output that feeds the DiT.""" import torch with torch.inference_mode(): try: out = te(input_ids = ids, attention_mask = mask, output_hidden_states = True) except TypeError: out = te(ids, output_hidden_states = True) hs = getattr(out, "last_hidden_state", None) if hs is None: hidden = getattr(out, "hidden_states", None) hs = hidden[-1] if hidden else (out[0] if isinstance(out, (tuple, list)) else out) m = mask.unsqueeze(-1).to(hs.dtype) v = (hs * m).sum(1) / m.sum(1).clamp(min = 1) return v.float().flatten() def _te_hidden_refs(bag, toks, device): """{te_attr: [pooled hidden vector per prompt]} across the present encoders.""" import torch refs: dict[str, list] = {} for attr in _TE_ATTRS: te = getattr(bag, attr, None) tok = toks.get(attr) if te is None or tok is None: continue vecs = [] for p in _PROMPT_SUITE: enc = tok(p, return_tensors = "pt", padding = "max_length", truncation = True, max_length = 64) ids = enc["input_ids"].to(device) mask = enc.get("attention_mask") mask = mask.to(device) if mask is not None else torch.ones_like(ids) vecs.append(_encoder_hidden(te, ids, mask)) refs[attr] = vecs return refs def _te_param_sig(te) -> tuple: """Weight-storage fingerprint (per-param class + dtype). Changes iff a caster rewrote this encoder's weights. quantize_text_encoders returns ONE scheme for the whole bag even when a per-encoder cast failed, so this is how the bench tells which encoders really engaged.""" return tuple((type(p).__name__, str(p.dtype)) for p in te.parameters()) def measure_te_accuracy( family: str, *, schemes = ("fp8", "fp8_dynamic"), logger = None, ) -> list[dict]: """Hidden-state cosine / relL2 of each TE scheme vs the dense bf16 encoder (per encoder), on real prompts. Bar (PR#150): cosine >= 0.99 and min_cosine >= 0.98.""" import torch import torch.nn.functional as F from core.inference.diffusion_precision import quantize_text_encoders repo = _FAMILIES[family]["repo"] device = "cuda" toks = _load_tokenizers(repo) _empty() bag = _load_text_encoders(repo, device) refs = _te_hidden_refs(bag, toks, device) del bag _empty() rows: list[dict] = [] for scheme in schemes: bag = _load_text_encoders(repo, device) dense_sigs = { attr: _te_param_sig(getattr(bag, attr)) for attr in _TE_ATTRS if getattr(bag, attr, None) is not None } engaged = quantize_text_encoders(bag, _target(), mode = scheme, family = family, logger = logger) if engaged is None: # The scheme was skipped, so the encoder is still dense. Comparing it against the # dense reference would score ~1.0 and falsely certify a no-op; record NOT engaged. rows.append( { "family": family, "encoder": "*", "scheme": scheme, "cosine": None, "min_cosine": None, "relL2": None, "pass": None, "engaged": False, } ) del bag _empty() continue cur = _te_hidden_refs(bag, toks, device) for attr, ref_vecs in refs.items(): # The caster is best-effort per encoder, yet the returned scheme covers the whole # bag. Record any encoder whose weight storage didn't change as NOT engaged, else a # still-dense encoder scores ~1.0 and falsely certifies the scheme. if _te_param_sig(getattr(bag, attr)) == dense_sigs.get(attr): rows.append( { "family": family, "encoder": attr, "scheme": engaged, "cosine": None, "min_cosine": None, "relL2": None, "pass": None, "engaged": False, } ) continue q_vecs = cur.get(attr, []) if not q_vecs: continue cosines = [F.cosine_similarity(r, q, dim = 0).item() for r, q in zip(ref_vecs, q_vecs)] rell2 = [ ((q - r).norm() / r.norm().clamp(min = 1e-8)).item() for r, q in zip(ref_vecs, q_vecs) ] mean_cos = sum(cosines) / len(cosines) min_cos = min(cosines) rows.append( { "family": family, "encoder": attr, "scheme": engaged, "cosine": round(mean_cos, 5), "min_cosine": round(min_cos, 5), "relL2": round(sum(rell2) / len(rell2), 4), "pass": bool(mean_cos >= 0.99 and min_cos >= 0.98), } ) del bag _empty() return rows def _encode_once(bag) -> None: """One forward through every present text encoder on a fixed short token batch. Uses a length within each encoder's max positions (CLIP caps at 77) so position embeddings never overflow.""" import torch with torch.inference_mode(): for attr in _TE_ATTRS: te = getattr(bag, attr, None) if te is None: continue cfg = getattr(te, "config", None) vocab = int(getattr(cfg, "vocab_size", 30000) or 30000) maxpos = int(getattr(cfg, "max_position_embeddings", 64) or 64) length = max(8, min(64, maxpos)) ids = torch.randint(1, min(vocab, 30000), (1, length), device = "cuda") mask = torch.ones_like(ids) try: te(input_ids = ids, attention_mask = mask) except TypeError: te(ids) # ── VAE loading + latent shape (reused from quant_accuracy_sweep) ────────────── def _load_vae( repo: str, device: str, *, force_fp32: bool = False, ): import torch diffusers = _import_diffusers() # vae_force_fp32 families (Wan) store AND run the VAE in fp32, and quantize_vae stays dense # for them; loading bf16 here would measure a truncated VAE production never runs. dtype = torch.float32 if force_fp32 else torch.bfloat16 vae = diffusers.AutoModel.from_pretrained(repo, subfolder = "vae", torch_dtype = dtype) return vae.to(device).eval() def _latent_spec(vae) -> tuple[int, bool]: from torch import nn conv = None dec = getattr(vae, "decoder", None) for mod in (dec, vae): if mod is None: continue for m in mod.modules(): if isinstance(m, (nn.Conv2d, nn.Conv3d)): conv = m break if conv is not None: break is_3d = isinstance(conv, nn.Conv3d) channels = None for key in ("latent_channels", "z_dim", "in_channels"): v = getattr(getattr(vae, "config", object()), key, None) if isinstance(v, int): channels = v break if conv is not None: channels = conv.in_channels return int(channels), bool(is_3d) def _decode_once(vae, z) -> None: import torch with torch.inference_mode(): try: out = vae.decode(z) except TypeError: out = vae.decode(z, return_dict = True) _ = out.sample if hasattr(out, "sample") else out[0] def _make_latent(vae, device: str): """A fixed modest latent at the VAE's native shape (spatial 64 -> ~512px; 3D uses a few frames).""" import torch channels, is_3d = _latent_spec(vae) g = torch.Generator().manual_seed(1234) shape = (1, channels, 3, 32, 32) if is_3d else (1, channels, 64, 64) z = torch.randn(shape, generator = g, dtype = torch.float32) # Match the VAE's parameter dtype (fp32 for vae_force_fp32 families, bf16 otherwise); # an fp32 conv rejects a bf16 latent. dtype = next(vae.parameters()).dtype return z.to(device = device, dtype = dtype) # ── measurement primitives ──────────────────────────────────────────────────── def _time_median(fn, *, warmup: int, iters: int) -> float: for _ in range(warmup): fn() _sync() dts = [] for _ in range(iters): _sync() t0 = time.perf_counter() fn() _sync() dts.append((time.perf_counter() - t0) * 1000.0) # ms return _median(dts) # ── mode: te ────────────────────────────────────────────────────────────────── def measure_te( family: str, *, warmup: int, iters: int, scheme: str = "auto", logger = None, ) -> list[dict]: from core.inference.diffusion_precision import quantize_text_encoders repo = _FAMILIES[family]["repo"] device = "cuda" rows: list[dict] = [] # dense _empty() _reset_peak() bag = _load_text_encoders(repo, device) _sync() mem_dense = _alloc_gb() _reset_peak() lat_dense = _time_median(lambda: _encode_once(bag), warmup = warmup, iters = iters) peak_dense = _peak_gb() # quant in place (scheme="auto" resolves the ladder; else force an explicit scheme to compare) engaged = quantize_text_encoders(bag, _target(), mode = scheme, family = family, logger = logger) _empty() _sync() mem_quant = _alloc_gb() _reset_peak() lat_quant = _time_median(lambda: _encode_once(bag), warmup = warmup, iters = iters) peak_quant = _peak_gb() del bag _empty() rows.append( { "family": family, "component": "text_encoder", "scheme": engaged or "dense(none-engaged)", "mem_dense_gb": round(mem_dense, 3), "mem_quant_gb": round(mem_quant, 3), "saved_gb": round(mem_dense - mem_quant, 3), "mem_ratio": round(mem_quant / mem_dense, 3) if mem_dense else None, "peak_dense_gb": round(peak_dense, 3), "peak_quant_gb": round(peak_quant, 3), "lat_dense_ms": round(lat_dense, 2), "lat_quant_ms": round(lat_quant, 2), "lat_delta_pct": round((lat_quant - lat_dense) / lat_dense * 100.0, 1) if lat_dense else None, } ) return rows # ── mode: vae ─────────────────────────────────────────────────────────────── def _measure_vae_scheme( family: str, repo: str, mode: str, *, warmup: int, iters: int, logger = None, ) -> dict: import types from core.inference.diffusion_vae_quant import quantize_vae device = "cuda" force_fp32 = bool(_FAMILIES.get(family, {}).get("vae_force_fp32", False)) _empty() _reset_peak() vae = _load_vae(repo, device, force_fp32 = force_fp32) z = _make_latent(vae, device) _sync() mem_dense = _alloc_gb() _reset_peak() lat_dense = _time_median(lambda: _decode_once(vae, z), warmup = warmup, iters = iters) peak_dense = _peak_gb() # quantize_vae reads pipe.vae, so hand it a bag exposing .vae (it mutates that module in place). # Pass the family's force_fp32 (Wan) so the real dense-only behaviour is reflected. engaged = quantize_vae( types.SimpleNamespace(vae = vae), _target(), mode = mode, family = family, force_fp32 = force_fp32, logger = logger, ) _empty() _sync() mem_quant = _alloc_gb() _reset_peak() lat_quant = _time_median(lambda: _decode_once(vae, z), warmup = warmup, iters = iters) peak_quant = _peak_gb() del vae, z _empty() return { "family": family, "component": "vae", "requested": mode, "scheme": engaged or "dense(none-engaged)", "mem_dense_gb": round(mem_dense, 3), "mem_quant_gb": round(mem_quant, 3), "saved_gb": round(mem_dense - mem_quant, 3), "mem_ratio": round(mem_quant / mem_dense, 3) if mem_dense else None, "peak_dense_gb": round(peak_dense, 3), "peak_quant_gb": round(peak_quant, 3), "lat_dense_ms": round(lat_dense, 2), "lat_quant_ms": round(lat_quant, 2), "lat_delta_pct": round((lat_quant - lat_dense) / lat_dense * 100.0, 1) if lat_dense else None, } def measure_vae( family: str, *, warmup: int, iters: int, logger = None, ) -> list[dict]: repo = _FAMILIES[family]["repo"] rows = [_measure_vae_scheme(family, repo, "auto", warmup = warmup, iters = iters, logger = logger)] if family == "flux.2": # the one image family where explicit fp8_dynamic conv is measured in-bar. rows.append( _measure_vae_scheme( family, repo, "fp8_dynamic", warmup = warmup, iters = iters, logger = logger ) ) return rows # ── mode: e2e (qwen-image) ──────────────────────────────────────────────────── def _e2e_run( repo: str, *, quant: bool, steps: int, res: int, seed: int, iters: int, family: str, logger = None, ): import torch from core.inference.diffusion_precision import quantize_text_encoders from core.inference.diffusion_vae_quant import quantize_vae diffusers = _import_diffusers() _empty() _reset_peak() # Mirror the loader order: quantize the CPU-built pipeline BEFORE placement, so the auto # row loads like production and load_peak_gb records the quantized config, not the dense one. pipe = diffusers.DiffusionPipeline.from_pretrained(repo, torch_dtype = torch.bfloat16) te_scheme = vae_scheme = None if quant: te_scheme = quantize_text_encoders( pipe, _target(), mode = "auto", family = family, logger = logger ) vae_scheme = quantize_vae(pipe, _target(), mode = "auto", family = family, logger = logger) pipe = pipe.to("cuda") _sync() load_peak = _peak_gb() if quant: _empty() weights_gb = _alloc_gb() def _gen(): g = torch.Generator(device = "cuda").manual_seed(seed) step_ts: list[float] = [] last = [0.0] def _cb(pp, i, t, kw): torch.cuda.synchronize() now = time.perf_counter() if last[0]: step_ts.append((now - last[0]) * 1000.0) last[0] = now return kw _sync() t0 = time.perf_counter() img = pipe( prompt = PROMPT, width = res, height = res, num_inference_steps = steps, generator = g, callback_on_step_end = _cb, ).images[0] _sync() return img, (time.perf_counter() - t0), step_ts img, _, _ = _gen() # warmup _reset_peak() dts, steps_ms, last_img = [], [], img for _ in range(iters): last_img, dt, st = _gen() dts.append(dt) steps_ms.append(_median(st) if st else 0.0) gen_peak = _peak_gb() del pipe _empty() return { "variant": "auto" if quant else "dense", # te_scheme / vae_scheme are the ACTUAL engaged modes: quantize_* returns None when the # component stays dense, so report "dense" then, NOT the requested "auto", else a no-op is # mislabelled as an auto-quantised run. "te_scheme": te_scheme or "dense", "vae_scheme": vae_scheme or "dense", "load_peak_gb": round(load_peak, 2), "weights_gb": round(weights_gb, 2), "gen_peak_gb": round(gen_peak, 2), "gen_latency_s": round(_median(dts), 3), "per_step_ms": round(_median(steps_ms), 1), }, last_img def measure_e2e( family: str, *, steps: int, res: int, seed: int, iters: int, out: Path, variant: str = "both", logger = None, ) -> list[dict]: repo = _FAMILIES[family]["repo"] # "both" runs dense then auto in one process (fast, but the 2nd run is on a hotter GPU). # "dense"/"auto" run a single variant per process to separate a real effect from thermal drift. variants = {"both": (False, True), "dense": (False,), "auto": (True,)}[variant] rows = [] for quant in variants: row, img = _e2e_run( repo, quant = quant, steps = steps, res = res, seed = seed, iters = iters, family = family, logger = logger, ) try: img.save(out / f"e2e_{family}_{row['variant']}.png") except Exception: pass rows.append(row) print(f" e2e {row['variant']:5s}: {json.dumps(row)}", flush = True) return rows # ── mode: dit (transformer quant, the real speed lever) ─────────────────────── def _compile_blocks(transformer) -> bool: """Regional block compile (the real feature path); both dense and quant variants are compiled (torchao dynamic quant is ~30x slower eager) for a fair speedup comparison.""" fn = getattr(transformer, "compile_repeated_blocks", None) if not callable(fn): return False for kw in ({"dynamic": True}, {}): try: fn(**kw) return True except Exception: continue return False def _dit_run( repo: str, family: str, *, dit_quant: str, steps: int, res: int, seed: int, iters: int, compile_blocks: bool = True, logger = None, ): """Load the full pipeline dense, quantise ONLY the transformer (TE + VAE stay dense), regional-compile it, then measure per-step + total latency and peak memory. Returns row + image.""" import torch from core.inference.diffusion_transformer_quant import quantize_transformer diffusers = _import_diffusers() _empty() _reset_peak() pipe = diffusers.DiffusionPipeline.from_pretrained(repo, torch_dtype = torch.bfloat16).to("cuda") load_peak = _peak_gb() engaged = None if dit_quant and dit_quant != "none": engaged = quantize_transformer( pipe, _target(), mode = dit_quant, family = family, logger = logger ) _empty() weights_gb = _alloc_gb() compiled = _compile_blocks(getattr(pipe, "transformer", None)) if compile_blocks else False img, _, _ = _timed_generate(pipe, steps = steps, res = res, seed = seed) # warmup (triggers compile) _reset_peak() dts, steps_ms, last_img = [], [], img for _ in range(iters): last_img, dt, st = _timed_generate(pipe, steps = steps, res = res, seed = seed) dts.append(dt) steps_ms.append(_median(st) if st else 0.0) gen_peak = _peak_gb() del pipe _empty() return { "family": family, "dit_quant": dit_quant, "dit_scheme": engaged or "dense", "compiled": compiled, "load_peak_gb": round(load_peak, 2), "weights_gb": round(weights_gb, 2), "gen_peak_gb": round(gen_peak, 2), "gen_latency_s": round(_median(dts), 3), "per_step_ms": round(_median(steps_ms), 1), }, last_img def measure_dit( family: str, *, schemes, steps: int, res: int, seed: int, iters: int, out: Path, logger = None, ): """Dense reference + each DiT scheme (auto/fp8/int8/mxfp8), reporting speedup, peak-memory drop, and LPIPS(AlexNet) vs the dense render.""" import numpy as np repo = _FAMILIES[family]["repo"] dense_row, dense_img = _dit_run( repo, family, dit_quant = "none", steps = steps, res = res, seed = seed, iters = iters, logger = logger ) try: dense_img.save(out / f"dit_{family}_dense.png") except Exception: pass ref_arr = np.array(dense_img) dense_row["lpips_vs_dense"] = 0.0 dense_row["speedup_vs_dense"] = 1.0 rows = [dense_row] print(f" dit dense: {json.dumps(dense_row)}", flush = True) base_lat = dense_row["gen_latency_s"] or 1.0 for scheme in schemes: row, img = _dit_run( repo, family, dit_quant = scheme, steps = steps, res = res, seed = seed, iters = iters, logger = logger, ) row["lpips_vs_dense"] = _lpips_alex(ref_arr, np.array(img)) row["speedup_vs_dense"] = ( round(base_lat / row["gen_latency_s"], 3) if row["gen_latency_s"] else None ) try: img.save(out / f"dit_{family}_{scheme}.png") except Exception: pass rows.append(row) print(f" dit {scheme:5s}: {json.dumps(row)}", flush = True) return rows # ── main ────────────────────────────────────────────────────────────────────── def main(argv = None) -> int: ap = argparse.ArgumentParser(description = __doc__) ap.add_argument("--family", required = True, choices = sorted(_FAMILIES)) ap.add_argument("--mode", required = True, choices = ("te", "vae", "e2e", "teacc", "dit")) ap.add_argument( "--dit-schemes", default = "auto", help = "dit mode: comma list e.g. auto,fp8,int8,mxfp8" ) ap.add_argument("--warmup", type = int, default = 2) ap.add_argument("--iters", type = int, default = 5) ap.add_argument("--steps", type = int, default = 20, help = "e2e denoise steps") ap.add_argument("--res", type = int, default = 1024, help = "e2e image size") ap.add_argument("--seed", type = int, default = 42) ap.add_argument("--e2e-iters", type = int, default = 3) ap.add_argument( "--variant", choices = ("both", "dense", "auto"), default = "both", help = "e2e variant(s)" ) ap.add_argument("--te-scheme", default = "auto", help = "te mode: auto | fp8_dynamic | fp8 | int8") ap.add_argument("--out", default = "outputs/quant_speedmem") args = ap.parse_args(argv) import logging logging.basicConfig(level = logging.INFO, format = "%(message)s") logger = logging.getLogger("speedmem") out = Path(args.out) out.mkdir(parents = True, exist_ok = True) print(f"== speed+mem bench: family={args.family} mode={args.mode} ==", flush = True) if args.mode in ("dit", "e2e") and _FAMILIES.get(args.family, {}).get("video"): ap.error( f"--mode {args.mode} drives the image pipeline interface (.images[0]); " f"'{args.family}' is a video family -- measure its DiT with " "scripts/video_speedmem_bench.py (frame kwargs, video call path). " "The isolated te/vae modes still support video families here." ) if args.mode == "dit": schemes = [s.strip() for s in args.dit_schemes.split(",") if s.strip()] rows = measure_dit( args.family, schemes = schemes, steps = args.steps, res = args.res, seed = args.seed, iters = args.e2e_iters, out = out, logger = logger, ) elif args.mode == "teacc": rows = measure_te_accuracy(args.family, logger = logger) elif args.mode == "te": rows = measure_te( args.family, warmup = args.warmup, iters = args.iters, scheme = args.te_scheme, logger = logger ) elif args.mode == "vae": rows = measure_vae(args.family, warmup = args.warmup, iters = args.iters, logger = logger) else: rows = measure_e2e( args.family, steps = args.steps, res = args.res, seed = args.seed, iters = args.e2e_iters, out = out, variant = args.variant, logger = logger, ) for r in rows: print(" " + json.dumps(r), flush = True) suffix = "" if args.mode == "e2e" and args.variant != "both": suffix = f"_{args.variant}" elif args.mode == "te" and args.te_scheme != "auto": suffix = f"_{args.te_scheme}" dest = out / f"{args.mode}_{args.family}{suffix}.json" with open(dest, "w", encoding = "utf-8") as fh: json.dump(rows, fh, indent = 2) print(f"wrote {dest}", flush = True) return 0 if __name__ == "__main__": raise SystemExit(main())