Spaces:
Sleeping
Sleeping
Download scripts/bench_compiler.py from spitfire4794/test1111111: direct link, hf CLI and curl.
- Browser
- Download file 29.6 kB
-
https://huggingface.co/spaces/spitfire4794/test1111111/resolve/main/scripts/bench_compiler.py
- Command line
-
hf download hf://spaces/spitfire4794/test1111111/scripts/bench_compiler.py
-
curl -L -o bench_compiler.py https://huggingface.co/spaces/spitfire4794/test1111111/resolve/main/scripts/bench_compiler.py
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()) | |