[pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
This commit is contained in:
parent
a5928064a0
commit
390bfae9e2
17 changed files with 1172 additions and 609 deletions
|
|
@ -69,33 +69,28 @@ _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()
|
||||
|
|
@ -116,12 +111,11 @@ def _lpips_alex(ref_arr, arr):
|
|||
|
||||
fn = _LP.get("fn")
|
||||
if fn is None:
|
||||
fn = lpips.LPIPS(net="alex", verbose=False).eval()
|
||||
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():
|
||||
|
|
@ -137,7 +131,7 @@ def _timed_generate(pipe, *, steps, res, seed):
|
|||
|
||||
import torch
|
||||
|
||||
g = torch.Generator(device="cuda").manual_seed(seed)
|
||||
g = torch.Generator(device = "cuda").manual_seed(seed)
|
||||
step_ts: list[float] = []
|
||||
last = [0.0]
|
||||
|
||||
|
|
@ -152,8 +146,12 @@ def _timed_generate(pipe, *, steps, res, seed):
|
|||
_sync()
|
||||
t0 = _time.perf_counter()
|
||||
img = pipe(
|
||||
prompt=PROMPT, width=res, height=res, num_inference_steps=steps, generator=g,
|
||||
callback_on_step_end=_cb,
|
||||
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
|
||||
|
|
@ -177,15 +175,14 @@ def _target():
|
|||
|
||||
import torch
|
||||
|
||||
return types.SimpleNamespace(device="cuda", dtype=torch.bfloat16)
|
||||
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:
|
||||
with open(path, "r", encoding = "utf-8") as fh:
|
||||
return json.load(fh)
|
||||
|
||||
|
||||
|
|
@ -203,7 +200,7 @@ def _load_named(repo: str, name: str):
|
|||
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)
|
||||
return klass.from_pretrained(repo, subfolder = name, torch_dtype = torch.bfloat16)
|
||||
|
||||
|
||||
def _load_text_encoders(repo: str, device: str):
|
||||
|
|
@ -244,7 +241,7 @@ def _load_tokenizers(repo: str):
|
|||
("text_encoder_3", "tokenizer_3"),
|
||||
):
|
||||
try:
|
||||
toks[te_attr] = AutoTokenizer.from_pretrained(repo, subfolder=tk_sub)
|
||||
toks[te_attr] = AutoTokenizer.from_pretrained(repo, subfolder = tk_sub)
|
||||
except Exception:
|
||||
toks[te_attr] = None
|
||||
return toks
|
||||
|
|
@ -256,15 +253,15 @@ def _encoder_hidden(te, ids, mask):
|
|||
|
||||
with torch.inference_mode():
|
||||
try:
|
||||
out = te(input_ids=ids, attention_mask=mask, output_hidden_states=True)
|
||||
out = te(input_ids = ids, attention_mask = mask, output_hidden_states = True)
|
||||
except TypeError:
|
||||
out = te(ids, output_hidden_states=True)
|
||||
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)
|
||||
v = (hs * m).sum(1) / m.sum(1).clamp(min = 1)
|
||||
return v.float().flatten()
|
||||
|
||||
|
||||
|
|
@ -280,7 +277,7 @@ def _te_hidden_refs(bag, toks, device):
|
|||
continue
|
||||
vecs = []
|
||||
for p in _PROMPT_SUITE:
|
||||
enc = tok(p, return_tensors="pt", padding="max_length", truncation=True, max_length=64)
|
||||
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)
|
||||
|
|
@ -289,7 +286,12 @@ def _te_hidden_refs(bag, toks, device):
|
|||
return refs
|
||||
|
||||
|
||||
def measure_te_accuracy(family: str, *, schemes=("fp8", "fp8_dynamic"), logger=None) -> list[dict]:
|
||||
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
|
||||
|
|
@ -310,14 +312,16 @@ def measure_te_accuracy(family: str, *, schemes=("fp8", "fp8_dynamic"), logger=N
|
|||
rows: list[dict] = []
|
||||
for scheme in schemes:
|
||||
bag = _load_text_encoders(repo, device)
|
||||
engaged = quantize_text_encoders(bag, _target(), mode=scheme, family=family, logger=logger)
|
||||
engaged = quantize_text_encoders(bag, _target(), mode = scheme, family = family, logger = logger)
|
||||
cur = _te_hidden_refs(bag, toks, device)
|
||||
for attr, ref_vecs in refs.items():
|
||||
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)]
|
||||
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(
|
||||
|
|
@ -340,7 +344,6 @@ 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)
|
||||
|
|
@ -350,10 +353,10 @@ def _encode_once(bag) -> 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")
|
||||
ids = torch.randint(1, min(vocab, 30000), (1, length), device = "cuda")
|
||||
mask = torch.ones_like(ids)
|
||||
try:
|
||||
te(input_ids=ids, attention_mask=mask)
|
||||
te(input_ids = ids, attention_mask = mask)
|
||||
except TypeError:
|
||||
te(ids)
|
||||
|
||||
|
|
@ -363,7 +366,7 @@ def _load_vae(repo: str, device: str):
|
|||
import torch
|
||||
|
||||
diffusers = _import_diffusers()
|
||||
vae = diffusers.AutoModel.from_pretrained(repo, subfolder="vae", torch_dtype=torch.bfloat16)
|
||||
vae = diffusers.AutoModel.from_pretrained(repo, subfolder = "vae", torch_dtype = torch.bfloat16)
|
||||
return vae.to(device).eval()
|
||||
|
||||
|
||||
|
|
@ -395,12 +398,11 @@ def _latent_spec(vae) -> tuple[int, bool]:
|
|||
|
||||
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 = vae.decode(z, return_dict = True)
|
||||
_ = out.sample if hasattr(out, "sample") else out[0]
|
||||
|
||||
|
||||
|
|
@ -411,8 +413,8 @@ def _make_latent(vae, device: str):
|
|||
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)
|
||||
return z.to(device=device, dtype=torch.bfloat16)
|
||||
z = torch.randn(shape, generator = g, dtype = torch.float32)
|
||||
return z.to(device = device, dtype = torch.bfloat16)
|
||||
|
||||
|
||||
# ── measurement primitives ────────────────────────────────────────────────────
|
||||
|
|
@ -431,7 +433,14 @@ def _time_median(fn, *, warmup: int, iters: int) -> float:
|
|||
|
||||
|
||||
# ── mode: te ──────────────────────────────────────────────────────────────────
|
||||
def measure_te(family: str, *, warmup: int, iters: int, scheme: str = "auto", logger=None) -> list[dict]:
|
||||
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"]
|
||||
|
|
@ -439,20 +448,22 @@ def measure_te(family: str, *, warmup: int, iters: int, scheme: str = "auto", lo
|
|||
rows: list[dict] = []
|
||||
|
||||
# dense
|
||||
_empty(); _reset_peak()
|
||||
_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)
|
||||
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()
|
||||
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)
|
||||
lat_quant = _time_median(lambda: _encode_once(bag), warmup = warmup, iters = iters)
|
||||
peak_quant = _peak_gb()
|
||||
|
||||
del bag
|
||||
|
|
@ -470,38 +481,55 @@ def measure_te(family: str, *, warmup: int, iters: int, scheme: str = "auto", lo
|
|||
"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,
|
||||
"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:
|
||||
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"
|
||||
_empty(); _reset_peak()
|
||||
_empty()
|
||||
_reset_peak()
|
||||
vae = _load_vae(repo, device)
|
||||
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)
|
||||
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.
|
||||
force_fp32 = bool(_FAMILIES.get(family, {}).get("vae_force_fp32", False))
|
||||
engaged = quantize_vae(
|
||||
types.SimpleNamespace(vae=vae), _target(), mode=mode, family=family, force_fp32=force_fp32, logger=logger
|
||||
types.SimpleNamespace(vae = vae),
|
||||
_target(),
|
||||
mode = mode,
|
||||
family = family,
|
||||
force_fp32 = force_fp32,
|
||||
logger = logger,
|
||||
)
|
||||
_empty(); _sync()
|
||||
_empty()
|
||||
_sync()
|
||||
mem_quant = _alloc_gb()
|
||||
_reset_peak()
|
||||
lat_quant = _time_median(lambda: _decode_once(vae, z), warmup=warmup, iters=iters)
|
||||
lat_quant = _time_median(lambda: _decode_once(vae, z), warmup = warmup, iters = iters)
|
||||
peak_quant = _peak_gb()
|
||||
|
||||
del vae, z
|
||||
|
|
@ -519,40 +547,65 @@ def _measure_vae_scheme(family: str, repo: str, mode: str, *, warmup: int, iters
|
|||
"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,
|
||||
"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]:
|
||||
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)]
|
||||
rows = [_measure_vae_scheme(family, repo, "auto", warmup = warmup, iters = iters, logger = logger)]
|
||||
if family == "flux.2":
|
||||
# the one image family where the explicit fp8_dynamic conv opt-in is measured in-bar.
|
||||
rows.append(_measure_vae_scheme(family, repo, "fp8_dynamic", warmup=warmup, iters=iters, logger=logger))
|
||||
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):
|
||||
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()
|
||||
pipe = diffusers.DiffusionPipeline.from_pretrained(repo, torch_dtype=torch.bfloat16)
|
||||
_empty()
|
||||
_reset_peak()
|
||||
pipe = diffusers.DiffusionPipeline.from_pretrained(repo, torch_dtype = torch.bfloat16)
|
||||
pipe = pipe.to("cuda")
|
||||
load_peak = _peak_gb()
|
||||
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)
|
||||
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)
|
||||
_empty()
|
||||
weights_gb = _alloc_gb()
|
||||
|
||||
def _gen():
|
||||
g = torch.Generator(device="cuda").manual_seed(seed)
|
||||
g = torch.Generator(device = "cuda").manual_seed(seed)
|
||||
step_ts: list[float] = []
|
||||
last = [0.0]
|
||||
|
||||
|
|
@ -567,12 +620,12 @@ def _e2e_run(repo: str, *, quant: bool, steps: int, res: int, seed: int, iters:
|
|||
_sync()
|
||||
t0 = time.perf_counter()
|
||||
img = pipe(
|
||||
prompt=PROMPT,
|
||||
width=res,
|
||||
height=res,
|
||||
num_inference_steps=steps,
|
||||
generator=g,
|
||||
callback_on_step_end=_cb,
|
||||
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
|
||||
|
|
@ -600,7 +653,15 @@ def _e2e_run(repo: str, *, quant: bool, steps: int, res: int, seed: int, iters:
|
|||
|
||||
|
||||
def measure_e2e(
|
||||
family: str, *, steps: int, res: int, seed: int, iters: int, out: Path, variant: str = "both", logger=None
|
||||
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 / a
|
||||
|
|
@ -610,14 +671,21 @@ def measure_e2e(
|
|||
rows = []
|
||||
for quant in variants:
|
||||
row, img = _e2e_run(
|
||||
repo, quant=quant, steps=steps, res=res, seed=seed, iters=iters, family=family, logger=logger
|
||||
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)
|
||||
print(f" e2e {row['variant']:5s}: {json.dumps(row)}", flush = True)
|
||||
return rows
|
||||
|
||||
|
||||
|
|
@ -638,8 +706,16 @@ def _compile_blocks(transformer) -> bool:
|
|||
|
||||
|
||||
def _dit_run(
|
||||
repo: str, family: str, *, dit_quant: str, steps: int, res: int, seed: int, iters: int,
|
||||
compile_blocks: bool = True, logger=None,
|
||||
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 to isolate
|
||||
the DiT), regional-compile it (the real feature path), then measure per-step + total latency and
|
||||
|
|
@ -649,21 +725,24 @@ def _dit_run(
|
|||
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")
|
||||
_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)
|
||||
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)
|
||||
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)
|
||||
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()
|
||||
|
|
@ -682,14 +761,24 @@ def _dit_run(
|
|||
}, last_img
|
||||
|
||||
|
||||
def measure_dit(family: str, *, schemes, steps: int, res: int, seed: int, iters: int, out: Path, logger=None):
|
||||
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 (the whole-image accuracy metric)."""
|
||||
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
|
||||
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")
|
||||
|
|
@ -699,78 +788,105 @@ def measure_dit(family: str, *, schemes, steps: int, res: int, seed: int, iters:
|
|||
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)
|
||||
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
|
||||
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
|
||||
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)
|
||||
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")
|
||||
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")
|
||||
logging.basicConfig(level = logging.INFO, format = "%(message)s")
|
||||
logger = logging.getLogger("speedmem")
|
||||
|
||||
out = Path(args.out)
|
||||
out.mkdir(parents=True, exist_ok=True)
|
||||
out.mkdir(parents = True, exist_ok = True)
|
||||
|
||||
print(f"== speed+mem bench: family={args.family} mode={args.mode} ==", flush=True)
|
||||
print(f"== speed+mem bench: family={args.family} mode={args.mode} ==", flush = True)
|
||||
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
|
||||
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)
|
||||
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)
|
||||
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)
|
||||
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
|
||||
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)
|
||||
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)
|
||||
with open(dest, "w", encoding = "utf-8") as fh:
|
||||
json.dump(rows, fh, indent = 2)
|
||||
print(f"wrote {dest}", flush = True)
|
||||
return 0
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue