test1111111 / scripts /bench_compiler.py
spitfire4794's picture
CISM remote autobench: full source + fleet runner, serve results on 7860
28a1a01
Raw History Blame Contribute Delete
29.6 kB
"""Compiler / quant / kernel benchmark for Surjo-50m on CPU.
Covers three questions without adding hard dependencies to pyproject.toml:
1. torch.compile (inductor, CPU) on the naive-GDN Surjo path:
warmup seconds + steady-state tok/s vs HF eager vs CISM.
2. torchao int4 weight-only (optional): quantize + PPL + tok/s.
3. flash-linear-attention on CPU: do chunk_gdn2 / fused_recurrent run
on CPU? Parity vs naive_gdn2_reference + prefill T=512 timing.
Surjo note: SurjoModel/SurjoForCausalLM use a bespoke SurjoCache, NOT
transformers StaticCache. Passing a StaticCache is silently discarded
(modeling_surjo.py: ``if not isinstance(past, SurjoCache): past = None``
then a fresh cache is allocated per forward call, so state never carries).
All manual loops here therefore thread ``past_key_values`` through the
returned ``outputs.past_key_values`` instead of pre-allocating any cache.
Usage:
.venv/Scripts/python scripts/bench_compiler.py <model-id-or-path>
[--tokens 64] [--warmup 16] [--threads 4] [--prefill-tokens 512]
[--modes eager,compile,torchao,fla,cism] [--dynamic]
[--json-out report.json]
All of torch/transformers/torchao/fla/cism are optional at import time;
missing pieces record a ``skipped`` section instead of failing the run.
Exit 0 on success (even with skips), 2 on hard error.
"""
from __future__ import annotations
import argparse
import json
import math
import os
import sys
import time
from contextlib import suppress
from pathlib import Path
BENCH_PROMPT = "The future of artificial intelligence is"
try:
from cism.autobench import CORPUS # single source of truth for PPL scoring
except Exception:
CORPUS = (
"The history of computing is a history of abstractions. Every "
"generation of machines has hidden the complexity of the previous one behind a "
"cleaner interface, and every generation of programmers has used that interface "
"to build something the machine designers never imagined. Mechanical adding "
"machines abstracted arithmetic from the human hand. Stored-program computers "
"abstracted the program itself into data, so that a single machine could become "
"a loom, a calculator, or a chess opponent simply by loading different "
"instructions. Operating systems abstracted the hardware into processes, files, "
"and pipes. High-level languages abstracted registers and memory addresses into "
"variables and expressions. Each layer trades a small amount of performance for "
"an enormous increase in what a single mind can hold at once. "
"Language models continue this tradition in an unexpected direction. Instead of "
"abstracting machine details for human convenience, they abstract human "
"expressions into vectors, and then predict which expression comes next. "
"Running these models on consumer hardware is an exercise in memory bandwidth. "
"A decoder step reads every weight exactly once, so the maximum generation "
"speed is roughly the bandwidth of the memory bus divided by the size of the "
"model. Quantization shrinks the weights and therefore raises the ceiling, but "
"it also perturbs every dot product in the network."
)
def _ensure_msvc() -> str | None:
"""Best-effort vcvars import so inductor finds cl.exe on Windows."""
if os.name != "nt":
return None
vswhere = Path(r"C:\Program Files (x86)\Microsoft Visual Studio\Installer\vswhere.exe")
if not vswhere.exists():
return None
import subprocess
import tempfile
try:
root = subprocess.run(
[str(vswhere), "-latest", "-products", "*",
"-requires", "Microsoft.VisualStudio.Component.VC.Tools.x86.x64",
"-property", "installationPath"],
capture_output=True, text=True, timeout=30).stdout.strip()
except (OSError, subprocess.TimeoutExpired):
return None
if not root:
return None
vcvars = Path(root) / "VC" / "Auxiliary" / "Build" / "vcvars64.bat"
if not vcvars.exists():
return None
try:
with tempfile.NamedTemporaryFile("w", suffix=".bat", delete=False) as handle:
handle.write(f'@echo off\r\ncall "{vcvars}" >nul 2>&1\r\nset\r\n')
script = handle.name
completed = subprocess.run(["cmd.exe", "/c", script],
capture_output=True, text=True, timeout=120)
except (OSError, subprocess.TimeoutExpired):
return None
finally:
with suppress(OSError):
os.unlink(script)
for line in completed.stdout.splitlines():
name, sep, value = line.partition("=")
if sep and name.upper() in ("PATH", "INCLUDE", "LIB", "LIBPATH"):
if name.upper() == "PATH":
os.environ["PATH"] = value + os.pathsep + os.environ.get("PATH", "")
else:
os.environ[name.upper()] = value
return root
def _timed_count(fn, warmup: int, measured: int) -> dict:
fn(warmup)
start = time.perf_counter()
got = fn(measured)
elapsed = time.perf_counter() - start
return {"completion_tokens": got, "seconds": round(elapsed, 4),
"tokens_per_second": round(got / elapsed, 2) if elapsed > 0 else 0.0}
def _hf_nll(model, tokenizer, corpus: str, window: int) -> dict:
import torch
ids = tokenizer(corpus, return_tensors="pt").input_ids[0]
total, count = 0.0, 0
step = window - 1
with torch.inference_mode():
for start in range(0, len(ids) - 1, step):
chunk = ids[start:start + window].unsqueeze(0)
targets = ids[start + 1:start + window]
logits = model(chunk).logits[0]
logprobs = torch.log_softmax(logits, dim=-1)
picked = -logprobs[torch.arange(len(targets)), targets]
total += float(picked.sum())
count += len(targets)
mean = total / count
try:
ppl = math.exp(mean)
except OverflowError:
ppl = float("inf")
return {"nll": mean, "perplexity": ppl, "scored_tokens": count}
def _dynamo_snapshot() -> dict:
try:
import torch._dynamo as dynamo
return {k: sum(v.values()) for k, v in dynamo.utils.counters.items()}
except Exception:
return {}
def _dynamo_delta(before: dict, after: dict) -> dict:
keys = set(before) | set(after)
return {k: after.get(k, 0) - before.get(k, 0) for k in keys
if after.get(k, 0) != before.get(k, 0)}
def _versions() -> dict:
out = {"python": sys.version.split()[0], "platform": sys.platform}
for mod in ("torch", "transformers", "torchao", "fla", "flash_attn", "triton", "cism"):
try:
out[mod] = __import__("importlib.metadata").metadata.version(mod.replace("_", "-"))
except Exception:
try:
m = __import__(mod)
out[mod] = getattr(m, "__version__", "installed-unknown")
except Exception as e:
out[mod] = f"missing ({type(e).__name__})"
try:
import torch
out["torch_version_detail"] = torch.__version__
out["torch_threads"] = torch.get_num_threads()
except Exception:
pass
return out
def run_eager_generate(model, ids, eos_pad, limit: int, warmup: int) -> dict:
def loop(count: int) -> int:
model.generate(ids, attention_mask=None, max_new_tokens=count,
min_new_tokens=count, do_sample=False, num_beams=1,
use_cache=True, pad_token_id=eos_pad)
return count
out = _timed_count(loop, min(warmup, limit), limit)
out["method"] = "generate(use_cache=True) eager"
return out
def run_eager_manual(model, ids, limit: int, warmup: int) -> dict:
"""Correct Surjo decode loop: thread returned past_key_values through."""
import torch
def loop(count: int) -> int:
with torch.inference_mode():
out = model(ids, past_key_values=None, use_cache=True)
past = out.past_key_values
token = int(out.logits[0, -1].argmax())
generated = 1
one = torch.tensor([[token]])
while generated < count:
out = model(one, past_key_values=past, use_cache=True)
past = out.past_key_values
token = int(out.logits[0, -1].argmax())
one = torch.tensor([[token]])
generated += 1
return generated
out = _timed_count(loop, min(warmup, limit), limit)
out["method"] = "manual SurjoCache loop eager (past threaded)"
return out
def run_compiled(model, ids, limit: int, warmup: int, dynamic: bool) -> dict:
import torch
try:
import torch._dynamo as dynamo
dynamo.config.recompile_limit = 256
except Exception:
pass
_ensure_msvc()
compiled = torch.compile(model, dynamic=dynamic)
def loop_ids(count: int) -> list[int]:
with torch.inference_mode():
out = compiled(ids, past_key_values=None, use_cache=True)
past = out.past_key_values
toks = [int(out.logits[0, -1].argmax())]
one = torch.tensor([[toks[0]]])
while len(toks) < count:
out = compiled(one, past_key_values=past, use_cache=True)
past = out.past_key_values
toks.append(int(out.logits[0, -1].argmax()))
one = torch.tensor([[toks[-1]]])
return toks
def loop(count: int) -> int:
return len(loop_ids(count))
t0 = time.perf_counter()
loop(min(warmup, limit))
loop(min(warmup, limit))
warmup_s = time.perf_counter() - t0
before = _dynamo_snapshot()
out = _timed_count(loop, min(warmup, limit), limit)
after = _dynamo_snapshot()
out["compile_warmup_seconds"] = round(warmup_s, 1)
out["dynamo_counter_delta_during_measurement"] = _dynamo_delta(before, after)
out["method"] = f"torch.compile(dynamic={dynamic}) + threaded SurjoCache loop"
try:
out["sample_token_ids"] = loop_ids(min(8, limit))
except Exception as e:
out["sample_token_ids_error"] = f"{type(e).__name__}: {e}"
return out
def run_prefill(model, tokenizer, tokens: int) -> dict:
"""Single forward at length T: the GDN-chunk path (prefill) timing."""
import torch
vocab = model.config.vocab_size
ids = (torch.arange(tokens).unsqueeze(0) % vocab).to("cpu")
# Warmup once (allocations), then time 3 reps, report best.
with torch.inference_mode():
model(ids[:, :16])
best = float("inf")
out_tokens = 0
for _ in range(3):
t0 = time.perf_counter()
with torch.inference_mode():
out = model(ids)
dt = time.perf_counter() - t0
best = min(best, dt)
out_tokens = int(out.logits.shape[1])
return {"sequence_length": tokens, "seconds_best_of_3": round(best, 4),
"tokens_per_second_prefill": round(tokens / best, 2) if best > 0 else 0.0,
"output_length": out_tokens}
def run_torchao_section(model_id, tokenizer, ids, eos_pad, limit, warmup,
window, local_files_only, trust: bool) -> dict:
try:
from torchao.quantization.quant_api import Int4WeightOnlyConfig as Int4Cfg
from torchao.quantization.quant_api import Int8WeightOnlyConfig as Int8Cfg
from torchao.quantization.quant_api import quantize_ as tao_quantize
except ImportError as e:
return {"skipped": f"torchao not installed ({e}); pip install torchao --index-url https://download.pytorch.org/whl/cpu"}
import torch
from transformers import AutoModelForCausalLM
section: dict = {"attempts": {}}
def _bench_variant(label: str, cfg) -> dict:
entry: dict = {"config": str(cfg)}
try:
m = AutoModelForCausalLM.from_pretrained(
model_id, dtype=torch.float32, local_files_only=local_files_only,
trust_remote_code=trust).to("cpu").eval()
except Exception as e:
return {"failed": f"reload failed: {type(e).__name__}: {e}"}
t0 = time.perf_counter()
try:
tao_quantize(m, cfg)
except Exception as e:
entry["quantize_failed"] = f"{type(e).__name__}: {e}"
return entry
entry["quantize_seconds"] = round(time.perf_counter() - t0, 2)
def loop(count: int) -> int:
m.generate(ids, attention_mask=None, max_new_tokens=count,
min_new_tokens=count, do_sample=False, num_beams=1,
use_cache=True, pad_token_id=eos_pad)
return count
try:
entry["decode"] = _timed_count(loop, min(warmup, limit), limit)
except Exception as e:
entry["decode"] = {"failed": f"{type(e).__name__}: {e}"}
try:
entry["perplexity"] = _hf_nll(m, tokenizer, CORPUS, window)
except Exception as e:
entry["perplexity"] = {"failed": f"{type(e).__name__}: {e}"}
try:
entry["decode_manual"] = run_eager_manual(m, ids, limit, warmup)
except Exception as e:
entry["decode_manual"] = {"failed": f"{type(e).__name__}: {e}"}
return entry
# Primary: default int4 weight-only (group 128, PLAIN/TINYGEMM). On
# Windows CPU this currently fails with "Requires mslk >= 1.0.0"
# (no Windows wheel); recorded verbatim, not hidden.
section["attempts"]["int4_default_g128"] = _bench_variant(
"int4_default_g128", Int4Cfg())
# Fallback that works on CPU without mslk: int8 weight-only.
try:
section["attempts"]["int8wo_fallback"] = _bench_variant(
"int8wo_fallback", Int8Cfg())
except Exception as e:
section["attempts"]["int8wo_fallback"] = {"failed": f"{type(e).__name__}: {e}"}
section["note"] = ("int4_default is the requested torchao path; int8wo_fallback "
"is the CPU-working torchao control")
return section
def run_fla_section(prefill_tokens: int) -> dict:
section: dict = {}
try:
import torch
from fla.ops.gdn2 import chunk_gdn2, fused_recurrent_gdn2
section["import"] = "ok"
except Exception as e:
return {"skipped": f"fla.ops.gdn2 unavailable ({type(e).__name__}: {e}); pip install 'flash-linear-attention[cpu]'"}
# Naive reference lives in the Surjo checkpoint's modeling file; fall back
# to an inline reimplementation if it cannot be imported.
naive = None
try:
import importlib.util
# Resolve via the loaded model's modeling file is overkill here;
# naive_gdn2_reference has a stable 20-line definition - reimplement.
import torch.nn.functional as F
def naive(q, k, v, g, b, w, initial_state=None):
B, T, H, K = q.shape
S = (torch.zeros(B, H, K, v.shape[-1], dtype=torch.float32)
if initial_state is None else initial_state.float())
scale = K ** -0.5
outs = []
for t in range(T):
q_t = F.normalize(q[:, t], dim=-1)
k_t = F.normalize(k[:, t], dim=-1)
v_t, g_t, b_t, w_t = v[:, t], g[:, t], b[:, t], w[:, t]
S = S * torch.exp(g_t).unsqueeze(-1)
e = b_t * k_t
r = torch.einsum("bhkv,bhk->bhv", S, e)
z = w_t * v_t
S = S + torch.einsum("bhk,bhv->bhkv", k_t, z - r)
outs.append(scale * torch.einsum("bhkv,bhk->bhv", S, q_t))
return torch.stack(outs, dim=1), S
section["naive_source"] = "inline (mirrors modeling_surjo.naive_gdn2_reference)"
except Exception as e:
return {"failed": f"naive reference unavailable: {e}"}
# Surjo-50m GDN geometry: H=8 heads, K=64, V=64, Hv=8.
B, H, K, Vd = 1, 8, 64, 64
torch.manual_seed(0)
section["cpu_kernel_probe"] = {}
# 1) fused_recurrent, T=1 (decode step) on CPU.
try:
q = torch.randn(B, 1, H, K), torch.randn(B, 1, H, K)
qq, kk = q
vv = torch.randn(B, 1, H, Vd)
gg = -torch.rand(B, 1, H, K) # decay log-rates (negative)
bb = torch.rand(B, 1, H, K).sigmoid()
ww = torch.rand(B, 1, H, Vd).sigmoid()
ref_o, ref_s = naive(qq.float(), kk.float(), vv.float(), gg, bb.float(), ww.float(), None)
t0 = time.perf_counter()
fla_o, fla_s = fused_recurrent_gdn2(
qq, kk, vv, gg, bb, ww, initial_state=None, output_final_state=True,
use_qk_l2norm_in_kernel=True)
dt = time.perf_counter() - t0
err = float((fla_o.float() - ref_o).abs().max())
section["cpu_kernel_probe"]["fused_recurrent_T1"] = {
"ran_on_cpu": True, "seconds": round(dt, 4),
"max_abs_err_vs_naive": err,
"parity": bool(err < 1e-3),
}
except Exception as e:
section["cpu_kernel_probe"]["fused_recurrent_T1"] = {
"ran_on_cpu": False, "error": f"{type(e).__name__}: {e}"}
# 2) chunk, T=prefill_tokens (prefill) on CPU + timing vs naive.
T = int(prefill_tokens)
# Naive at T=512 is ~seconds; cap parity at T=32, time chunk at full T.
for t_len, key in ((min(T, 32), "chunk_T32_parity"), (T, f"chunk_T{T}_timing")):
try:
qq = torch.randn(B, t_len, H, K)
kk = torch.randn(B, t_len, H, K)
vv = torch.randn(B, t_len, H, Vd)
gg = -torch.rand(B, t_len, H, K)
bb = torch.rand(B, t_len, H, K).sigmoid()
ww = torch.rand(B, t_len, H, Vd).sigmoid()
t0 = time.perf_counter()
fla_o, _ = chunk_gdn2(
q=qq, k=kk, v=vv, g=gg, b=bb, w=ww,
initial_state=None, output_final_state=False,
use_qk_l2norm_in_kernel=True, cu_seqlens=None)
dt_fla = time.perf_counter() - t0
entry: dict = {"ran_on_cpu": True, "seconds_fla": round(dt_fla, 4)}
if key.endswith("parity"):
t0 = time.perf_counter()
ref_o, _ = naive(qq.float(), kk.float(), vv.float(), gg,
bb.float(), ww.float(), None)
dt_naive = time.perf_counter() - t0
err = float((fla_o.float() - ref_o).abs().max())
entry.update({"seconds_naive": round(dt_naive, 4),
"max_abs_err_vs_naive": err,
"parity": bool(err < 1e-3),
"speedup_fla_vs_naive": round(dt_naive / dt_fla, 2) if dt_fla > 0 else None})
section["cpu_kernel_probe"][key] = entry
except Exception as e:
section["cpu_kernel_probe"][key] = {
"ran_on_cpu": False, "error": f"{type(e).__name__}: {e}"}
return section
def run_cism_section(model_id, threads: int, limit: int, warmup: int,
window: int, local_files_only: bool) -> dict:
try:
from cism import Engine
from cism.autobench import CORPUS as ABCORPUS
except Exception as e:
return {"skipped": f"cism unavailable: {type(e).__name__}: {e}"}
section: dict = {}
for prec in ("fp32", "int8"):
try:
eng = Engine.from_pretrained(model_id, precision=prec, threads=threads,
local_files_only=local_files_only)
except Exception as e:
section[prec] = {"failed": f"{type(e).__name__}: {e}"}
continue
prompt = BENCH_PROMPT
n = min(limit, eng.context_length - len(eng.encode(prompt)))
def run(count: int, _e=eng) -> int:
return len(_e.generate(prompt, max_new_tokens=count, temperature=0).token_ids)
timed = _timed_count(run, min(warmup, n), n)
try:
timed["perplexity"] = eng.nll(ABCORPUS, window=window)
except Exception as e:
timed["perplexity"] = {"failed": f"{type(e).__name__}: {e}"}
timed["weight_mb"] = round(int(eng.info.get("weight_bytes", 0) or 0) / (1024 * 1024), 2)
section[prec] = timed
return section
def main() -> int:
ap = argparse.ArgumentParser(description=__doc__)
ap.add_argument("model", help="HF id or local snapshot path (e.g. SurjoLabs/Surjo-50m)")
ap.add_argument("--tokens", type=int, default=64)
ap.add_argument("--warmup", type=int, default=16)
ap.add_argument("--threads", type=int, default=4)
ap.add_argument("--prefill-tokens", type=int, default=512)
ap.add_argument("--window", type=int, default=128)
ap.add_argument("--modes", default="eager,compile,torchao,fla,cism,prefill",
help="comma list of eager,compile,torchao,fla,cism,prefill")
ap.add_argument("--dynamic", action="store_true",
help="also compile with dynamic=True (doubles warmup cost)")
ap.add_argument("--local-files-only", action="store_true", default=True)
ap.add_argument("--json-out", default=None)
args = ap.parse_args()
modes = {m.strip().lower() for m in args.modes.split(",") if m.strip()}
report: dict = {"model": args.model, "threads": args.threads,
"measured_tokens": args.tokens, "warmup_tokens": args.warmup,
"prefill_tokens": args.prefill_tokens, "modes": sorted(modes),
"versions": _versions()}
try:
import torch
torch.set_num_threads(args.threads)
try:
torch.set_num_interop_threads(1)
except RuntimeError:
pass
report["versions"]["torch_threads_after_set"] = torch.get_num_threads()
except Exception as e:
report["torch_thread_error"] = f"{type(e).__name__}: {e}"
print(json.dumps(report, indent=2))
return 2
from transformers import AutoModelForCausalLM, AutoTokenizer
trust = "surjo" in str(args.model).lower()
if not trust:
try:
import json as _j
cfg = _j.loads(Path(str(args.model)).joinpath("config.json").read_text(encoding="utf-8"))
trust = cfg.get("model_type") == "surjo"
except Exception:
pass
try:
tokenizer = AutoTokenizer.from_pretrained(
args.model, local_files_only=args.local_files_only, trust_remote_code=trust)
kw: dict = {"dtype": torch.float32, "local_files_only": args.local_files_only,
"trust_remote_code": trust} if 'torch' in dir() else {}
import torch as _t
kw = {"dtype": _t.float32, "local_files_only": args.local_files_only,
"trust_remote_code": trust}
model = AutoModelForCausalLM.from_pretrained(args.model, **kw).to("cpu").eval()
except Exception as e:
report["load_failed"] = f"{type(e).__name__}: {e}"
print(json.dumps(report, indent=2))
return 2
report["surjo"] = bool(getattr(getattr(model, "config", None), "model_type", "") == "surjo")
ids = tokenizer(BENCH_PROMPT, return_tensors="pt").input_ids
eos_pad = tokenizer.eos_token_id or 0
try:
limit = min(args.tokens, model.config.max_position_embeddings - len(ids[0]))
except Exception:
limit = args.tokens
report["decode_limit"] = limit
if "eager" in modes:
try:
report["eager_generate"] = run_eager_generate(model, ids, eos_pad, limit, args.warmup)
except Exception as e:
report["eager_generate"] = {"failed": f"{type(e).__name__}: {e}"}
try:
report["eager_manual"] = run_eager_manual(model, ids, limit, args.warmup)
except Exception as e:
report["eager_manual"] = {"failed": f"{type(e).__name__}: {e}"}
try:
import torch as _t
with _t.inference_mode():
report["eager_ppl"] = _hf_nll(model, tokenizer, CORPUS,
min(args.window, model.config.max_position_embeddings))
except Exception as e:
report["eager_ppl"] = {"failed": f"{type(e).__name__}: {e}"}
if "prefill" in modes:
try:
report["eager_prefill"] = run_prefill(model, tokenizer, args.prefill_tokens)
except Exception as e:
report["eager_prefill"] = {"failed": f"{type(e).__name__}: {e}"}
if "compile" in modes:
# SurjoCache grows every step (XSA KV length), so dynamic=False
# specializes per length and recompiles per token (unusable). dynamic=True
# is the only viable inductor mode; both are measured for the report.
try:
report["compile_dynamic_true"] = run_compiled(model, ids, limit, args.warmup, dynamic=True)
except Exception as e:
report["compile_dynamic_true"] = {"failed": f"{type(e).__name__}: {e}"}
if args.dynamic:
try:
report["compile_dynamic_false"] = run_compiled(model, ids, limit, args.warmup, dynamic=False)
except Exception as e:
report["compile_dynamic_false"] = {"failed": f"{type(e).__name__}: {e}"}
else:
report["compile_dynamic_false"] = {"skipped": "pass --dynamic to also measure dynamic=False (recompiles per token on SurjoCache)"}
# Keep legacy key: the viable Surjo number is dynamic=True.
report["compile_fp32"] = report.get("compile_dynamic_true", {})
try:
eager_ids = None
with torch.inference_mode():
out = model(ids, past_key_values=None, use_cache=True)
past = out.past_key_values
eager_ids = [int(out.logits[0, -1].argmax())]
one = torch.tensor([[eager_ids[0]]])
while len(eager_ids) < min(8, limit):
out = model(one, past_key_values=past, use_cache=True)
past = out.past_key_values
eager_ids.append(int(out.logits[0, -1].argmax()))
one = torch.tensor([[eager_ids[-1]]])
comp_ids = (report.get("compile_dynamic_true") or {}).get("sample_token_ids")
report["compile_parity"] = {"eager_first8": eager_ids,
"compiled_first8": comp_ids,
"match": bool(eager_ids == comp_ids) if comp_ids else None}
except Exception as e:
report["compile_parity"] = {"failed": f"{type(e).__name__}: {e}"}
if "prefill" in modes and isinstance(report.get("compile_dynamic_true"), dict) and "seconds" in report.get("compile_dynamic_true", {}):
try:
import torch as _t
compiled_pf = _t.compile(model, dynamic=True)
# Warm the prefill shape once outside timing.
with _t.inference_mode():
_voc = model.config.vocab_size
_warm = (torch.arange(16).unsqueeze(0) % _voc).to("cpu")
compiled_pf(_warm)
report["compile_prefill"] = run_prefill(compiled_pf, tokenizer, args.prefill_tokens)
except Exception as e:
report["compile_prefill"] = {"failed": f"{type(e).__name__}: {e}"}
if "torchao" in modes:
report["torchao_int4"] = run_torchao_section(
args.model, tokenizer, ids, eos_pad, limit, args.warmup,
min(args.window, model.config.max_position_embeddings),
args.local_files_only, trust)
if "fla" in modes:
report["fla_cpu"] = run_fla_section(args.prefill_tokens)
if "cism" in modes:
report["cism"] = run_cism_section(args.model, args.threads, args.tokens,
args.warmup, args.window, args.local_files_only)
# Summary speedups where numbers exist.
try:
eager_rate = (report.get("eager_manual") or {}).get("tokens_per_second") or \
(report.get("eager_generate") or {}).get("tokens_per_second")
compile_rate = (report.get("compile_dynamic_true") or {}).get("tokens_per_second")
compile_rate_f = (report.get("compile_dynamic_false") or {}).get("tokens_per_second")
cism_int8 = ((report.get("cism") or {}).get("int8") or {}).get("tokens_per_second")
cism_fp32 = ((report.get("cism") or {}).get("fp32") or {}).get("tokens_per_second")
summary = {}
if eager_rate and compile_rate:
summary["compile_dynamic_true_speedup_vs_eager"] = round(compile_rate / eager_rate, 2)
if eager_rate and compile_rate_f:
summary["compile_dynamic_false_speedup_vs_eager"] = round(compile_rate_f / eager_rate, 2)
if eager_rate and cism_int8:
summary["cism_int8_speedup_vs_eager"] = round(cism_int8 / eager_rate, 2)
if compile_rate and cism_int8:
summary["cism_int8_speedup_vs_compile_dynamic_true"] = round(cism_int8 / compile_rate, 2)
if eager_rate and cism_fp32:
summary["cism_fp32_speedup_vs_eager"] = round(cism_fp32 / eager_rate, 2)
# Legacy aliases for autobench comparison.
if eager_rate and compile_rate:
summary["compile_speedup_vs_eager"] = summary["compile_dynamic_true_speedup_vs_eager"]
report["summary"] = summary
except Exception:
pass
text = json.dumps(report, indent=2)
print(text)
if args.json_out:
Path(args.json_out).write_text(text, encoding="utf-8")
return 0
if __name__ == "__main__":
raise SystemExit(main())