From 5c3e7a89a79784b9951e81bab194594d58d0b353 Mon Sep 17 00:00:00 2001 From: Jeongkyu Shin Date: Sat, 12 Sep 2026 12:16:31 +0900 Subject: [PATCH 01/13] perf(cuda): bucket the cuDNN SDPA plan-cache key MLX keys its cuDNN SDPA execution-plan cache on the exact shapes and strides of q, k, v and the mask, and a speculative verify round appends keys every round, so every round of every attention-layer class missed that cache and rebuilt a plan on the host: about 22 ms per build on GB10, 67 to 76 ms per round on the Laguna DFlash pairing, and a fatal `Cache thrashing` abort once lifetime misses passed twice the capacity. #1799 routed those calls off cuDNN entirely; this makes the key reusable so they can stay on it. Measured on the pairing #1799 attributed the defect on (GB10, block 4, 25 verify rounds, `MLXCEL_SDPA_PLAN_DEBUG=1`): three shape classes, three plan builds per round, 82 in the process. Only three key fields move round over round, the k/v sequence length and the mask's column count and row stride; the k/v strides and the mask's two leading strides (0, from the broadcast `fast::scaled_dot_product_attention` builds) do not. With bucketing each class builds one plan for the whole run, 11 in the process, and the round's accept counts are unchanged. The canonicalization is MLX's own one-row decode path extended to a small array-masked multi-row call: k and v are widened to a bucket, the mask is widened to the same width with the new columns set to `-inf`, and the true lengths reach cuDNN through `set_padding_mask` with `set_seq_len_q` / `set_seq_len_kv`, which is what keeps the widened region out of the result. k and v reach the bucket by unslicing when they are a leading slice of one cache buffer with room to spare (the target's dense and speculative-buffered caches), and by a zero-padded copy when they are not, which is the third class: the drafter concatenates its proposal keys onto the cache window, so its k and v are freshly built arrays. The copy is bounded by `MLXCEL_SDPA_PLAN_BUCKET_MAX_MB` (64 per tensor); above it the call keeps its exact shape. `MLXCEL_SDPA_PLAN_BUCKET_MAX_QUERIES` (default 32, 0 disables) is the kill switch, and it narrows #1799's gate rather than leaving both mechanisms in place: `MLXCEL_SDPA_FALLBACK_MAX_QUERIES` no longer claims a call whose key can be bucketed, so what it still covers is a causal-mode block with no array mask, where there is no mask to widen. `MLXCEL_SDPA_PLAN_DEBUG=1` traces the key fields and plan builds per call. A cuDNN that cannot plan the widened shape degrades to per-length plans instead of aborting, per sinks setting so one refusal cannot disable the other shape. Refs #1820 --- .../harness/bench_cli.py | 132 ++ .../harness/bench_srv.py | 101 ++ .../harness/hostgate.py | 36 + .../harness/identity.sh | 15 + .../harness/idle_gate.sh | 10 + .../harness/prompt_code0.txt | 17 + .../harness/prompt_code_long.txt | 494 ++++++ .../harness/summarize_cli.py | 54 + .../harness/summarize_srv.py | 22 + .../harness/summarize_trace.py | 44 + .../harness/sweep_laguna.sh | 24 + .../harness/sweep_nonspec.sh | 18 + .../harness/trace_keys.sh | 19 + .../trace_upstream_w4.txt | 1401 +++++++++++++++++ docs/environment-variables.md | 4 +- .../cuda/scaled_dot_product_attention.cpp | 497 +++++- src/lib/mlxcel-core/src/layers.rs | 68 +- src/lib/mlxcel-core/src/lib.rs | 8 + .../mlxcel-core/src/sdpa_plan_bucket_tests.rs | 188 +++ 19 files changed, 3083 insertions(+), 69 deletions(-) create mode 100755 docs/benchmark_results/data/sdpa-plan-bucket-gb10-2026-09-12/harness/bench_cli.py create mode 100755 docs/benchmark_results/data/sdpa-plan-bucket-gb10-2026-09-12/harness/bench_srv.py create mode 100755 docs/benchmark_results/data/sdpa-plan-bucket-gb10-2026-09-12/harness/hostgate.py create mode 100755 docs/benchmark_results/data/sdpa-plan-bucket-gb10-2026-09-12/harness/identity.sh create mode 100755 docs/benchmark_results/data/sdpa-plan-bucket-gb10-2026-09-12/harness/idle_gate.sh create mode 100644 docs/benchmark_results/data/sdpa-plan-bucket-gb10-2026-09-12/harness/prompt_code0.txt create mode 100644 docs/benchmark_results/data/sdpa-plan-bucket-gb10-2026-09-12/harness/prompt_code_long.txt create mode 100755 docs/benchmark_results/data/sdpa-plan-bucket-gb10-2026-09-12/harness/summarize_cli.py create mode 100755 docs/benchmark_results/data/sdpa-plan-bucket-gb10-2026-09-12/harness/summarize_srv.py create mode 100755 docs/benchmark_results/data/sdpa-plan-bucket-gb10-2026-09-12/harness/summarize_trace.py create mode 100755 docs/benchmark_results/data/sdpa-plan-bucket-gb10-2026-09-12/harness/sweep_laguna.sh create mode 100755 docs/benchmark_results/data/sdpa-plan-bucket-gb10-2026-09-12/harness/sweep_nonspec.sh create mode 100755 docs/benchmark_results/data/sdpa-plan-bucket-gb10-2026-09-12/harness/trace_keys.sh create mode 100644 docs/benchmark_results/data/sdpa-plan-bucket-gb10-2026-09-12/trace_upstream_w4.txt create mode 100644 src/lib/mlxcel-core/src/sdpa_plan_bucket_tests.rs diff --git a/docs/benchmark_results/data/sdpa-plan-bucket-gb10-2026-09-12/harness/bench_cli.py b/docs/benchmark_results/data/sdpa-plan-bucket-gb10-2026-09-12/harness/bench_cli.py new file mode 100755 index 000000000..95cd79234 --- /dev/null +++ b/docs/benchmark_results/data/sdpa-plan-bucket-gb10-2026-09-12/harness/bench_cli.py @@ -0,0 +1,132 @@ +#!/usr/bin/env python3 +"""CLI-driven A/B sweep for issue #1799 (Laguna DFlash verify cost, GB10). + +Runs `mlxcel generate` once per (config, round), round-robin over the configs +so drift spreads across the table instead of pooling on one arm. Same binary +on every arm; the classic arm is the same command without --draft-model. +Records decode tok/s as the CLI reports it, the `DFlash:` diagnostics line, +the generated token ids (greedy identity), and the 1-minute load average +before each run (host idleness evidence). + +Usage: + bench_cli.py --bin B --target T --draft D --prompt-file P --out results.jsonl \ + --rounds 3 --configs off,b2,b4,b8,b16 [--env K=V ...] [--max-tokens 200] + config grammar: off | b | -b | -off + prefix maps through --preset NAME=K=V,K=V (e.g. ops100=MLX_MAX_OPS_PER_BUFFER=100) +""" +import argparse +import json +import os +import re +import shlex +import subprocess +import sys +import time +import hostgate + +ANSI = re.compile(r"\x1b\[[0-9;]*m") +GEN_RE = re.compile(r"\[Generated (\d+) tokens in ([0-9.]+)s = ([0-9.]+) tok/s\]") +DFLASH_RE = re.compile(r"^DFlash: (.*)$", re.M) +KV_RE = re.compile(r"(\w+)=([-0-9.]+)") +IDS_RE = re.compile(r"\[token ids \((\d+)\): ([0-9 ]*)\]") + + +def run_once(a, cfg, presets, prompt, wrap=""): + env = dict(os.environ) + env["MLXCEL_PRINT_TOKEN_IDS"] = "1" + env["MLXCEL_MTP_ALLOW_INEXACT"] = "1" + for kv in a.env: + k, v = kv.split("=", 1) + env[k] = v + prefix, _, arm = cfg.rpartition("-") if "-" in cfg else ("", "", cfg) + if prefix: + for kv in presets[prefix]: + k, v = kv.split("=", 1) + env[k] = v + cmd = shlex.split(wrap) + [a.bin, "generate", "-m", a.target, "-p", prompt, + "-n", str(a.max_tokens), "--temp", "0"] + block = None + if arm != "off": + block = int(arm[1:]) + cmd += ["--draft-model", a.draft, "--draft-kind", "dflash", + "--draft-block-size", str(block)] + gate_wait_s = hostgate.wait_quiet(log=sys.stderr) + load1 = os.getloadavg()[0] + ci_job = hostgate.ci_job_running() + t0 = time.perf_counter() + p = subprocess.run(cmd, env=env, capture_output=True, text=True, timeout=1800) + wall = time.perf_counter() - t0 + out = ANSI.sub("", p.stdout) + err = ANSI.sub("", p.stderr) + rec = {"cfg": cfg, "block": block, "load1_before": load1, "ci_job_running": ci_job, "gate_wait_s": gate_wait_s, "wall_s": wall, + "rc": p.returncode} + m = GEN_RE.search(out) + if m: + rec["gen_tokens"] = int(m.group(1)) + rec["decode_s"] = float(m.group(2)) + rec["tok_s"] = float(m.group(3)) + m = DFLASH_RE.search(out) + if m: + rec["diag"] = {k: float(v) for k, v in KV_RE.findall(m.group(1))} + m = IDS_RE.search(err) + if m: + rec["ids"] = m.group(2).strip() + if p.returncode != 0 or "tok_s" not in rec: + rec["stderr_tail"] = err[-2000:] + rec["stdout_tail"] = out[-1000:] + return rec + + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("--bin", required=True) + ap.add_argument("--target", required=True) + ap.add_argument("--draft", required=True) + ap.add_argument("--prompt-file", required=True) + ap.add_argument("--out", required=True) + ap.add_argument("--configs", required=True) + ap.add_argument("--rounds", type=int, default=3) + ap.add_argument("--warmup", type=int, default=1) + ap.add_argument("--max-tokens", type=int, default=200) + ap.add_argument("--env", action="append", default=[]) + ap.add_argument("--preset", action="append", default=[]) + ap.add_argument("--wrap", default="") + ap.add_argument("--tag", default="") + a = ap.parse_args() + presets = {} + for p in a.preset: + name, _, kvs = p.partition("=") + presets[name] = kvs.split(",") + prompt = open(a.prompt_file).read() + cfgs = a.configs.split(",") + with open(a.out, "a") as f: + for i in range(a.warmup): + r = run_once(a, cfgs[0], presets, prompt, a.wrap) + r["warmup"] = True + r["tag"] = a.tag + print(f"[warmup {i}] {cfgs[0]} tok/s={r.get('tok_s')} rc={r['rc']}", file=sys.stderr, flush=True) + f.write(json.dumps(r) + "\n"); f.flush() + for rd in range(a.rounds): + n = len(cfgs) + for k in range(n): + cfg = cfgs[(rd + k) % n] + r = run_once(a, cfg, presets, prompt, a.wrap) + r["warmup"] = False + r["round"] = rd + r["tag"] = a.tag + d = r.get("diag", {}) + extra = "" + if d.get("rounds"): + n_r = d["rounds"] + extra = (f" rounds={n_r:.0f} acc={d['accepted']/n_r:.2f}" + f" draft/r={d['draft_ms']/n_r:.1f} vgraph/r={d['verify_graph_ms']/n_r:.1f}" + f" sync/r={d['verify_sync_ms']/n_r:.1f}") + print(f"[round {rd}] {cfg} tok/s={r.get('tok_s')} load1={r['load1_before']:.2f}" + f" rc={r['rc']}{extra}", file=sys.stderr, flush=True) + if r["rc"] != 0: + print(r.get("stderr_tail", "")[-800:], file=sys.stderr, flush=True) + f.write(json.dumps(r) + "\n"); f.flush() + + +if __name__ == "__main__": + main() diff --git a/docs/benchmark_results/data/sdpa-plan-bucket-gb10-2026-09-12/harness/bench_srv.py b/docs/benchmark_results/data/sdpa-plan-bucket-gb10-2026-09-12/harness/bench_srv.py new file mode 100755 index 000000000..8ac73a1eb --- /dev/null +++ b/docs/benchmark_results/data/sdpa-plan-bucket-gb10-2026-09-12/harness/bench_srv.py @@ -0,0 +1,101 @@ +#!/usr/bin/env python3 +"""Batched-serving A/B sweep for issue #1798: one mlxcel-server per (config, round), +`scripts/bench_serving_concurrency.py` at the given concurrency ladder, round-robin +over configs. Records each level's row (ttft mean/p95, per-request decode tok/s, +aggregate tok/s) and load1 before the server start. Same binary on every arm. +Usage: bench_srv.py --server BIN --client scripts/bench_serving_concurrency.py --model M + --out X.jsonl --configs default,both --preset both=K=V,K=V --rounds 3 + [--concurrency 1,4] [--max-tokens 200] [--prompt-tokens 128] [--port 18798] +""" +import argparse, json, os, re, shlex, signal, subprocess, sys, time, urllib.request +import hostgate + +ROW = re.compile(r"^\s*(\d+)\s+(\d+)\s+(\d+)\s+([0-9.]+)\s+([0-9.]+)\s+([0-9.]+)\s+([0-9.]+)\s*$", re.M) + + +def wait_health(port, timeout=600): + t0 = time.time() + while time.time() - t0 < timeout: + try: + with urllib.request.urlopen(f"http://127.0.0.1:{port}/health", timeout=2) as r: + if r.status == 200: + return True + except Exception: + pass + time.sleep(1) + return False + + +def run_once(a, cfg, presets): + env = dict(os.environ) + env.setdefault("MLX_ENABLE_TF32", "0") + if cfg != "default": + for kv in presets[cfg]: + k, v = kv.split("=", 1) + env[k] = v + gate_wait_s = hostgate.wait_quiet(log=sys.stderr) + load1 = os.getloadavg()[0] + ci_job = hostgate.ci_job_running() + cmd = shlex.split(a.wrap) + [a.server, "-m", a.model, "--port", str(a.port), + "--max-batch-size", str(a.max_batch), "--ignore-eos"] + log = open(f"{a.out}.{cfg}.server.log", "a") + t0 = time.time() + srv = subprocess.Popen(cmd, env=env, stdout=log, stderr=subprocess.STDOUT, start_new_session=True) + rec = {"cfg": cfg, "load1_before": load1, "ci_job_running": ci_job, "gate_wait_s": gate_wait_s, "model": a.model, + "env": {k: env[k] for k in env if k.startswith("MLX_")}} + try: + if not wait_health(a.port): + rec["error"] = "server never became healthy" + return rec + rec["startup_s"] = time.time() - t0 + cl = ["python3", a.client, "--port", str(a.port), "--concurrency", a.concurrency, + "--prompt-tokens", str(a.prompt_tokens), "--max-tokens", str(a.max_tokens)] + p = subprocess.run(cl, capture_output=True, text=True, timeout=1800) + rec["rc"] = p.returncode + levels = [] + for m in ROW.finditer(p.stdout): + levels.append({"conc": int(m.group(1)), "ok": int(m.group(2)), "fail": int(m.group(3)), + "ttft_ms_mean": float(m.group(4)), "ttft_ms_p95": float(m.group(5)), + "decode_tok_s_mean": float(m.group(6)), "aggregate_tok_s": float(m.group(7))}) + rec["levels"] = levels + if not levels or p.returncode != 0: + rec["stdout_tail"] = p.stdout[-1500:]; rec["stderr_tail"] = p.stderr[-1500:] + finally: + try: + os.killpg(srv.pid, signal.SIGTERM) + srv.wait(timeout=60) + except Exception: + os.killpg(srv.pid, signal.SIGKILL) + log.close() + return rec + + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("--server", required=True); ap.add_argument("--client", required=True) + ap.add_argument("--model", required=True); ap.add_argument("--out", required=True) + ap.add_argument("--configs", required=True); ap.add_argument("--preset", action="append", default=[]) + ap.add_argument("--rounds", type=int, default=3); ap.add_argument("--concurrency", default="1,4") + ap.add_argument("--max-tokens", type=int, default=200); ap.add_argument("--prompt-tokens", type=int, default=128) + ap.add_argument("--max-batch", type=int, default=8); ap.add_argument("--port", type=int, default=18798) + ap.add_argument("--tag", default=""); ap.add_argument("--wrap", default="") + a = ap.parse_args() + presets = {} + for p in a.preset: + name, _, kvs = p.partition("="); presets[name] = kvs.split(",") + cfgs = a.configs.split(",") + with open(a.out, "a") as f: + for rd in range(a.rounds): + n = len(cfgs) + for k in range(n): + cfg = cfgs[(rd + k) % n] + r = run_once(a, cfg, presets); r["round"] = rd; r["tag"] = a.tag + lv = " ".join(f"c{l['conc']}:agg={l['aggregate_tok_s']:.1f}/dec={l['decode_tok_s_mean']:.1f}/ttft={l['ttft_ms_mean']:.0f}" for l in r.get("levels", [])) + print(f"[round {rd}] {cfg} load1={r['load1_before']:.2f} start={r.get('startup_s', 0):.0f}s {lv} {r.get('error', '')}", file=sys.stderr, flush=True) + if not r.get("levels"): + print(r.get("stderr_tail", "")[-600:], file=sys.stderr, flush=True) + f.write(json.dumps(r) + "\n"); f.flush() + + +if __name__ == "__main__": + main() diff --git a/docs/benchmark_results/data/sdpa-plan-bucket-gb10-2026-09-12/harness/hostgate.py b/docs/benchmark_results/data/sdpa-plan-bucket-gb10-2026-09-12/harness/hostgate.py new file mode 100755 index 000000000..cefa95ffa --- /dev/null +++ b/docs/benchmark_results/data/sdpa-plan-bucket-gb10-2026-09-12/harness/hostgate.py @@ -0,0 +1,36 @@ +"""Per-run host busy gate for the #1798 harnesses: block while compiler-like +processes (the self-hosted CI runner's builds, or anyone's) are alive, and +report the CI runner's job state so each record can state host idleness.""" +import subprocess, time, os + +BUSY = ("rustc", "cc1plus", "cc1", "nvcc", "cicc", "ptxas", "cmake", "ninja", "ld", "ld.lld", "lld", + "clang", "clang++", "gcc", "g++", "fatbinary", "cargo-clippy", "clippy-driver") + + +def _procs(): + out = subprocess.run(["ps", "-eo", "comm="], capture_output=True, text=True).stdout.split() + return out + + +def ci_job_running(): + return "Runner.Worker" in _procs() + + +def busy_procs(): + return sorted({p for p in _procs() if p in BUSY}) + + +def wait_quiet(samples=2, interval=5, log=None): + n = 0; waited = 0 + while n < samples: + b = busy_procs() + if b: + n = 0 + if log and waited % 60 == 0: + print(f"[gate] busy: {b} ci_job={ci_job_running()} load1={os.getloadavg()[0]:.2f}", file=log, flush=True) + time.sleep(interval); waited += interval + else: + n += 1 + if n < samples: + time.sleep(interval) + return waited diff --git a/docs/benchmark_results/data/sdpa-plan-bucket-gb10-2026-09-12/harness/identity.sh b/docs/benchmark_results/data/sdpa-plan-bucket-gb10-2026-09-12/harness/identity.sh new file mode 100755 index 000000000..33173c76e --- /dev/null +++ b/docs/benchmark_results/data/sdpa-plan-bucket-gb10-2026-09-12/harness/identity.sh @@ -0,0 +1,15 @@ +#!/usr/bin/env bash +# Greedy token-id identity, compared as ids, between two binaries or two env +# settings of one binary. Prints the id line for each arm and whether they match. +# +# identity.sh