From 3201b21eb18fdf9ca9f908aa808ec735840d64ac Mon Sep 17 00:00:00 2001 From: davidtai Date: Mon, 7 Sep 2026 07:36:58 -0500 Subject: [PATCH 01/30] perf(qwen4): three PR 391 remainder lanes for Qwen3.8 Flash-Next onto 2.11.2 MTPLX 2.11.2 re-landed nine of PR 391's twelve Flash-Next decode keys; this ports three of the unlanded five onto upstream's re-landed structure, wired into the fixed-M4 auto-arm block with the MTPLX_QWEN4_*/MTPLX_QSA_* namespace and a per-key =0 opt-out (the old MTPLX_FABLE_* names kept as aliases). - HC_M4 (MTPLX_QWEN4_HC_M4): the verify-width (2..8 row) hyper-connection read run as one multi-threadgroup GEMV (kernels/qwen4_m4_hyper_read). Reader in runtime_options read once at import; GatedResidual gains the geometry eligibility check, pack validation and the fused read; install validation runs after the M4-stage3 install and reports at /health qwen4_install_reports.hc_m4. Rounding-class. - prefill causal-mask fuse (MTPLX_QWEN4_PREFILL_MASK_FUSE): the dense QSA prefill chunk goes through MLX's fused SDPA instead of a materialized score tensor MLX's head-dim-256 heuristic declines; a per-shape-class capability cache keeps a verify step MLX refuses from disarming a wide chunk. Rounding-class (exact visible set). - QSA prefill query tile (MTPLX_QSA_PREFILL_QUERY_TILE): tiles only the dense QSA attention query rows so a wider prefill chunk keeps the narrow chunk's attention peak and cost. Value companion, default 2048 (inert at the production 2,048 chunk width). Rounding-class (exact visible set). Each is default-on for a served fixed-M4 Flash-Next pack (server auto-arm lane_defaults, gated on the fixed-M4 config predicate) with a per-key kill switch through the existing pop loop, and registered in the boot-time runtime-env validator. All three are rounding-class, so quality-gated on HumanEval. Two of the five remain and are documented in docs/perf/qwen38-391-remainder.md: the QSA sparse split-K decode (a native kernel whose build and parity probe need the GPU) and the graph-build overlap (its prefix/suffix split of the fixed-M4 verify has no substrate on 2.11.2's single-graph verify). CPU tests (venv mlx 0.32.2, no GPU): tests/test_qwen4_hc_m4.py 53 passed, tests/test_qwen4_prefill_mask_fuse.py 40 passed. --- docs/perf/qwen38-391-remainder.md | 31 + mtplx/kernels/qwen4_m4_hyper_read.py | 567 +++++++++++++++++ mtplx/models/qwen4_exp.py | 623 ++++++++++++++++++- mtplx/profiles.py | 7 + mtplx/qwen4_prefill_chunk.py | 77 +++ mtplx/runtime.py | 11 + mtplx/runtime_options.py | 28 + mtplx/server/openai.py | 27 + tests/test_qwen4_hc_m4.py | 602 ++++++++++++++++++ tests/test_qwen4_prefill_mask_fuse.py | 842 ++++++++++++++++++++++++++ 10 files changed, 2797 insertions(+), 18 deletions(-) create mode 100644 docs/perf/qwen38-391-remainder.md create mode 100644 mtplx/kernels/qwen4_m4_hyper_read.py create mode 100644 mtplx/qwen4_prefill_chunk.py create mode 100644 tests/test_qwen4_hc_m4.py create mode 100644 tests/test_qwen4_prefill_mask_fuse.py diff --git a/docs/perf/qwen38-391-remainder.md b/docs/perf/qwen38-391-remainder.md new file mode 100644 index 000000000..8fe9011ff --- /dev/null +++ b/docs/perf/qwen38-391-remainder.md @@ -0,0 +1,31 @@ +# Qwen3.8 Flash-Next: the PR 391 remainder onto MTPLX 2.11.2 + +MTPLX 2.11.2 re-landed nine of PR 391's twelve Flash-Next decode keys under the +`MTPLX_QWEN4_*` / `MTPLX_QSA_*` namespace. On the 16K decode cell the release +default reaches 71.17 tok/s against 80.92 for the full 391 stack and 57.65 for +2.10.2 on the same instrument — the nine keys recovered about 60% of 391's +gain. This change ports the remaining lanes so the ~9.75 tok/s difference is +reachable. Each is default-on for a served fixed-M4 Flash-Next pack (stamped in +the `mtplx/server/openai.py` auto-arm block, gated on the fixed-M4 config +predicate) with a per-key `=0` opt-out through the existing pop loop, and each +old `MTPLX_FABLE_*` name is honoured as an alias when the new key is unset. + +All measurements run under the Sustained profile (the release self-selects it +for these packs); the numbers below are from 391's own campaign on the +production geometry and are re-measured in the GPU phase. + +| Lane | Key | What it does | Exactness | 391 recovered | +|------|-----|--------------|-----------|---------------| +| HC_M4 | `MTPLX_QWEN4_HC_M4` | Runs the verify-width (2..8 row) hyper-connection read as one multi-threadgroup GEMV over a kernel-private 8-bit pack instead of the eager norm/down/up chain. | rounding-class | ~-1.1% cycle | +| prefill mask fuse | `MTPLX_QWEN4_PREFILL_MASK_FUSE` | Sends the dense QSA prefill chunk through MLX's fused SDPA (causal string or bool selection) instead of a materialized `[H,S,T]` score tensor MLX's head-dim-256 heuristic otherwise declines. | rounding-class (exact visible set) | part of the prefill share | +| QSA prefill query tile | `MTPLX_QSA_PREFILL_QUERY_TILE` | Tiles only the dense QSA attention query rows so a wide prefill chunk keeps the narrow chunk's attention peak and `sum(rows x context)` cost. | rounding-class (exact visible set) | part of the prefill share | +| QSA sparse split-K decode | `MTPLX_QSA_SPARSE_DECODE` | Native split-K sparse-GQA decode kernel that reads the selected KV rows once instead of materializing a gathered `[1,2,4,2052,256]` K/V pair per QSA layer per verify cycle. | rounding-class | ~-1.46 ms/cycle at 16K | +| graph-build overlap | `MTPLX_QWEN4_GRAPH_BUILD_OVERLAP` | Submits the PLE-independent prefix of the fixed-M4 verify graph early so its GPU work overlaps the ~1.4 ms/cycle host build of the rest. | exact | ~1.4 ms/cycle (~4.5% at N=3) | + +Status in this change: HC_M4, the prefill mask fuse, and the QSA prefill query +tile are ported and CPU-tested. The QSA sparse split-K decode and the graph-build +overlap are covered in the port report (`.benchmark-artifacts/over100-reports/ +remainder-port-report.md`) — the decode lane's substrate is present upstream and +is a native+wiring port whose kernel needs the GPU lock to build and parity-probe; +the graph-build overlap is blocked because upstream re-landed the fixed-M4 verify +as a single compiled graph, without the prefix/suffix split the lane rides. diff --git a/mtplx/kernels/qwen4_m4_hyper_read.py b/mtplx/kernels/qwen4_m4_hyper_read.py new file mode 100644 index 000000000..03f5b7bb5 --- /dev/null +++ b/mtplx/kernels/qwen4_m4_hyper_read.py @@ -0,0 +1,567 @@ +"""Verify-width (rows 2..8) fused hyper-connection READ for Qwen3.8 Flash-Next. + +WHY THIS EXISTS (and why ``hyper_connection.fused_hyper_read`` does not do it) +----------------------------------------------------------------------------- +``GatedResidual.__call__`` is read ~97 times per fixed-M4 verify forward (2 per +layer + the trunk mixer). Eagerly it is 11 dispatches -- grouped rms_norm, the +gamma multiply, the ``[320, 10240]`` down GEMV, silu, the ``[10240, 320]`` up +GEMV, sigmoid, the gated multiply, the hc-mean sum, the mean scale, the +``[4, 10240]`` inject GEMV and its sigmoid -- so 1,067 dispatches/cycle, the +largest zero-byte dispatch family in the compiled verify graph. It is also the +largest non-MoE *byte* consumer: 13.19 MB of bf16 mix weights per read, +1.28 GB/cycle, and the down GEMV measured ~385 GB/s. + +``fused_hyper_read`` (kernels/hyper_connection.py) collapses all 11 into ONE +dispatch with ``grid=(1024, S, 1)``: one threadgroup per row. That shape is a +latency trap at S=4 -- 4 threadgroups on a 40-core GPU, each thread walking +~1.6 KB of weights serially, and every weight element re-read once PER ROW +(4 x 13.19 MB = 52.7 MB at S=4). Measured on the M4 verifier 2026-09-01: +13.2 tok/s with ``MTPLX_FUSED_HC=1`` vs 67.8 control, every GPU phase ~5x +slower because the underfilled kernel backs the queue up. + +This module is the same arithmetic laid out as a *bandwidth-bound GEMV*: +threadgroups tile the OUTPUT COLUMNS, the row dimension R lives in registers, +and each weight element is read EXACTLY ONCE per call regardless of R. + +SHAPE OF THE THREE DISPATCHES +----------------------------- +K0 ``norm`` x[R,10240], gamma[10240] -> normed[R,10240] (bf16) + one threadgroup per (row, hc group); 4R threadgroups of + ``norm_threads``. Materialising ``normed`` once costs 80 KB of + writes and removes the per-threadgroup re-derivation that the + v1/v3 kernels pay (v3 recomputes ``x*wn*rms`` inside every + simdgroup's dot loop -- ~13 MB of L2 traffic at R=1). + +K1 ``down`` normed, wd[320,10240], wi[4,10240] -> mixv[R,320], inject[R,4] + The inject rows are FOLDED IN as virtual output rows 320..323 + (v3's trick), so one kernel streams a [324, 10240] matrix. + ``out_per_tg`` output rows per threadgroup, K split across the + threadgroup's threads: thread t owns k = t, t+NT, t+2NT, ... + Consecutive threads read consecutive weight addresses (fully + coalesced), each thread keeps ``normed[r][k]`` for all R rows in + registers and reuses it for all ``out_per_tg`` weight rows, so + the activation:weight load ratio is R/out_per_tg, not R:1. + wd + wi are read once: 6.63 MB per call at any R. + +K2 ``up`` normed, mixv, wu[10240,320] -> mixed[R,2560] + One simdgroup per (hc group, output d); a threadgroup is + HC=4 simdgroups covering the four hc partners of the same d, so + the hc-mean closes inside threadgroup memory with no second + pass. ``d_per_block`` d's per threadgroup amortises the + per-lane register copy of ``mixv`` (KUP/32 = 10 j's x R). + wu is read once: 6.55 MB per call at any R. + +Total: 3 dispatches (from 11) and 13.19 MB of weight reads (from R x 13.19 MB +under the (1024, S, 1) kernel) -- the DRAM floor for this read, which at the +M5 Max's 614 GB/s ceiling is ~21.5 us/call, ~2.1 ms/cycle over 97 calls. +There is no arrangement of this arithmetic that goes faster, because every +GatedResidual owns private weights and nothing is reused across the 97 calls. + +NUMERICS: WHAT "MATCHES THE EAGER CHAIN" MEANS HERE +--------------------------------------------------- +Target is the EAGER bf16 chain (what the compiled verifier runs today), not +the fused v1/v3 kernels. Every op boundary the eager chain rounds at is +rounded here too, in the same order: + + normed = (T)( (float)(T)(x * rsqrt(ss/D + eps)) * (float)gamma ) + lin = (T)(dot(normed, wd_row)) # nn.Linear output cast + t0 = (T)(lin * 0.25f) # / hc_count (exact in bf16) + mixv = (T)( (float)(T)sigmoid(t0) * t0 ) # nn.silu = x * sigmoid(x) + inject = (T)( 2.0f * (float)(T)sigmoid((T)((T)dot(normed, wi_row)*0.25f)) ) + up = (T)(dot(mixv, wu_row)) + gate = (T)sigmoid(up) + prod = (T)( (float)gate * (float)normed[p] ) + s = prod[0]; for g in 1..3: s = (T)(s + prod[g]) # mx.sum, bf16 acc + mixed = (T)((float)s * 0.25f) # mx.mean scale + +The op-boundary contract above was checked against MLX 0.32.2 on the CPU +stream (all four assertions live in tests/test_qwen4_hc_m4.py): + + * ``bf16 / 4`` and ``2.0 * bf16`` both stay bf16 (weak scalar promotion), + so the ``/ hc_count`` and the inject's ``2.0 *`` round in bf16. + * ``mx.mean(a, axis=-2)`` on a length-4 bf16 axis is bit-identical to + ``(mx.sum(a, -2) * 0.25)``, and that ``mx.sum`` is bit-identical to a + SEQUENTIAL bf16 accumulation -- an fp32 accumulation of the same four + terms differs on 33% of random inputs. Hence the bf16 loop above. + * ``nn.silu(a)`` is exactly ``a * mx.sigmoid(a)``. + +EXPECTED DIFFERENCE CLASS -- rounding only, three named sources: + + 1. GEMV REDUCTION ORDER. MLX's ``gemv_wide_bfloat16_nv4_kl32`` tiles + K=10,240 its own way; this kernel accumulates fp32 per thread over a + strided k subset, then a simd butterfly (``simd_sum``), then a sequential + fp32 walk over the NT/32 simdgroup partials. Both are fp32 accumulations + of the same 10,240 products (MLX's Metal GEMV uses an fp32 AccT for half + types), so they differ only by fp32 reassociation -- a few ulp in fp32, + which then usually rounds to the SAME bf16. The residual is a + 1-ulp-of-bf16 flip on outputs sitting near a rounding boundary. Same + argument for the rms sum of squares. + + 2. SIGMOID FLAVOUR -- the dominant term. MLX's own bf16 ``Sigmoid`` is not + reproducible from any fp32 model: on the CPU backend it differs from + ``bf16(stable-fp32 sigmoid)`` on ~14% of random bf16 inputs, from + ``bf16(naive-fp32)`` on ~14%, and from an fp64 evaluation on the same + ~14% -- i.e. it rounds somewhere inside its own decomposition. This + kernel uses the stable fp32 form (``y = 1/(1+exp(-|x|))`` mirrored for + x<0) and rounds once. Expect O(10%) of the ``mixv``/``gate``/``inject`` + elements to land one bf16 ulp away, which is a relative error of at most + 2^-8 on a value in (0, 1) that then multiplies a normed activation. + + 3. rsqrt FLAVOUR. ``metal::precise::rsqrt`` here against whatever + ``mx.fast.rms_norm`` picked -- a bf16-ulp class on ``normed``. + +So: NOT bit-identical, and not claimed to be. Adoption gates on acceptance +parity on the real verifier (the same bar `_fused_read_applies` was always +going to need), plus the microbench's max-abs-diff / differing-element counts +against the eager module in the hyper-read microbenchmark. + +No silent fallback: ``mtplx.models.qwen4_exp`` raises when MTPLX_QWEN4_HC_M4 +is armed and the module geometry or weight dtypes do not match. +""" + +from __future__ import annotations + +import struct +from functools import lru_cache + +import mlx.core as mx + +HC = 4 +D_HIDDEN = 2560 +HCD = HC * D_HIDDEN # 10240 +R_LOWRANK = 320 +N_FOLDED = R_LOWRANK + HC # 324 virtual stage-1 output rows + +#: Rows this kernel family is wired for. rows==1 keeps the v3 draft path. +MIN_ROWS = 2 +MAX_ROWS = 8 + +DEFAULT_NORM_THREADS = 256 +DEFAULT_OUT_PER_TG = 4 +DEFAULT_D_PER_BLOCK = 8 + +#: Dispatches per read, for the microbench's dispatch-count column. +DISPATCHES_PER_READ = 3 +EAGER_DISPATCHES_PER_READ = 11 + +_HEADER = """ +#include +using namespace metal; + +// MLX's Sigmoid op, in fp32: y = 1/(1+exp(-|x|)) mirrored for x < 0. +inline float mlx_sigmoid_f(float x) { + const float y = 1.0f / (1.0f + metal::exp(-metal::abs(x))); + return (x < 0.0f) ? (1.0f - y) : y; +} +""" + +# -------------------------------------------------------------------------- +# K0 -- grouped RMS norm + gamma, materialised once. +# +# One threadgroup per (row, hc group): grid.x = norm_threads * (R * HC). +# ``EPS_LITERAL`` is substituted per module eps so the kernel needs no extra +# input array (and therefore no extra array to eval) on the hot path. +# -------------------------------------------------------------------------- +_SRC_NORM = """ + constexpr int HC = 4; + constexpr int D = 2560; + constexpr int HCD = HC * D; + constexpr int NT = NTHREADS; + constexpr int NSG = NT / 32; + + const uint blk = threadgroup_position_in_grid.x; + const uint r = blk / (uint)HC; + const uint g = blk % (uint)HC; + const uint tid = thread_position_in_threadgroup.x; + const uint sg = tid / 32; + const uint lane = tid % 32; + + device const T* xg = x + (size_t)r * HCD + (size_t)g * D; + device const T* wg = gamma + (size_t)g * D; + device T* og = normed + (size_t)r * HCD + (size_t)g * D; + + threadgroup float part[NSG]; + + float ss = 0.0f; + for (int i = (int)tid; i < D; i += NT) { + const float v = (float)xg[i]; + ss += v * v; + } + ss = simd_sum(ss); + if (lane == 0) part[sg] = ss; + threadgroup_barrier(mem_flags::mem_threadgroup); + + // Every thread walks the NSG partials in the same order, so every thread + // lands on the identical fp32 scale -- one barrier, and no broadcast slot + // to leave uninitialized. + float tot = 0.0f; + for (int s = 0; s < NSG; ++s) tot += part[s]; + const float sc = metal::precise::rsqrt(tot / (float)D + EPS_LITERAL); + for (int i = (int)tid; i < D; i += NT) { + // mx.fast.rms_norm(grouped, None, eps) rounds to T, then the module + // multiplies by the full-width weight in T. + const float nv = (float)((T)((float)xg[i] * sc)); + og[i] = (T)(nv * (float)wg[i]); + } +""" + +# -------------------------------------------------------------------------- +# K1 -- folded [NTOT, 10240] down/inject GEMV over R rows. +# +# grid.x = down_threads * ceil(NTOT / OUT_PER_TG). Threadgroup ``blk`` owns +# output rows [blk*OUT_PER_TG, +OUT_PER_TG); thread ``t`` owns the k stripe +# {t, t+NT, ...}. Each weight element is loaded once by exactly one thread of +# exactly one threadgroup. +# -------------------------------------------------------------------------- +_SRC_DOWN = """ + constexpr int HC = 4; + constexpr int D = 2560; + constexpr int HCD = HC * D; + constexpr int NDOWN = 320; + constexpr int NTOT = HAS_INJECT ? (NDOWN + HC) : NDOWN; + constexpr int NT = NTHREADS; + constexpr int NSG = NT / 32; + constexpr int OPT = OUT_PER_TG; + constexpr int RR = ROWS; + + const uint blk = threadgroup_position_in_grid.x; + const uint tid = thread_position_in_threadgroup.x; + const uint sg = tid / 32; + const uint lane = tid % 32; + const int o0 = (int)blk * OPT; + + threadgroup float red[NSG * OPT * RR]; + + float acc[OPT][RR]; + #pragma clang loop unroll(full) + for (int o = 0; o < OPT; ++o) { + #pragma clang loop unroll(full) + for (int r = 0; r < RR; ++r) acc[o][r] = 0.0f; + } + + for (int k = (int)tid; k < HCD; k += NT) { + float nv[RR]; + #pragma clang loop unroll(full) + for (int r = 0; r < RR; ++r) nv[r] = (float)normed[(size_t)r * HCD + k]; + #pragma clang loop unroll(full) + for (int o = 0; o < OPT; ++o) { + // Loop-invariant in k: LICM hoists these OPT row pointers into + // registers. Out-of-range rows (only the tail block, and only + // when OPT does not divide NTOT) alias row 0 so the inner loop + // stays branch-free; their partials are dropped in the epilogue. + const int orow = o0 + o; + const int safe = (orow < NTOT) ? orow : 0; + const device T* wrow = (safe < NDOWN) + ? (wd + (size_t)safe * HCD) + : (wi + (size_t)(safe - NDOWN) * HCD); + const float w = (float)wrow[k]; + #pragma clang loop unroll(full) + for (int r = 0; r < RR; ++r) acc[o][r] += nv[r] * w; + } + } + + #pragma clang loop unroll(full) + for (int o = 0; o < OPT; ++o) { + #pragma clang loop unroll(full) + for (int r = 0; r < RR; ++r) { + const float v = simd_sum(acc[o][r]); + if (lane == 0) red[(size_t)sg * (OPT * RR) + o * RR + r] = v; + } + } + threadgroup_barrier(mem_flags::mem_threadgroup); + + constexpr int NPAIR = OPT * RR; + for (int idx = (int)tid; idx < NPAIR; idx += NT) { + const int o = idx / RR; + const int r = idx % RR; + const int orow = o0 + o; + if (orow >= NTOT) continue; + float tot = 0.0f; + for (int s = 0; s < NSG; ++s) tot += red[(size_t)s * NPAIR + idx]; + const float lin = (float)((T)tot); // nn.Linear output + const float t0 = (float)((T)(lin * 0.25f)); // / hc_count + const float s0 = (float)((T)mlx_sigmoid_f(t0)); + if (orow < NDOWN) { + mixv[(size_t)r * NDOWN + orow] = (T)(s0 * t0); // nn.silu + } else { + inject[(size_t)r * HC + (orow - NDOWN)] = (T)(2.0f * s0); + } + } +""" + +# -------------------------------------------------------------------------- +# K2 -- up GEMV + sigmoid gate + hc-mean. +# +# Threadgroup = HC simdgroups (128 threads); simdgroup ``g`` owns hc group g, +# the threadgroup owns d in [blk*DPB, +DPB). Every wu row is read once. +# -------------------------------------------------------------------------- +_SRC_UP = """ + constexpr int HC = 4; + constexpr int D = 2560; + constexpr int HCD = HC * D; + constexpr int KUP = 320; + constexpr int JPL = KUP / 32; // 10 j's per lane + constexpr int RR = ROWS; + constexpr int DPB = D_PER_BLOCK; + constexpr int NT = HC * 32; + + const uint blk = threadgroup_position_in_grid.x; + const uint tid = thread_position_in_threadgroup.x; + const uint sg = tid / 32; // == hc group + const uint lane = tid % 32; + const int d0 = (int)blk * DPB; + + threadgroup float prodv[HC * DPB * RR]; + + // Per-lane register copy of the R mixv rows: reused for all DPB d's. + float m[RR][JPL]; + #pragma clang loop unroll(full) + for (int r = 0; r < RR; ++r) { + #pragma clang loop unroll(full) + for (int jj = 0; jj < JPL; ++jj) { + m[r][jj] = (float)mixv[(size_t)r * KUP + (int)lane + jj * 32]; + } + } + + for (int dd = 0; dd < DPB; ++dd) { + const int d = d0 + dd; + if (d >= D) break; + const int p = (int)sg * D + d; + const device T* wrow = wu + (size_t)p * KUP; + float a[RR]; + #pragma clang loop unroll(full) + for (int r = 0; r < RR; ++r) a[r] = 0.0f; + #pragma clang loop unroll(full) + for (int jj = 0; jj < JPL; ++jj) { + const float w = (float)wrow[(int)lane + jj * 32]; + #pragma clang loop unroll(full) + for (int r = 0; r < RR; ++r) a[r] += m[r][jj] * w; + } + #pragma clang loop unroll(full) + for (int r = 0; r < RR; ++r) { + const float tot = simd_sum(a[r]); // uniform: all lanes + if ((int)lane == r) { + const float lin = (float)((T)tot); // nn.Linear output + const float gate = (float)((T)mlx_sigmoid_f(lin)); + const float nv = (float)normed[(size_t)r * HCD + p]; + prodv[(size_t)sg * (DPB * RR) + dd * RR + r] = + (float)((T)(gate * nv)); + } + } + } + threadgroup_barrier(mem_flags::mem_threadgroup); + + for (int idx = (int)tid; idx < DPB * RR; idx += NT) { + const int dd = idx / RR; + const int r = idx % RR; + const int d = d0 + dd; + if (d >= D) continue; + // mx.mean(..., axis=-2) is sum-then-scale, and MLX's col_reduce + // instantiates T == U, so the 4-term sum accumulates IN bf16 and + // rounds at every add, in axis order. Verified against the CPU + // backend (mx.sum over a length-4 bf16 axis is bit-identical to a + // sequential bf16 accumulation; an fp32 accumulation differs on 33% + // of random inputs, so this ordering is not cosmetic). + float s = prodv[idx]; // g == 0 + for (int g = 1; g < HC; ++g) { + s = (float)((T)(s + prodv[(size_t)g * (DPB * RR) + idx])); + } + mixed[(size_t)r * D + d] = (T)(s * 0.25f); // exact: 2^-2 + } +""" + + +def _eps_tag(eps: float) -> str: + """Stable short tag for an eps value, so a kernel name never collides + across two different eps (MLX caches compiled libraries by name).""" + + return struct.pack(" int: + return struct.unpack(" int: + """Validate the family contract and return R. Raises -- never returns + False -- so an armed flag on a mismatched pack fails at the call site + instead of silently reverting to the eager chain. + + ``x`` may carry any leading dims (the module sees ``[B, S, 10240]``); R is + their product. Taking the unreshaped array keeps this off the graph: no + reshape node is created just to validate. + """ + + if x.ndim < 1 or x.shape[-1] != HCD: + raise ValueError( + f"hyper input must be [..., {HCD}] for the M4 hyper read; got " + f"{tuple(x.shape)}" + ) + rows = 1 + for s in x.shape[:-1]: + rows *= int(s) + if not (rows_min <= rows <= rows_max): + raise ValueError( + f"M4 hyper read is wired for {rows_min}..{rows_max} rows; got {rows}" + ) + check_weight_shapes(gamma, wd, wu, wi, dtype=x.dtype) + return rows + + +def check_weight_shapes(gamma, wd, wu, wi, *, dtype=None) -> None: + """Validate the WEIGHT half of the family contract. + + Split out of :func:`check_shapes` so the same list can run at INSTALL + time, where the weights exist but no activation does: a pack the kernel + cannot read is a deployment error, and it should stop the server coming up + rather than fail the first request that happens to reach verify width. + ``dtype`` is the activation dtype when there is one to compare against; + ``None`` skips only that check. One list, two callers. + """ + + checks = ( + ("hc_norm.weight", gamma, (HCD,)), + ("input_mix_weight_down.weight", wd, (R_LOWRANK, HCD)), + ("input_mix_weight_up.weight", wu, (HCD, R_LOWRANK)), + ) + if wi is not None: + checks = checks + (("block_inject_weight.weight", wi, (HC, HCD)),) + for name, arr, want in checks: + if tuple(arr.shape) != want: + raise ValueError( + f"M4 hyper read: {name} must be {want}; got {tuple(arr.shape)}" + ) + if dtype is not None and arr.dtype != dtype: + raise ValueError( + f"M4 hyper read: {name} dtype {arr.dtype} != hyper input dtype " + f"{dtype} (the kernel reads the module's unquantized weights)" + ) + + +def fused_hc_read_m4( + x2d, + gamma, + wd, + wu, + wi=None, + *, + eps: float = 1e-6, + norm_threads: int = DEFAULT_NORM_THREADS, + down_threads: int = DEFAULT_NORM_THREADS, + out_per_tg: int = DEFAULT_OUT_PER_TG, + d_per_block: int = DEFAULT_D_PER_BLOCK, +): + """Fused GatedResidual read for R rows. + + ``x2d`` [R, 10240], ``gamma`` [10240], ``wd`` [320, 10240], + ``wu`` [10240, 320], ``wi`` [4, 10240] or None (the no-combine trunk + mixer). Returns ``(mixed [R, 2560], inject [R, 4] or None)``. + """ + + if x2d.ndim != 2: + raise ValueError( + f"fused_hc_read_m4 wants a 2-D [R, {HCD}] view; got {tuple(x2d.shape)}" + ) + rows = check_shapes(x2d, gamma, wd, wu, wi) + if norm_threads % 32 or not (32 <= norm_threads <= 1024): + raise ValueError(f"norm_threads={norm_threads}: want a multiple of 32 in [32,1024]") + if down_threads % 32 or not (32 <= down_threads <= 1024): + raise ValueError(f"down_threads={down_threads}: want a multiple of 32 in [32,1024]") + if out_per_tg < 1 or out_per_tg > 32: + raise ValueError("out_per_tg must be in 1..32") + if d_per_block < 1 or d_per_block > 64: + raise ValueError("d_per_block must be in 1..64") + + dt = x2d.dtype + has_inject = wi is not None + n_tot = N_FOLDED if has_inject else R_LOWRANK + + (normed,) = _kernel_norm(_eps_bits(eps))( + inputs=[x2d, gamma], + template=[("T", dt), ("NTHREADS", norm_threads)], + grid=(norm_threads * rows * HC, 1, 1), + threadgroup=(norm_threads, 1, 1), + output_shapes=[(rows, HCD)], + output_dtypes=[dt], + ) + + n_blk = (n_tot + out_per_tg - 1) // out_per_tg + mixv, inject = _kernel_down()( + inputs=[normed, wd, wi if has_inject else wd], + template=[ + ("T", dt), + ("ROWS", rows), + ("NTHREADS", down_threads), + ("OUT_PER_TG", out_per_tg), + ("HAS_INJECT", 1 if has_inject else 0), + ], + grid=(down_threads * n_blk, 1, 1), + threadgroup=(down_threads, 1, 1), + output_shapes=[(rows, R_LOWRANK), (rows, HC)], + output_dtypes=[dt, dt], + ) + + d_blk = (D_HIDDEN + d_per_block - 1) // d_per_block + (mixed,) = _kernel_up()( + inputs=[normed, mixv, wu], + template=[("T", dt), ("ROWS", rows), ("D_PER_BLOCK", d_per_block)], + grid=(HC * 32 * d_blk, 1, 1), + threadgroup=(HC * 32, 1, 1), + output_shapes=[(rows, D_HIDDEN)], + output_dtypes=[dt], + ) + return mixed, (inject if has_inject else None) + + +def weight_bytes_per_read(dtype_size: int = 2, *, has_inject: bool = True) -> int: + """Weight bytes this read streams, independent of R -- the number the + GB/s column in the hyper-read microbenchmark divides by.""" + + n = R_LOWRANK * HCD + HCD * R_LOWRANK + HCD # down + up + gamma + if has_inject: + n += HC * HCD + return n * dtype_size diff --git a/mtplx/models/qwen4_exp.py b/mtplx/models/qwen4_exp.py index e8e0e04fe..9623a8ad2 100644 --- a/mtplx/models/qwen4_exp.py +++ b/mtplx/models/qwen4_exp.py @@ -44,6 +44,7 @@ import struct import time from dataclasses import dataclass, field +from functools import lru_cache from pathlib import Path from typing import Any, Dict, List, Optional, Union @@ -61,7 +62,11 @@ ) from mtplx.attention_context import current_attention_phase -from mtplx.runtime_options import qwen4_opdiet_enabled, qwen4_verify_glue_enabled +from mtplx.runtime_options import ( + qwen4_hc_m4_enabled, + qwen4_opdiet_enabled, + qwen4_verify_glue_enabled, +) @dataclass @@ -861,6 +866,111 @@ def __init__(self, args: TextArgs, use_combine: bool = True): self.input_mix_weight_up = nn.Linear(args.hc_lowrank, hc_hidden, bias=False) if use_combine: self.block_inject_weight = nn.Linear(hc_hidden, self.hc_count, bias=False) + # Construction-time half of the MTPLX_QWEN4_HC_M4 eligibility check: + # the config-level geometry the kernel hardcodes is knowable now, so a + # mis-armed flag fails at model build rather than mid-forward. The + # weight-level half (dtype, quantization, shapes) needs loaded weights + # and runs on the first verify-width read. + if qwen4_hc_m4_enabled() and ( + self.hc_count != 4 or self.hidden_size != 2560 + ): + raise RuntimeError( + "MTPLX_QWEN4_HC_M4 is armed but this GatedResidual is not the " + f"Flash-Next family shape: hc_count={self.hc_count} (want 4), " + f"hidden_size={self.hidden_size} (want 2560). Unset the flag " + "for this model; the kernel has no other geometry." + ) + + def validate_hc_m4_pack(self, label: str = "GatedResidual") -> None: + """The PACK half of the MTPLX_QWEN4_HC_M4 contract, at install time. + + Everything here is a property of the loaded weights, so it is the same + answer for every request this process will ever serve. Checking it + when the weights land (``install_hc_m4_pack_validation``, called from + the runtime's qwen4 install section) means a mis-armed flag stops the + server coming up with a precise reason, instead of turning the first + request that reaches verify width into an HTTP 500. + + Not checked here: the weight/activation dtype agreement, which needs + an activation -- ``_hc_m4_applies`` still checks it, and it is likewise + process-invariant, so it cannot single out one request either. + """ + + from mtplx.kernels import qwen4_m4_hyper_read as hcm4 + + down = self.input_mix_weight_down + up = self.input_mix_weight_up + for name, proj in ( + ("input_mix_weight_down", down), + ("input_mix_weight_up", up), + ): + if hasattr(proj, "scales"): + raise RuntimeError( + f"MTPLX_QWEN4_HC_M4: {label}.{name} is quantized; the " + "kernel reads unquantized bf16 mix weights. Unset the " + "flag for this pack." + ) + wi = self.block_inject_weight.weight if "block_inject_weight" in self else None + try: + hcm4.check_weight_shapes( + self.hc_norm.weight, down.weight, up.weight, wi + ) + except ValueError as exc: + raise RuntimeError(f"MTPLX_QWEN4_HC_M4: {label}: {exc}") from exc + + def _hc_m4_applies(self, hyper_input: mx.array) -> bool: + """Verify-width fused read gate (MTPLX_QWEN4_HC_M4). + + Rows 2..8 only -- rows == 1 keeps whatever the draft path already + uses (v3, or the eager chain). The row count is the one REQUEST-shaped + term here and it routes (returns False) rather than raising; every + other term is a property of the pack, re-checked here but already + settled at install by :meth:`validate_hc_m4_pack`. + """ + + if not qwen4_hc_m4_enabled(): + return False + from mtplx.kernels import qwen4_m4_hyper_read as hcm4 + + rows = 1 + for s in hyper_input.shape[:-1]: + rows *= s + if rows < hcm4.MIN_ROWS or rows > hcm4.MAX_ROWS: + return False + down = self.input_mix_weight_down + up = self.input_mix_weight_up + for name, proj in (("input_mix_weight_down", down), ("input_mix_weight_up", up)): + if hasattr(proj, "scales"): + raise RuntimeError( + f"MTPLX_QWEN4_HC_M4: {name} is quantized; the kernel reads " + "unquantized bf16 mix weights. Unset the flag for this pack." + ) + wi = self.block_inject_weight.weight if "block_inject_weight" in self else None + # Raises with the offending shape/dtype named. Takes the unreshaped + # hyper state so validation adds no node to the traced graph. + hcm4.check_shapes( + hyper_input, self.hc_norm.weight, down.weight, up.weight, wi + ) + return True + + def _hc_m4_read(self, hyper_input: mx.array): + from mtplx.kernels.qwen4_m4_hyper_read import fused_hc_read_m4 + + combine = "block_inject_weight" in self + x2 = hyper_input.reshape(-1, self.hc_count * self.hidden_size) + mixed, inject = fused_hc_read_m4( + x2, + self.hc_norm.weight, + self.input_mix_weight_down.weight, + self.input_mix_weight_up.weight, + self.block_inject_weight.weight if combine else None, + eps=float(self.hc_norm.eps), + ) + mixed = mixed.reshape(*hyper_input.shape[:-1], self.hidden_size) + if not combine: + return mixed + inject = inject.reshape(*hyper_input.shape[:-1], self.hc_count) + return mixed, hyper_input, inject def _fused_read_applies(self, hyper_input: mx.array) -> bool: # The fused kernel hardcodes the family geometry and reads bf16 @@ -910,6 +1020,12 @@ def _v3_read_applies(self, hyper_input: mx.array) -> bool: return True def __call__(self, hyper_input: mx.array): + # Verify-width (2..8 rows) fused read first: it is the only one of the + # three fused paths laid out as a multi-threadgroup GEMV, and it is + # gated to widths the others handle badly (the (1024, S, 1) v1 kernel + # re-reads every weight once per row and runs 4 threadgroups at S=4). + if self._hc_m4_applies(hyper_input): + return self._hc_m4_read(hyper_input) if self._v3_read_applies(hyper_input): from mtplx.kernels.hyper_connection_v3 import fused_hyper_read_v3 @@ -950,18 +1066,72 @@ def __call__(self, hyper_input: mx.array): -def _named_gated_residuals(owner: Any): - """``(attribute name, module)`` for every GatedResidual ``owner`` holds.""" +def install_hc_m4_pack_validation(model: Any) -> dict[str, Any]: + """Validate every ``GatedResidual`` against the MTPLX_QWEN4_HC_M4 contract. - for name in dir(owner): - if name.startswith("__"): - continue - try: - value = getattr(owner, name) - except Exception: # pragma: no cover - defensive: properties may raise - continue - if isinstance(value, GatedResidual): - yield name, value + Called from the runtime's qwen4 install section once the weights are + loaded. A no-op (``{"armed": False}``) when the flag is off, so an + unarmed process pays one attribute read. + + This is the INSTALL-time half of the flag's contract. Every property it + checks -- quantized mix weights, weight shapes -- belongs to the pack, so + a failure means the flag was armed for a model the kernel cannot read. + That is a deployment error and it should stop the server, not fail + whichever request first reaches verify width. + + Discovery goes through :func:`_named_gated_residuals`, which walks the + model the way MLX itself does. The first cut of this function walked + ``dir(layer)`` instead and found NOTHING on the real pack -- an + ``nn.Module``'s children live in its dict and are served by + ``__getattr__``, so ``dir()`` does not list them -- which turned an armed + flag into a dead server (2026-09-02, both the HumanEval screen and the + ABBA lane). Never enumerate an MLX module with ``dir()``. + """ + + if not qwen4_hc_m4_enabled(): + return {"armed": False, "validated": 0} + + validated = 0 + for path, module in _named_gated_residuals(model): + module.validate_hc_m4_pack(path) + validated += 1 + if not validated: + raise RuntimeError( + "MTPLX_QWEN4_HC_M4 is armed but no GatedResidual hyper-connection " + "module was found anywhere in " + f"{type(model).__name__}.named_modules(); the flag cannot do " + "anything here. Unset it for this model." + ) + return {"armed": True, "validated": validated} + + +def _named_gated_residuals(model: Any): + """``(dotted path, module)`` for every GatedResidual in ``model``. + + ``nn.Module.named_modules`` is the model's OWN traversal -- the one that + finds its parameters -- so it reaches children held directly + (``model.hyper_connection_mixer``), inside lists (``model.layers[i] + .attn_hyper_connection``) and behind a published sub-tree + (``language_model.mtp.hyper_connection_mixer``) alike. If the forward can + reach a module, this finds it. + + Sorted by path so the failure message names layers in a stable order + rather than MLX's dict order. + """ + + named = getattr(model, "named_modules", None) + if named is None: # pragma: no cover - every qwen4 model is an nn.Module + raise RuntimeError( + "MTPLX_QWEN4_HC_M4 pack validation needs an mlx.nn.Module; got " + f"{type(model).__name__}" + ) + found = [ + (path, module) + for path, module in named() + if isinstance(module, GatedResidual) + ] + found.sort(key=lambda item: item[0]) + return found class SparseMoeBlock(_Qwen3NextSparseMoeBlock): @@ -1747,6 +1917,411 @@ def _qsa_prefill_gather_tile_rows() -> int: return 64 +@lru_cache(maxsize=1) +def _prefill_mask_fuse_enabled() -> bool: + """MTPLX_QWEN4_PREFILL_MASK_FUSE: ask MLX for the fused masked SDPA. + + (Old name ``MTPLX_FABLE_PREFILL_MASK_FUSE`` still works as an alias.) + + + At ``head_dim`` 256 MLX's own heuristic + (``ScaledDotProductAttention::use_fallback``) refuses the fused steel + attention kernel and takes the unfused route: QK^T into a materialized + ``[H, S, T]`` bf16 tensor, ``mx.where`` against the bool mask, softmax, + then P@V. The census confirms it -- ``steel_attention`` appears **zero** + times in 2.1 M dispatches, while ``g2_Selectbfloat16`` runs at grids + ``2048 x T`` for every QSA layer of every chunk (337 ms of pure mask + apply at 687 GB/s, on a 1.61 GB transient at the last chunk). + + There is nothing to fuse by hand: the unfused route is MLX's C++ + fallback lambda, and swapping the bool mask for an additive one only + turns the ``where`` into an ``add`` over the same bytes. What MLX 0.32 + added instead is ``force_fused=True``, whose own docs say it "would + result in slower kernel getting used but can reduce memory + consumption" -- and the shipped metallib does carry + ``steel_attention_bfloat16_bq32_bk16_bd256_wm4_wn1_maskbool_``. So the + exact-visible-set fused kernel exists at this geometry; only the + heuristic declines it. + + Two arms ride this one flag, both fused, chosen by what the indexer + returned: + + * **causal** -- no selection came back, so every key a row can see IS + visible. Nothing is built: the string ``"causal"`` goes to MLX, whose + lower-right alignment matches this lane's ``[1, 1, S, T]`` mask exactly + when ``T == pos_start + S`` (see ``_prefill_causal_mask_is_exact``). + This is the regime below the indexer's own budget -- + ``T <= (block_topk + 1) * ratio - 1``, i.e. 2,051 tokens on the + production pack -- plus vision requests, where QSA is bypassed. With + the retained 4,096-token prefill width it therefore fires for NO chunk + of a 16K or 32K prompt; the win at those cells is entirely the bool + arm below. + * **bool** -- a real top-k selection, which no string can express; the + array is handed to the fused kernel untouched. + + Off by default: it trades a materialized score tensor (and its mask + apply and softmax passes) for a flash kernel at a head dimension MLX + considers unfavourable. It is also NOT bit-identical -- online softmax + reassociates the same visible set -- so it is a quality-gated arm. + + Counters (``MTPLX_QSA_PREFILL_DEBUG=1``): ``mask_causal_eligible`` (the + lane saw an exactly-causal visible set, flag-independent), + ``mask_fuse_causal``, ``mask_fuse_bool``, ``mask_fuse_unavailable`` + (one per SHAPE CLASS this MLX has no fused kernel for -- see + :data:`_PREFILL_MASK_FUSE_UNAVAILABLE`; it is NOT a process-wide + disarm), and ``mask_fuse_dense_causal`` / ``mask_fuse_dense_bool`` + (calls sent to the dense route because their own class was refused, + while every other class stays fused). A serving process runs without + that debug flag, so the first class MLX fuses also prints one + ``engaged:`` line -- the absence of a refusal is not a receipt. + """ + + # Cached (maxsize=1) so the hot lane does not touch os.environ every + # call; ``cache_clear()`` is the escape hatch a test or an in-process A/B + # uses to re-read. The renamed key wins when set to any non-empty value + # (including "0" for the per-key opt-out); the old + # MTPLX_FABLE_PREFILL_MASK_FUSE name is honoured as an alias only when + # the new key is unset. + raw = os.environ.get("MTPLX_QWEN4_PREFILL_MASK_FUSE") + if raw is None or not str(raw).strip(): + raw = os.environ.get("MTPLX_FABLE_PREFILL_MASK_FUSE") + raw = (raw or "0").strip().lower() + return raw in {"1", "true", "yes", "on"} + + +#: MLX's only string mask mode. ``mx.fast.scaled_dot_product_attention`` +#: documents the alignment in MLX 0.32.2 itself (``mlx/core/fast.pyi``): +#: *"The ``causal`` mask uses lower-right alignment where the last query +#: aligns with the last key."* That is exactly the offset case this lane +#: needs -- chunk k carries ``pos_start`` keys of prior context and ``S`` +#: queries, so its last query IS its last key -- which is why a chunked +#: prefill can pass the string instead of a ``[1, 1, S, T]`` tensor. +_CAUSAL_MASK = "causal" + + +def _prefill_causal_mask_is_exact( + *, pos_start: int, rows: int, total_keys: int +) -> bool: + """Is the dense mask this lane would build exactly MLX's ``"causal"``? + + The lane builds ``tpos[None, :] <= qpos[:, None]`` over + ``qpos = pos_start + arange(S)`` and ``tpos = arange(T)``: query row + ``i`` sees keys ``0 .. pos_start + i``. MLX's lower-right ``"causal"`` + gives row ``i`` keys ``0 .. T - S + i``. The two agree for every row + iff ``T == pos_start + S`` -- i.e. iff the KV the cache just handed back + ends at this chunk's last query. + + That is the normal case (``T`` comes straight out of + ``cache.kv.update_and_fetch``), but a cache that returned padded or + capacity-shaped keys would break it, and there the dense mask -- which + masks the pad columns off -- is the correct one. Checked, never assumed. + """ + + return int(rows) > 0 and int(total_keys) == int(pos_start) + int(rows) + + +#: Fused-SDPA capability, keyed by SHAPE CLASS -- never by mask kind alone. +#: +#: MLX admits the fused route through two different kernels with two +#: different rule sets, and it picks between them on the QUERY LENGTH +#: (``ScaledDotProductAttention::use_fallback``, +#: ``mlx/backend/metal/scaled_dot_product_attention.cpp``; the installed +#: 0.32.2 dylib carries exactly these refusal reasons as strings, with the +#: head-dim sets below): +#: +#: * the **full** kernel (``steel_attention``) takes the long query -- +#: ``query_sequence_length > 8`` -- with a head dim in +#: ``{64, 72, 80, 96, 128, 192, 256}`` shared by q and v and, when the +#: mask is the causal string, ``q_len <= k_len``; +#: * the **vector** kernel (``sdpa_vector``) takes the short query -- +#: ``q_len <= 8`` -- with a head dim in ``{64, 96, 128, 192, 256}``, +#: ``q_len <= k_len``, AND ``q_len * gqa_factor <= 32``. +#: +#: (The ``> 8`` split is read off MLX's own source; what this key relies on +#: is only that SOME such split exists, which the served refusal proves -- +#: MLX answered a 5-row query with the VECTOR kernel's reason at a head dim +#: the full kernel accepts. ``q_len`` is therefore kept exact below, so the +#: constant itself is never baked in.) +#: +#: This model is 24 q heads over 2 kv heads (GQA 12) at ``head_dim`` 256, +#: which leaves a DEAD BAND at ``q_len`` 3..8: too long for the vector +#: kernel (``3 * 12 = 36 > 32``), too short for the full one. A 4,096-row +#: prefill chunk is far above the band and IS fused; an MTP verify step (4 +#: rows, 5 with the extra) sits inside it and is not. Keying availability +#: by mask kind alone therefore let the FIRST verify step of the process +#: disarm the flag for every prefill chunk after it -- which is exactly +#: what the served process did (warmup ladder: verify before any wide +#: chunk) while the benchmark driver, whose first call IS a wide chunk, +#: kept the win. So the capability is keyed by the geometry MLX's own +#: rules read, and a negative learned at one class says nothing about any +#: other. +_PREFILL_MASK_FUSE_UNAVAILABLE: Dict[tuple, bool] = {} + +#: Cap on the class table. Distinct classes are structurally few (mask +#: kind x chunk width), but a build with no fused kernel at all would learn +#: one negative per distinct prefill width forever; evicting the oldest +#: costs one more host-side refusal, never a wrong answer. +_MASK_FUSE_CLASS_CACHE_MAX = 1024 + +#: stderr is a shared log, and a build that refuses everything would +#: otherwise print one line per prefill width for the life of the process. +_MASK_FUSE_REFUSAL_PRINT_LIMIT = 8 +_MASK_FUSE_REFUSALS_PRINTED = [0] + +#: One positive receipt per process, printed by the first class that MLX +#: actually fuses. ABSENCE of a refusal line is not evidence of engagement: +#: the served process printed two refusals (the 4- and 5-row verify steps) +#: and nothing at all for its chunks, which is exactly the log a process +#: with the flag unset writes. A serving process has no +#: ``MTPLX_QSA_PREFILL_DEBUG`` receipt to fall back on, so it needs this. +_MASK_FUSE_ENGAGED = [False] + + +def _prefill_mask_fuse_kind(mask) -> str: + """``"causal"`` for the string mode, ``"bool"`` for a real selection.""" + + return "causal" if isinstance(mask, str) else "bool" + + +def _prefill_mask_fuse_class(kind: str, q, k, v) -> tuple: + """The geometry MLX's fused-SDPA admission rules actually read. + + Everything in the tuple appears in one of MLX's own refusal reasons: + the query and value head dims, the query length (which selects the + kernel and, in the vector kernel, is multiplied by the GQA factor + against its 32 cap), the GQA factor, and whether the query is no longer + than the key sequence. ``dtype`` rides along because the fused kernels + are per-dtype specialisations. Two calls with the same class get the + same answer from MLX by construction -- which is what makes caching one + call's refusal safe for the others in it, and only for those. + """ + + kv_heads = max(int(k.shape[1]), 1) + q_len = int(q.shape[2]) + return ( + kind, + int(q.shape[-1]), + int(v.shape[-1]), + str(q.dtype).rsplit(".", 1)[-1], + q_len, + int(q.shape[1]) // kv_heads, + q_len <= int(k.shape[2]), + ) + + +def _prefill_mask_fuse_class_text(cls: tuple) -> str: + """The class as a log line reads it.""" + + kind, q_head_dim, v_head_dim, dtype, q_len, gqa, q_fits = cls + head_dim = ( + f"head_dim {q_head_dim}" + if q_head_dim == v_head_dim + else f"head_dim {q_head_dim} (value {v_head_dim})" + ) + return ( + f"{kind}-mask q_len {q_len} x GQA {gqa} at {head_dim} {dtype}" + + ("" if q_fits else ", query longer than the key sequence") + ) + + +def _prefill_mask_fuse_announce(cls: tuple) -> None: + """Say ONCE, on the first fused call, which shape class engaged.""" + + _MASK_FUSE_ENGAGED[0] = True + import sys as _sys + + print( + "[mtplx] MTPLX_QWEN4_PREFILL_MASK_FUSE engaged: fused SDPA for " + f"shape class [{_prefill_mask_fuse_class_text(cls)}]; every other " + "shape class is asked on its first call and reported here only if " + "it is refused", + file=_sys.stderr, + flush=True, + ) + + +def _prefill_mask_fuse_refuse(cls: tuple, exc: BaseException) -> None: + """Route ONE shape class to the dense path for the rest of the process. + + Loud, because an armed flag that silently measured the control under a + candidate label is the failure this receipt exists to prevent -- but + scoped, because the class that refused is usually not the class the + lane is armed for. + """ + + if len(_PREFILL_MASK_FUSE_UNAVAILABLE) >= _MASK_FUSE_CLASS_CACHE_MAX: + _PREFILL_MASK_FUSE_UNAVAILABLE.pop( + next(iter(_PREFILL_MASK_FUSE_UNAVAILABLE)), None + ) + _PREFILL_MASK_FUSE_UNAVAILABLE[cls] = True + _qsa_prefill_count("mask_fuse_unavailable") + if _MASK_FUSE_REFUSALS_PRINTED[0] >= _MASK_FUSE_REFUSAL_PRINT_LIMIT: + return + _MASK_FUSE_REFUSALS_PRINTED[0] += 1 + import sys as _sys + + tail = ( + "; further shape-class refusals are counted but not printed" + if _MASK_FUSE_REFUSALS_PRINTED[0] == _MASK_FUSE_REFUSAL_PRINT_LIMIT + else "" + ) + print( + "[mtplx] MTPLX_QWEN4_PREFILL_MASK_FUSE armed but this MLX has no " + "fused SDPA for shape class " + f"[{_prefill_mask_fuse_class_text(cls)}]; falling back to the " + "dense/unfused route for THAT shape class only -- per-class, NOT " + "process-wide: every other shape class (a wide prefill chunk " + f"included) is still asked and still fused{tail}: {exc}", + file=_sys.stderr, + flush=True, + ) + + +def _prefill_mask_fuse_sdpa(q, k, v, *, scale, mask): + """Masked SDPA, fused when armed and available, else stock. + + Two arms behind one door: + + * ``mask is _CAUSAL_MASK`` -- the visible set is exactly causal (proved + host-side by :func:`_prefill_causal_mask_is_exact`), so no tensor is + built at all and MLX generates the mask inside the kernel. + * ``mask`` is a bool array -- the real QSA block selection, which no + string can express; the fused kernel reads it directly + (``steel_attention_..._maskbool_``). + + **Exactness.** Neither arm is bit-identical to the unfused route, and + both are the same rounding class. The dense path materialises bf16 + ``QK^T``, applies the mask with ``mx.where``, runs a precise two-pass + softmax and contracts ``P@V`` in bf16; the fused kernel streams tiles + through an fp32 online softmax and accumulates in fp32. What is + identical is the VISIBLE SET -- causal-string and dense-causal mask + admit the same keys per row by construction, and the bool arm passes the + selection through untouched -- so this is reassociation of the same sum, + not an approximation, and it is gated by the agreement screen rather + than by a bit-parity assert. + """ + + kind = _prefill_mask_fuse_kind(mask) + if ( + mask is not None + # Prefill only. At S == 1 MLX already fuses (head_dim 256 IS in the + # sdpa_VECTOR supported set); the fallback this flag exists to + # replace is the S > 1 one. + and int(q.shape[2]) > 1 + and _prefill_mask_fuse_enabled() + ): + # The capability question is asked AT THE CALL, with the call's own + # operands -- there is no synthetic stand-in that can be wrong about + # the geometry, because it IS the geometry. ``force_fused=True`` is + # resolved while the op is BUILT (no encoder is opened, nothing is + # evaluated), so a class MLX refuses costs one host-side raise, once, + # and never a dispatch. The answer is then bound to that shape class + # alone: the hot lane never retries a class it has been refused, and + # a refusal at one class never touches another. + cls = _prefill_mask_fuse_class(kind, q, k, v) + if not _PREFILL_MASK_FUSE_UNAVAILABLE.get(cls, False): + try: + out = mx.fast.scaled_dot_product_attention( + q, k, v, scale=scale, mask=mask, force_fused=True + ) + except Exception as exc: # no fused kernel for THIS shape class + _prefill_mask_fuse_refuse(cls, exc) + else: + if not _MASK_FUSE_ENGAGED[0]: + _prefill_mask_fuse_announce(cls) + _qsa_prefill_count( + "mask_fuse_causal" if kind == "causal" else "mask_fuse_bool" + ) + return out + # Armed, and this call still went dense: its own shape class has no + # fused kernel in this MLX. Counted per call so a receipt can tell + # "the flag never fired" from "the flag fired for the chunks and not + # for the 4-row verify steps". + _qsa_prefill_count( + "mask_fuse_dense_causal" if kind == "causal" else "mask_fuse_dense_bool" + ) + return mx.fast.scaled_dot_product_attention(q, k, v, scale=scale, mask=mask) + + +def _prefill_qsa_query_tile_rows() -> int: + """MTPLX_QSA_PREFILL_QUERY_TILE: rows per attention query tile. + + (Old name ``MTPLX_FABLE_PREFILL_QSA_QUERY_TILE`` still works as an alias.) + + The middle path for a wide prefill chunk. The 36 GDN layers, the MoE + grouped GEMM and every projection want a WIDE chunk (better grouped-GEMM + tile occupancy, fewer per-chunk syncs); the 12 dense QSA layers want a + NARROW one (their score tensor is ``[H, rows, T]`` and their work term + ``rows x T`` grows with the chunk). Splitting only the attention into + query tiles gives both: tile A never reads tile B's keys, so a 4,096-row + chunk tiled at 2,048 has exactly the peak AND exactly the + ``sum(rows x context)`` of an 8 x 2,048 cut, while everything outside + attention still sees 4,096 rows. + + 0 (default) = whole chunk, i.e. today's behaviour. + """ + + from mtplx.qwen4_prefill_chunk import resolve_query_tile_rows + + return resolve_query_tile_rows() + + +def _qsa_dense_attention(q, k, v, *, mask, scale): + """Dense masked SDPA, optionally split into query tiles. + + Unarmed (the default) this is one ``_prefill_mask_fuse_sdpa`` call -- + byte-identical to the code it replaced when the mask-fuse flag is also + off. + + Armed, it cuts the chunk's query rows into tiles and truncates each + tile's keys to that tile's last visible position. Rows are independent + under attention, and the keys dropped were exactly the ones the mask + already set to ``finfo.min`` (whose ``exp`` is a hard zero), so the + visible set per row is unchanged. Reduction order is not: shorter + softmax rows and a shorter P@V contraction, the same class of difference + as the portable gather tier. + + ``mask`` may also be :data:`_CAUSAL_MASK`, the string mode. Tiling + composes with it without any slicing: ``query_tile_spans`` sets a tile's + key bound to ``context_before + row_end``, so the tile's own last query + is again its last key and MLX's lower-right ``"causal"`` describes the + sub-problem exactly. The string is therefore passed through unchanged + while the K/V slices narrow, which is the whole point -- there is no + ``[1, 1, S, T]`` tensor to slice in the first place. + """ + + S = int(q.shape[2]) + tile = _prefill_qsa_query_tile_rows() if S > 1 else 0 + total_keys = int(k.shape[2]) + causal_string = isinstance(mask, str) + if ( + tile <= 0 + or tile >= S + or mask is None + or total_keys < S + or (not causal_string and int(mask.shape[-1]) != total_keys) + ): + return _prefill_mask_fuse_sdpa(q, k, v, scale=scale, mask=mask) + + from mtplx.qwen4_prefill_chunk import query_tile_spans + + spans = query_tile_spans(S, context_before=total_keys - S, tile=tile) + if not spans: + return _prefill_mask_fuse_sdpa(q, k, v, scale=scale, mask=mask) + _qsa_prefill_count("query_tile") + parts = [ + _prefill_mask_fuse_sdpa( + q[:, :, r0:r1], + k[:, :, :keys], + v[:, :, :keys], + scale=scale, + mask=mask if causal_string else mask[..., r0:r1, :keys], + ) + for r0, r1, keys in spans + ] + return mx.concatenate(parts, axis=2) + + def _fused_hc_enabled() -> bool: raw = (os.environ.get("MTPLX_FUSED_HC") or "0").strip().lower() return raw in {"1", "true", "yes", "on"} @@ -3789,15 +4364,27 @@ def _qsa_gather_call(): elif sel_mask is not None: mask = sel_mask elif S > 1: - qpos = pos_start + mx.arange(S, dtype=mx.int32) - tpos = mx.arange(T, dtype=mx.int32) - mask = (tpos[None, :] <= qpos[:, None])[None, None] + # No selection came back, so every visible key is visible: either + # the chunk's whole history fits the indexer budget + # (``last_nb <= block_topk``, i.e. ``T <= (block_topk + 1) * ratio + # - 1`` = 2,051 tokens on the production pack, where top-512 of + # <= 512 candidate blocks IS all of them), or the request is a + # vision one where QSA is bypassed outright. In both cases the + # mask below is EXACTLY causal -- so when the flag is armed, hand + # MLX the string and never build the tensor. + _qsa_prefill_count("mask_causal_eligible") + if _prefill_mask_fuse_enabled() and _prefill_causal_mask_is_exact( + pos_start=pos_start, rows=S, total_keys=T + ): + mask = _CAUSAL_MASK + else: + qpos = pos_start + mx.arange(S, dtype=mx.int32) + tpos = mx.arange(T, dtype=mx.int32) + mask = (tpos[None, :] <= qpos[:, None])[None, None] else: mask = None - out = mx.fast.scaled_dot_product_attention( - q, k, v, scale=self.scale, mask=mask - ) + out = _qsa_dense_attention(q, k, v, mask=mask, scale=self.scale) out = out.transpose(0, 2, 1, 3).reshape(B, S, -1) return self.o_proj(out * mx.sigmoid(gate)) diff --git a/mtplx/profiles.py b/mtplx/profiles.py index 9699b306d..bff61f963 100644 --- a/mtplx/profiles.py +++ b/mtplx/profiles.py @@ -430,6 +430,13 @@ def announce_runtime_gated_env( "MTPLX_QWEN4_PLE_FIRST_GATHER_EARLY", "MTPLX_SESSION_BANK_SHED_BOUNDARIES", "MTPLX_SESSION_BANK_PROTECTED_TERMINAL", + # PR #391 remainder ports (davidtai), stamped by the Flash-Next lane + # defaults on the fixed-M4 geometry; registered so operator A/B + # launches and pack contracts pass the boot-time runtime-env + # validator. All rounding-class, quality-gated. + "MTPLX_QWEN4_HC_M4", + "MTPLX_QWEN4_PREFILL_MASK_FUSE", + "MTPLX_QSA_PREFILL_QUERY_TILE", "MTPLX_NGRAM_PREWARM_ORDER", "MTPLX_STRICT_CLAIMS", "MTPLX_QWEN4_COMPILED_MTP_PREPARE", diff --git a/mtplx/qwen4_prefill_chunk.py b/mtplx/qwen4_prefill_chunk.py new file mode 100644 index 000000000..6bc468255 --- /dev/null +++ b/mtplx/qwen4_prefill_chunk.py @@ -0,0 +1,77 @@ +"""The QSA prefill "middle path": a wide chunk with narrow attention. + +The 36 GDN layers, the MoE grouped GEMM and every projection want a WIDE +prefill chunk (better grouped-GEMM tile occupancy, fewer per-chunk syncs); +the 12 dense QSA layers want a NARROW one -- their score tensor is +``[H, rows, T]`` and their work term ``rows x T`` grows with the chunk. +Splitting only the attention into query tiles gives both: tile A never reads +tile B's keys, so a 4,096-row chunk tiled at 2,048 has exactly the peak AND +exactly the ``sum(rows x context)`` of an 8 x 2,048 cut, while everything +outside attention still sees 4,096 rows. + +``MTPLX_QSA_PREFILL_QUERY_TILE`` (old name ``MTPLX_FABLE_PREFILL_QSA_QUERY_TILE`` +still works as an alias) is the rows-per-tile knob; 0 / unset means the whole +chunk, i.e. today's behaviour, so flag-off is byte-identical. +""" + +from __future__ import annotations + +import os +from typing import Mapping + +#: Rows per QSA attention query tile. The new key wins when set to a +#: non-empty value; the old name is honoured only when the new one is unset. +QUERY_TILE_ENV = "MTPLX_QSA_PREFILL_QUERY_TILE" +QUERY_TILE_ALIAS_ENV = "MTPLX_FABLE_PREFILL_QSA_QUERY_TILE" + + +def _env(name: str, environ: Mapping[str, str] | None = None) -> str: + source = os.environ if environ is None else environ + return str(source.get(name) or "").strip() + + +def resolve_query_tile_rows(environ: Mapping[str, str] | None = None) -> int: + """Rows per QSA attention query tile; 0 (default) = whole chunk.""" + + raw = _env(QUERY_TILE_ENV, environ) or _env(QUERY_TILE_ALIAS_ENV, environ) + if not raw: + return 0 + try: + rows = int(raw) + except ValueError: + return 0 + return rows if rows > 0 else 0 + + +def query_tile_spans( + rows: int, *, context_before: int, tile: int +) -> list[tuple[int, int, int]]: + """``(row_start, row_end, keys_visible)`` for one chunk's query tiles. + + Attention rows are independent -- each row's softmax runs over its own + causal/selected key set -- so grouping rows differently cannot change + which keys a row sees. ``keys_visible`` is the exclusive key bound for + the tile's LAST row, and every earlier row in the tile is masked down to + its own bound exactly as before. Dropping the keys past that bound is + mathematically a no-op: the dense path fills them with the score dtype's + ``finfo.min``, whose ``exp`` underflows to a hard zero. + + The reduction *order* changes (shorter softmax rows, shorter P@V K), so + the result is exact-visible-set, not bit-identical -- the same class as + the portable gather tier. + + Empty list when tiling does not apply, so callers keep one code path. + """ + + span_rows = int(rows) + tile_rows = int(tile) + if span_rows <= 0 or tile_rows <= 0 or tile_rows >= span_rows: + return [] + base = max(0, int(context_before)) + spans: list[tuple[int, int, int]] = [] + row = 0 + while row < span_rows: + row_end = min(span_rows, row + tile_rows) + spans.append((row, row_end, base + row_end)) + row = row_end + return spans diff --git a/mtplx/runtime.py b/mtplx/runtime.py index bec879046..a150e56ef 100644 --- a/mtplx/runtime.py +++ b/mtplx/runtime.py @@ -1040,6 +1040,17 @@ def load( routed_glu_enabled=routed_glu_enabled, ) logger.info("[qwen4-M4-stage3] %s", qwen4_m4_stage3_report) + # MTPLX_QWEN4_HC_M4's PACK contract, checked here because the weights + # exist by now. Everything it validates is process-invariant, so a + # miss is a deployment error: it must stop the server coming up with a + # precise reason rather than turn the first request that reaches verify + # width into an HTTP 500. No-op (armed=False) when the flag is off. + from .models.qwen4_exp import install_hc_m4_pack_validation + + hc_m4_report = install_hc_m4_pack_validation(runtime.model) + runtime.qwen4_hc_m4_report = hc_m4_report + if hc_m4_report.get("armed"): + logger.info("[qwen4-hc-m4] %s", hc_m4_report) if whole_moe_plan is not None: if compiled_target_factory is None: from .a3b_whole_moe import A3BWholeMoeConfigError diff --git a/mtplx/runtime_options.py b/mtplx/runtime_options.py index d133534bc..bb4bf72b0 100644 --- a/mtplx/runtime_options.py +++ b/mtplx/runtime_options.py @@ -223,6 +223,34 @@ def reset_qwen4_verify_glue_cache(env: Mapping[str, str] | None = None) -> None: ) +#: Verify-width fused hyper-connection read (mtplx/kernels/qwen4_m4_hyper_read). +#: +#: Read ONCE at import so the hot verify path never touches ``os.environ`` and +#: two traces of the same compiled graph cannot disagree about which chain they +#: carry. ``MTPLX_QWEN4_HC_M4`` is the key; the old ``MTPLX_FABLE_HC_M4`` name +#: is honoured as an alias only when the new key is unset (the new key wins for +#: any non-empty value, including ``0`` for the per-key opt-out). The kernel +#: RAISES on a family-contract miss rather than falling back, so an +#: armed-but-inert lane is unreachable. +def _qwen4_hc_m4_import_default() -> bool: + raw = os.environ.get("MTPLX_QWEN4_HC_M4") + if raw is None or not str(raw).strip(): + return env_bool("MTPLX_FABLE_HC_M4", default=False) + return env_bool("MTPLX_QWEN4_HC_M4", default=False) + + +_QWEN4_HC_M4 = _qwen4_hc_m4_import_default() + + +def qwen4_hc_m4_enabled() -> bool: + """True when the HC_M4 flag armed this process at import. + + Armed by ``MTPLX_QWEN4_HC_M4`` (or the old ``MTPLX_FABLE_HC_M4`` alias). + """ + + return _QWEN4_HC_M4 + + @dataclass(frozen=True) class ResolvedAPIKey: value: str | None diff --git a/mtplx/server/openai.py b/mtplx/server/openai.py index 2fbcdd0e6..a17062a5a 100644 --- a/mtplx/server/openai.py +++ b/mtplx/server/openai.py @@ -927,6 +927,16 @@ def _server_runtime_env_overrides( "MTPLX_QWEN4_PLE_FIRST_GATHER_EARLY", "MTPLX_SESSION_BANK_SHED_BOUNDARIES", "MTPLX_SESSION_BANK_PROTECTED_TERMINAL", + # PR #391 remainder ports (davidtai), same fixed-M4 geometry. + # Decode: the verify-width fused hyper-connection read + # (mtplx/kernels/qwen4_m4_hyper_read; rounding-class, RAISES on + # a family-contract miss rather than falling back). + "MTPLX_QWEN4_HC_M4", + # Prefill: the causal-mask fuse routes the dense QSA prefill + # chunk through MLX's fused SDPA (rounding-class, exact visible + # set; per-shape-class capability cache, so a verify step MLX + # refuses never disarms a wide chunk). + "MTPLX_QWEN4_PREFILL_MASK_FUSE", ] if _qwen4_port_opt_in( overrides, "MTPLX_FUSED_GATE_UP" @@ -942,6 +952,14 @@ def _server_runtime_env_overrides( # leaves them inert. if os.environ.get("MTPLX_NGRAM_PREWARM") is None: overrides.setdefault("MTPLX_NGRAM_PREWARM", "auto") + # PR #391 remainder port (davidtai): the QSA prefill query tile. + # Caps the dense QSA attention peak to 2,048 rows so a wider + # prefill chunk keeps the 8x2,048 attention peak AND cost. Inert + # at the production 2,048 chunk width (tile >= chunk == no-op), so + # it only bites a 4,096-row chunk experiment; an explicit export + # (including 0 for whole-chunk) wins via the pop loop below. + if os.environ.get("MTPLX_QSA_PREFILL_QUERY_TILE") is None: + overrides.setdefault("MTPLX_QSA_PREFILL_QUERY_TILE", "2048") # The stage-3 child routes are consumed at model load and raise # unless stage 3 itself resolves on, so they are derived from the # resolved parent, never stamped alone: the routed-down reduction, @@ -1097,6 +1115,10 @@ def _served_model_type_is_qwen4_exp(args: argparse.Namespace) -> bool: "MTPLX_QWEN4_PLE_FIRST_GATHER_EARLY", "MTPLX_SESSION_BANK_SHED_BOUNDARIES", "MTPLX_SESSION_BANK_PROTECTED_TERMINAL", + # PR #391 remainder ports (davidtai): each is its own kill switch through + # the pop loop below. + "MTPLX_QWEN4_HC_M4", + "MTPLX_QWEN4_PREFILL_MASK_FUSE", "MTPLX_NGRAM_PREWARM", ) # Every key the fixed-M4 lane defaults may stamp; an explicit operator @@ -1105,6 +1127,8 @@ def _served_model_type_is_qwen4_exp(args: argparse.Namespace) -> bool: "MTPLX_FRSPEC_DRAFT", "MTPLX_FRSPEC_VOCAB", "MTPLX_QSA_GATHER_MAX_ROWS", + # PR #391 remainder port (davidtai): the QSA prefill query-tile value. + "MTPLX_QSA_PREFILL_QUERY_TILE", ) @@ -18625,6 +18649,9 @@ def _qwen4_install_reports(state: Any) -> dict[str, Any]: stage3 = getattr(runtime, "qwen4_m4_stage3_report", None) if isinstance(stage3, dict): out["m4_stage3"] = stage3 + hc_m4 = getattr(runtime, "qwen4_hc_m4_report", None) + if isinstance(hc_m4, dict): + out["hc_m4"] = hc_m4 glue = getattr(runtime, "_mtplx_qwen4_verify_glue", None) if isinstance(glue, dict): out["verify_glue"] = glue diff --git a/tests/test_qwen4_hc_m4.py b/tests/test_qwen4_hc_m4.py new file mode 100644 index 000000000..0607525c6 --- /dev/null +++ b/tests/test_qwen4_hc_m4.py @@ -0,0 +1,602 @@ +"""MTPLX_QWEN4_HC_M4 — wiring, eligibility, and the eager chain's op contract. + +Two things are proved here, both entirely on the CPU stream with tiny tensors +(no Metal, no kernel dispatch, no model): + +1. THE OP CONTRACT the Metal kernel encodes. ``mtplx/kernels/qwen4_m4_hyper_read`` + reproduces the eager chain op by op, and every rounding boundary it copies + is an assumption about how MLX types and rounds a specific expression. If + MLX ever changes one of those (``2.0 * bf16`` promoting to fp32, ``mx.mean`` + accumulating in fp32, ``nn.silu`` growing its own kernel), the kernel goes + silently wrong. These tests fail instead. + +2. THE GATE. Flag off: nothing changes and nothing is imported. Flag on: rows + 1 keeps the old path, rows 2..8 either take the kernel or RAISE with the + offending field named. There is no silent fallback — that is the failure + mode that left MTPLX_FUSED_HC_V3 armed but structurally inert at M=4. +""" + +from __future__ import annotations + +import mlx.core as mx +import mlx.nn as nn +import pytest + +import mtplx.kernels.qwen4_m4_hyper_read as hcm4 +import mtplx.models.qwen4_exp as qwen4_exp +import mtplx.runtime_options as runtime_options + +HCD = hcm4.HCD +LOWRANK = hcm4.R_LOWRANK +HC = hcm4.HC + + +@pytest.fixture(autouse=True) +def _cpu_stream(): + """Confine every op in this module to the CPU stream.""" + + with mx.stream(mx.cpu): + yield + + +@pytest.fixture(autouse=True) +def _flag_off(monkeypatch): + """Default every test to the shipped state, whatever the session env is.""" + + monkeypatch.setattr(runtime_options, "_QWEN4_HC_M4", False) + + +@pytest.fixture +def armed(monkeypatch): + monkeypatch.setattr(runtime_options, "_QWEN4_HC_M4", True) + return True + + +def _bits(value: mx.array) -> mx.array: + widths = {mx.bfloat16: mx.uint16, mx.float16: mx.uint16, mx.float32: mx.uint32} + return value.view(widths[value.dtype]) + + +def _same_bits(a: mx.array, b: mx.array) -> bool: + return bool(mx.all(_bits(a) == _bits(b)).item()) + + +# -------------------------------------------------------------------------- +# 1. the eager chain's op contract, as encoded in the Metal source +# -------------------------------------------------------------------------- + + +def test_divide_by_hc_count_stays_bf16(): + """``mix / self.hc_count`` must not widen: the kernel rounds it to bf16.""" + + a = (mx.random.normal((64,)) * 3).astype(mx.bfloat16) + assert (a / 4).dtype == mx.bfloat16 + assert _same_bits(a / 4, (a.astype(mx.float32) * 0.25).astype(mx.bfloat16)) + + +def test_inject_scale_stays_bf16(): + """``2.0 * mx.sigmoid(...)`` — a Python float is a weak scalar in MLX, so + the inject value the kernel writes is bf16, not fp32.""" + + a = (mx.random.normal((64,)) * 3).astype(mx.bfloat16) + s = mx.sigmoid(a) + assert (2.0 * s).dtype == mx.bfloat16 + + +def test_silu_is_x_times_sigmoid_x(): + """The kernel writes ``(T)(sigmoid(t0) * t0)``; nn.silu must be that.""" + + a = (mx.random.normal((256,)) * 3).astype(mx.bfloat16) + assert nn.silu(a).dtype == mx.bfloat16 + assert _same_bits(nn.silu(a), a * mx.sigmoid(a)) + + +def test_hc_mean_is_bf16_sum_times_quarter(): + """``mx.mean(mix * grouped, axis=-2)`` == ``mx.sum(...) * 0.25``, and that + sum accumulates IN bf16 — which is why the kernel's hc reduction rounds at + every add instead of accumulating in fp32.""" + + a = (mx.random.normal((512, HC, 4)) * 3).astype(mx.bfloat16) + m = mx.mean(a, axis=-2) + s = mx.sum(a, axis=-2) + assert m.dtype == mx.bfloat16 and s.dtype == mx.bfloat16 + assert _same_bits(m, (s.astype(mx.float32) * 0.25).astype(mx.bfloat16)) + + seq = a[:, 0, :] + for g in range(1, HC): + seq = seq + a[:, g, :] + assert _same_bits(s, seq), ( + "mx.sum over a length-4 bf16 axis is no longer a sequential bf16 " + "accumulation; qwen4_m4_hyper_read's hc reduction must follow it" + ) + + +def test_hc_mean_is_not_fp32_accumulated(): + """Guard the *reason* the test above is not cosmetic: an fp32 accumulation + of the same four terms is a genuinely different answer.""" + + a = (mx.random.normal((4096, HC)) * 3).astype(mx.bfloat16) + s = mx.sum(a, axis=-1) + f32 = a.astype(mx.float32).sum(axis=-1).astype(mx.bfloat16) + assert not _same_bits(s, f32) + + +def test_grouped_rms_norm_decomposition(): + """``GroupedRMSNorm`` == per-group rms_norm rounded to bf16, then a bf16 + multiply by the full-width weight — the two-step the kernel's K0 copies.""" + + dims, group = 32, 8 + norm = qwen4_exp.GroupedRMSNorm(dims, group, eps=1e-6) + norm.weight = (mx.random.normal((dims,)) * 0.5 + 1.0).astype(mx.bfloat16) + x = (mx.random.normal((3, dims)) * 2).astype(mx.bfloat16) + + got = norm(x) + grouped = x.reshape(3, -1, group) + step1 = mx.fast.rms_norm(grouped, None, 1e-6).reshape(3, dims) + assert step1.dtype == mx.bfloat16 + assert _same_bits(got, step1 * norm.weight) + + +# -------------------------------------------------------------------------- +# 2. check_shapes — returns rows, or raises with the field named +# -------------------------------------------------------------------------- + + +def _family_weights(dtype=mx.bfloat16, *, down_rows=LOWRANK): + return { + "gamma": mx.zeros((HCD,), dtype=dtype), + "wd": mx.zeros((down_rows, HCD), dtype=dtype), + "wu": mx.zeros((HCD, LOWRANK), dtype=dtype), + "wi": mx.zeros((HC, HCD), dtype=dtype), + } + + +@pytest.mark.parametrize("rows", [2, 3, 4, 5, 8]) +def test_check_shapes_accepts_verify_widths(rows): + w = _family_weights() + x = mx.zeros((rows, HCD), dtype=mx.bfloat16) + assert hcm4.check_shapes(x, w["gamma"], w["wd"], w["wu"], w["wi"]) == rows + + +@pytest.mark.parametrize("rows", [1, 9, 16]) +def test_check_shapes_rejects_other_widths(rows): + w = _family_weights() + x = mx.zeros((rows, HCD), dtype=mx.bfloat16) + with pytest.raises(ValueError, match="wired for 2..8 rows"): + hcm4.check_shapes(x, w["gamma"], w["wd"], w["wu"], w["wi"]) + + +def test_check_shapes_rejects_wrong_hidden(): + w = _family_weights() + with pytest.raises(ValueError, match=r"must be \[\.\.\., 10240\]"): + hcm4.check_shapes( + mx.zeros((4, 4096), dtype=mx.bfloat16), + w["gamma"], w["wd"], w["wu"], w["wi"], + ) + + +def test_check_shapes_names_the_bad_weight(): + w = _family_weights(down_rows=256) + x = mx.zeros((4, HCD), dtype=mx.bfloat16) + with pytest.raises(ValueError, match="input_mix_weight_down.weight"): + hcm4.check_shapes(x, w["gamma"], w["wd"], w["wu"], w["wi"]) + + +def test_check_shapes_rejects_dtype_mismatch(): + w = _family_weights() + w["wu"] = mx.zeros((HCD, LOWRANK), dtype=mx.float16) + x = mx.zeros((4, HCD), dtype=mx.bfloat16) + with pytest.raises(ValueError, match="input_mix_weight_up.weight dtype"): + hcm4.check_shapes(x, w["gamma"], w["wd"], w["wu"], w["wi"]) + + +def test_check_shapes_takes_the_module_shape_unreshaped(): + """GatedResidual is called with [B, S, 10240]; R is the leading product, + and validating must not require materialising a reshape node.""" + + w = _family_weights() + x = mx.zeros((1, 4, HCD), dtype=mx.bfloat16) + assert hcm4.check_shapes(x, w["gamma"], w["wd"], w["wu"], w["wi"]) == 4 + x = mx.zeros((2, 3, HCD), dtype=mx.bfloat16) + assert hcm4.check_shapes(x, w["gamma"], w["wd"], w["wu"], w["wi"]) == 6 + + +def test_driver_requires_a_two_d_view(): + w = _family_weights() + with pytest.raises(ValueError, match="wants a 2-D"): + hcm4.fused_hc_read_m4( + mx.zeros((1, 4, HCD), dtype=mx.bfloat16), + w["gamma"], w["wd"], w["wu"], w["wi"], + ) + + +def test_check_shapes_allows_missing_inject(): + """The trunk mixer is built with ``use_combine=False``.""" + + w = _family_weights() + x = mx.zeros((4, HCD), dtype=mx.bfloat16) + assert hcm4.check_shapes(x, w["gamma"], w["wd"], w["wu"], None) == 4 + + +def test_weight_bytes_matches_the_census_figure(): + assert hcm4.weight_bytes_per_read(2) == 13_209_600 + assert hcm4.weight_bytes_per_read(2, has_inject=False) == 13_127_680 + + +# -------------------------------------------------------------------------- +# 3. the gate in GatedResidual +# -------------------------------------------------------------------------- + + +def _family_module(use_combine=True, dtype=mx.bfloat16): + mod = qwen4_exp.GatedResidual(qwen4_exp.TextArgs(), use_combine=use_combine) + w = _family_weights(dtype) + mod.hc_norm.weight = w["gamma"] + mod.input_mix_weight_down.weight = w["wd"] + mod.input_mix_weight_up.weight = w["wu"] + if use_combine: + mod.block_inject_weight.weight = w["wi"] + return mod + + +def _small_module(use_combine=True): + args = qwen4_exp.TextArgs(hidden_size=8, hc_lowrank=4, rms_norm_eps=1e-6) + mod = qwen4_exp.GatedResidual(args, use_combine=use_combine) + hcd = args.hc_count * args.hidden_size + mod.hc_norm.weight = (mx.random.normal((hcd,)) * 0.3 + 1.0).astype(mx.bfloat16) + mod.input_mix_weight_down.weight = ( + mx.random.normal((args.hc_lowrank, hcd)) * 0.1 + ).astype(mx.bfloat16) + mod.input_mix_weight_up.weight = ( + mx.random.normal((hcd, args.hc_lowrank)) * 0.1 + ).astype(mx.bfloat16) + if use_combine: + mod.block_inject_weight.weight = ( + mx.random.normal((args.hc_count, hcd)) * 0.1 + ).astype(mx.bfloat16) + return mod, args + + +def test_flag_off_never_applies(): + mod = _family_module() + for rows in (1, 2, 4, 8, 32): + x = mx.zeros((1, rows, HCD), dtype=mx.bfloat16) + assert mod._hc_m4_applies(x) is False + + +def test_flag_off_leaves_the_eager_chain_bit_identical(): + """The full eager expression, transcribed, on a tiny config.""" + + mod, args = _small_module() + hcd = args.hc_count * args.hidden_size + x = (mx.random.normal((1, 4, hcd)) * 2).astype(mx.bfloat16) + + mixed, passthrough, inject = mod(x) + + normed = mod.hc_norm(x) + mix = nn.silu(mod.input_mix_weight_down(normed) / args.hc_count) + mix = mx.sigmoid(mod.input_mix_weight_up(mix)) + mix = mix.reshape(*mix.shape[:-1], args.hc_count, args.hidden_size) + grouped = normed.reshape(*normed.shape[:-1], args.hc_count, args.hidden_size) + want_mixed = mx.mean(mix * grouped, axis=-2) + want_inject = 2.0 * mx.sigmoid(mod.block_inject_weight(normed) / args.hc_count) + + assert _same_bits(mixed, want_mixed) + assert _same_bits(inject, want_inject) + assert passthrough is x + + +def test_flag_off_tolerates_off_family_geometry(): + """Construction must not raise for a non-Flash-Next config when the flag + is unset — this module class is shared.""" + + qwen4_exp.GatedResidual(qwen4_exp.TextArgs(hidden_size=8, hc_lowrank=4)) + + +@pytest.mark.parametrize("rows", [2, 3, 4, 8]) +def test_armed_applies_at_verify_widths(armed, rows): + mod = _family_module() + assert mod._hc_m4_applies(mx.zeros((1, rows, HCD), dtype=mx.bfloat16)) is True + + +def test_armed_applies_to_the_noncombine_mixer(armed): + mod = _family_module(use_combine=False) + assert mod._hc_m4_applies(mx.zeros((1, 4, HCD), dtype=mx.bfloat16)) is True + + +@pytest.mark.parametrize("rows", [1]) +def test_armed_leaves_row_one_alone(armed, rows): + """rows == 1 is the draft path's business (v3 / eager); this gate must not + take it, and must not raise on it either.""" + + mod = _family_module() + assert mod._hc_m4_applies(mx.zeros((1, rows, HCD), dtype=mx.bfloat16)) is False + + +def test_armed_leaves_prefill_widths_alone(armed): + mod = _family_module() + assert mod._hc_m4_applies(mx.zeros((1, 64, HCD), dtype=mx.bfloat16)) is False + + +def test_armed_off_family_config_raises_at_construction(armed): + with pytest.raises(RuntimeError, match="not the Flash-Next family shape"): + qwen4_exp.GatedResidual(qwen4_exp.TextArgs(hidden_size=8, hc_lowrank=4)) + with pytest.raises(RuntimeError, match="hc_count=2"): + qwen4_exp.GatedResidual(qwen4_exp.TextArgs(hc_count=2)) + + +def test_armed_quantized_mix_weights_raise(armed): + mod = _family_module() + mod.input_mix_weight_down.scales = mx.zeros((LOWRANK, HCD // 64), mx.bfloat16) + with pytest.raises(RuntimeError, match="input_mix_weight_down is quantized"): + mod._hc_m4_applies(mx.zeros((1, 4, HCD), dtype=mx.bfloat16)) + + +def test_armed_dtype_mismatch_raises(armed): + """fp32 mix weights against a bf16 hyper state: raise, do not fall back.""" + + mod = _family_module(dtype=mx.float32) + with pytest.raises(ValueError, match="dtype"): + mod._hc_m4_applies(mx.zeros((1, 4, HCD), dtype=mx.bfloat16)) + + +def test_armed_wrong_down_rows_raise(armed): + mod = _family_module() + mod.input_mix_weight_down.weight = mx.zeros((256, HCD), dtype=mx.bfloat16) + with pytest.raises(ValueError, match="input_mix_weight_down.weight"): + mod._hc_m4_applies(mx.zeros((1, 4, HCD), dtype=mx.bfloat16)) + + +def test_env_flag_is_read_once(monkeypatch): + """A mid-run env change must not reach the hot path: two traces of the + same compiled verify graph would then disagree about which read they + contain.""" + + before = runtime_options.qwen4_hc_m4_enabled() + monkeypatch.setenv("MTPLX_QWEN4_HC_M4", "1") + assert runtime_options.qwen4_hc_m4_enabled() is before + monkeypatch.delenv("MTPLX_QWEN4_HC_M4", raising=False) + assert runtime_options.qwen4_hc_m4_enabled() is before + + +def test_env_flag_defaults_off(): + from mtplx.runtime_options import env_bool + + assert env_bool("MTPLX_QWEN4_HC_M4", default=False, env={}) is False + + +def test_env_flag_rejects_a_bad_spelling(): + from mtplx.runtime_options import env_bool + + with pytest.raises(ValueError): + env_bool("MTPLX_QWEN4_HC_M4", default=False, env={"MTPLX_QWEN4_HC_M4": "yep"}) + + +# -------------------------------------------------------------------------- +# 4. the driver's own guards (no dispatch: these all raise before the kernel) +# -------------------------------------------------------------------------- + + +@pytest.mark.parametrize( + "kwargs, match", + [ + ({"norm_threads": 100}, "norm_threads"), + ({"down_threads": 2048}, "down_threads"), + ({"out_per_tg": 0}, "out_per_tg"), + ({"d_per_block": 999}, "d_per_block"), + ], +) +def test_driver_rejects_bad_tuning(kwargs, match): + w = _family_weights() + x = mx.zeros((4, HCD), dtype=mx.bfloat16) + with pytest.raises(ValueError, match=match): + hcm4.fused_hc_read_m4( + x, w["gamma"], w["wd"], w["wu"], w["wi"], **kwargs + ) + + +def test_dispatch_budget_is_three(): + assert hcm4.DISPATCHES_PER_READ == 3 + assert hcm4.EAGER_DISPATCHES_PER_READ == 11 + + +# -------------------------------------------------------------------------- +# 5. the PACK contract moves to install time +# -------------------------------------------------------------------------- +# +# Every term below is a property of the loaded weights, so it is the same +# answer for every request the process will ever serve. Checking it when the +# weights land means a mis-armed flag stops the server coming up with a +# precise reason, instead of turning the first request that happens to reach +# verify width into an HTTP 500. + + +def _real_tree(num_layers: int = 2, *, with_mtp: bool = True): + """The REAL Qwen3.8 module tree, built from the model's own config. + + Same classes, same attribute names, same containers as a served pack -- + ``Model.language_model.model.layers[i].{attn,mlp}_hyper_connection``, the + text model's own ``hyper_connection_mixer``, and the MTP sub-tree + published on ``language_model`` by ``inject_qwen4_exp_mtp_support``. Only + the layer count and the vocabulary are cut, and neither is part of the + MTPLX_QWEN4_HC_M4 contract; the hyper-connection geometry the kernel + hardcodes (hc_count 4, hidden 2560, lowrank 320) is the shipped one. + + Cheap despite the real shapes: MLX arrays are lazy, and nothing here + evaluates one. + """ + + config = { + "model_type": "qwen4_exp", + "num_hidden_layers": num_layers, + # full_attention first: Qwen4ExpMTP picks the first non-linear layer + # as its own, exactly as the shipped config's layer_types allow. + "layer_types": (["full_attention", "linear_attention"] * num_layers)[ + :num_layers + ], + "vocab_size": 1024, + "tie_word_embeddings": False, + } + model = qwen4_exp.Model(qwen4_exp.ModelArgs.from_dict(config)) + if with_mtp: + # `Qwen4ExpMTP` is published on language_model, NOT on Model -- + # registering it on both trees would double-count its parameters. + model.language_model.mtp = qwen4_exp.Qwen4ExpMTP(model.language_model.args) + return model + + +def test_pack_validation_is_a_no_op_when_the_flag_is_off(): + report = qwen4_exp.install_hc_m4_pack_validation(_real_tree()) + assert report == {"armed": False, "validated": 0} + + +def test_pack_validation_finds_every_module_in_the_real_tree(armed): + """The 2026-09-02 dead-server regression, pinned. + + The first cut walked ``dir(layer)`` and found NOTHING on a served pack, + so an armed flag raised "no GatedResidual ... to replace" and killed the + load for both the HumanEval screen and the ABBA lane -- on a model where + the kernel had been installing and running bit-exact all day. + """ + + model = _real_tree(num_layers=2) + report = qwen4_exp.install_hc_m4_pack_validation(model) + assert report["armed"] is True + # 2 per decoder layer + the text model's mixer + the MTP's own layer pair + # and mixer. + assert report["validated"] == 8 + + +def test_pack_validation_reaches_all_three_places_the_model_keeps_one(armed): + """Held directly, inside a list, and behind a published sub-tree.""" + + model = _real_tree(num_layers=2) + paths = [path for path, _ in qwen4_exp._named_gated_residuals(model)] + assert paths == [ + "language_model.model.hyper_connection_mixer", + "language_model.model.layers.0.attn_hyper_connection", + "language_model.model.layers.0.mlp_hyper_connection", + "language_model.model.layers.1.attn_hyper_connection", + "language_model.model.layers.1.mlp_hyper_connection", + "language_model.mtp.hyper_connection_mixer", + "language_model.mtp.layers.0.attn_hyper_connection", + "language_model.mtp.layers.0.mlp_hyper_connection", + ] + + +def test_dir_cannot_see_an_mlx_modules_children(armed): + """Why the first cut found nothing -- pinned so it cannot come back. + + An ``nn.Module``'s children live in its dict and are served through + ``__getattr__``, so they are absent from ``dir()``. Discovery must go + through the model's own traversal. + """ + + layer = _real_tree(num_layers=1).language_model.model.layers[0] + assert isinstance(layer.attn_hyper_connection, qwen4_exp.GatedResidual) + seen_by_dir = [ + name + for name in dir(layer) + if isinstance(getattr(layer, name, None), qwen4_exp.GatedResidual) + ] + assert seen_by_dir == [] + seen_by_named_modules = [ + name + for name, module in layer.named_modules() + if isinstance(module, qwen4_exp.GatedResidual) + ] + assert sorted(seen_by_named_modules) == [ + "attn_hyper_connection", + "mlp_hyper_connection", + ] + + +def test_pack_validation_names_the_quantized_layer(armed): + model = _real_tree(num_layers=2) + victim = model.language_model.model.layers[1].attn_hyper_connection + victim.input_mix_weight_down.scales = mx.zeros((LOWRANK, HCD // 64), mx.bfloat16) + with pytest.raises( + RuntimeError, match=r"layers\.1\.attn_hyper_connection.*is quantized" + ): + qwen4_exp.install_hc_m4_pack_validation(model) + + +def test_pack_validation_names_the_wrong_weight_shape(armed): + model = _real_tree(num_layers=1) + victim = model.language_model.model.hyper_connection_mixer + victim.input_mix_weight_down.weight = mx.zeros((256, HCD), dtype=mx.bfloat16) + with pytest.raises(RuntimeError, match="input_mix_weight_down.weight"): + qwen4_exp.install_hc_m4_pack_validation(model) + + +def test_pack_validation_refuses_a_model_with_nothing_to_replace(armed): + """An armed flag that can never do anything is a deployment error.""" + + import mlx.nn as nn + + class _NoHyperConnections(nn.Module): + def __init__(self): + super().__init__() + self.embed = nn.Embedding(8, 8) + + with pytest.raises(RuntimeError, match="no GatedResidual"): + qwen4_exp.install_hc_m4_pack_validation(_NoHyperConnections()) + + +def test_pack_validation_does_not_check_the_activation_dtype(armed): + """That term needs an activation; `_hc_m4_applies` still owns it. + + It is process-invariant too (one pack, one activation dtype), so it cannot + single out one request either -- but it cannot be answered without a + forward, so it stays where the forward is. + """ + + model = _real_tree(num_layers=1, with_mtp=False) + assert qwen4_exp.install_hc_m4_pack_validation(model)["validated"] == 3 + with pytest.raises(ValueError, match="dtype"): + _family_module(dtype=mx.float32)._hc_m4_applies( + mx.zeros((1, 4, HCD), dtype=mx.bfloat16) + ) + + +def test_the_load_sequence_publishes_the_mtp_before_the_validation(armed): + """Replays runtime.load's module wiring in the order runtime.py uses. + + The order is read out of `runtime.py` itself rather than restated, so the + test fails if the hook is ever moved above the MTP injection or the qwen4 + installs -- which would put it back to validating a tree the model has not + finished building. + """ + + import inspect + + from mtplx import runtime + + source = inspect.getsource(runtime.load) + steps = ( + "inject_qwen4_exp_mtp_support(", # publishes language_model.mtp + "install_qwen4_fixed_verify_route(", + "install_qwen4_m4_stage3(", + "install_hc_m4_pack_validation(runtime.model)", + ) + positions = [source.index(step) for step in steps] + assert positions == sorted(positions), "runtime.load's qwen4 order moved" + + # Now the same order, for real: a model whose MTP is not yet published + # exposes fewer modules than one whose MTP is. + model = _real_tree(num_layers=2, with_mtp=False) + before = qwen4_exp.install_hc_m4_pack_validation(model)["validated"] + model.language_model.mtp = qwen4_exp.Qwen4ExpMTP(model.language_model.args) + after = qwen4_exp.install_hc_m4_pack_validation(model)["validated"] + assert (before, after) == (5, 8) + + +def test_the_runtime_validates_the_pack_at_install(armed): + """The install hook is wired into the runtime's qwen4 section.""" + + import inspect + + from mtplx import runtime + + source = inspect.getsource(runtime) + assert "install_hc_m4_pack_validation(runtime.model)" in source diff --git a/tests/test_qwen4_prefill_mask_fuse.py b/tests/test_qwen4_prefill_mask_fuse.py new file mode 100644 index 000000000..68448e764 --- /dev/null +++ b/tests/test_qwen4_prefill_mask_fuse.py @@ -0,0 +1,842 @@ +"""MTPLX_QWEN4_PREFILL_MASK_FUSE -- the causal arm of the dense QSA lane. + +Below the sparse crossover the QSA indexer hands attention **no** selection +whenever the chunk's whole history fits its block budget +(``last_nb <= block_topk``): top-512 of at most 512 candidate blocks is all +of them, so the visible set is exactly causal. The lane used to materialise +a dense ``[1, 1, S, T]`` bool mask for that case anyway, which is what pushes +MLX off the fused ``steel_attention`` route at ``head_dim`` 256. This module +pins the replacement: + +* **when** the selection is trivially complete (the ``T <= (block_topk + 1) * + ratio - 1`` frontier, 2,051 tokens on the production pack) and **when** the + mask is causal-with-offset (chunk k sees all prior context plus a causal + block within the chunk); +* that MLX 0.32.2's ``mask="causal"`` -- documented lower-right aligned -- + describes that offset case exactly, checked against the installed MLX + rather than assumed; +* that the visible set of the causal string, of the dense mask the lane + built, and of ``_qsa_blocks_to_dense_mask`` under a full selection are the + same set; +* the counters (``mask_fuse_causal`` / ``mask_fuse_bool`` / + ``mask_fuse_unavailable``), the query-tile composition, and the loud + one-shot refusal when the build has no fused kernel. + +Everything runs on tiny tensors on the CPU stream: the GPU on a development +box is usually holding a guarded benchmark, and ``force_fused=True`` is +resolved while the op is BUILT, so the refusal path is exercised natively +there (MLX raises "the fused kernels require a GPU (Metal) stream"). What +this cannot settle is the fused kernel's own numerics or its speed. +""" + +from __future__ import annotations + +import io +import contextlib + +import mlx.core as mx +import numpy as np +import pytest + +import mtplx.models.qwen4_exp as qwen4_exp +from mtplx.qwen4_prefill_chunk import QUERY_TILE_ENV, query_tile_spans + +MASK_FUSE_ENV = "MTPLX_QWEN4_PREFILL_MASK_FUSE" + +#: Production pack (``~/.mtplx/models/Youssofal--Qwen3.8-Flash-Next-MTPLX- +#: Optimized-Speed/config.json``): indexer_budget 2048, compress_ratio 4 => +#: block_topk 512. +BLOCK_TOPK = 512 +RATIO = 4 +#: Largest post-update context whose selection is trivially complete: +#: ``T // ratio <= block_topk``. +TRIVIAL_FRONTIER = (BLOCK_TOPK + 1) * RATIO - 1 + + +@pytest.fixture(autouse=True) +def _cpu_default_device(): + # set_default_device leaks into every later-collected module (pytest + # shares one process), so restore it. + previous = mx.default_device() + mx.set_default_device(mx.cpu) + try: + yield + finally: + mx.set_default_device(previous) + + +@pytest.fixture(autouse=True) +def _clean_lane_state(monkeypatch): + """Every knob and one-shot in this lane is process-global; reset them.""" + + monkeypatch.delenv(MASK_FUSE_ENV, raising=False) + monkeypatch.delenv(QUERY_TILE_ENV, raising=False) + qwen4_exp._prefill_mask_fuse_enabled.cache_clear() + saved_counts = dict(qwen4_exp._QSA_PREFILL_COUNTS) + saved_unavailable = dict(qwen4_exp._PREFILL_MASK_FUSE_UNAVAILABLE) + saved_printed = qwen4_exp._MASK_FUSE_REFUSALS_PRINTED[0] + saved_engaged = qwen4_exp._MASK_FUSE_ENGAGED[0] + qwen4_exp._QSA_PREFILL_COUNTS.clear() + qwen4_exp._PREFILL_MASK_FUSE_UNAVAILABLE.clear() + qwen4_exp._MASK_FUSE_REFUSALS_PRINTED[0] = 0 + qwen4_exp._MASK_FUSE_ENGAGED[0] = False + try: + yield + finally: + qwen4_exp._prefill_mask_fuse_enabled.cache_clear() + qwen4_exp._QSA_PREFILL_COUNTS.clear() + qwen4_exp._QSA_PREFILL_COUNTS.update(saved_counts) + qwen4_exp._PREFILL_MASK_FUSE_UNAVAILABLE.clear() + qwen4_exp._PREFILL_MASK_FUSE_UNAVAILABLE.update(saved_unavailable) + qwen4_exp._MASK_FUSE_REFUSALS_PRINTED[0] = saved_printed + qwen4_exp._MASK_FUSE_ENGAGED[0] = saved_engaged + + +def _arm(monkeypatch, value: str = "1") -> None: + monkeypatch.setenv(MASK_FUSE_ENV, value) + qwen4_exp._prefill_mask_fuse_enabled.cache_clear() + + +def _lane_mask(pos_start: int, rows: int, total: int) -> mx.array: + """The dense mask the lane builds at ``qwen4_exp.Attention.__call__``.""" + + qpos = pos_start + mx.arange(rows, dtype=mx.int32) + tpos = mx.arange(total, dtype=mx.int32) + return (tpos[None, :] <= qpos[:, None])[None, None] + + +def _lower_right_causal(rows: int, total: int) -> mx.array: + """MLX's documented ``"causal"``: last query aligns with the last key.""" + + qpos = (total - rows) + mx.arange(rows, dtype=mx.int32) + tpos = mx.arange(total, dtype=mx.int32) + return (tpos[None, :] <= qpos[:, None])[None, None] + + +# --------------------------------------------------------------------------- +# 1. WHEN is the selection trivially complete / the mask causal-with-offset +# --------------------------------------------------------------------------- + + +def test_indexer_short_circuits_exactly_at_the_block_budget(): + """``last_nb <= block_topk`` is the whole condition, in both routes.""" + + trivial = [t for t in range(1, 4 * TRIVIAL_FRONTIER) if t // RATIO <= BLOCK_TOPK] + assert max(trivial) == TRIVIAL_FRONTIER == 2051 + # Both the eager route (_call_rows) and the compiled route + # (_compiled_mode -> "update_only") use this one predicate, so the two + # cannot disagree about when attention gets no selection. + source = qwen4_exp.__file__ + text = open(source, encoding="utf-8").read() + assert text.count("last_nb <= self.block_topk") == 2 + assert "last_nb = T // self.ratio" in text + + +@pytest.mark.parametrize( + "chunk, prompt", + [(2048, 16_384), (4096, 16_384), (2048, 32_768), (4096, 32_768)], +) +def test_which_prefill_chunks_are_trivially_complete(chunk, prompt): + """Only a chunk whose POST-update context is <= 2,051 gets no selection. + + At the retained prefill width (4,096) that is no chunk at all, at 16K or + at 32K: the causal arm cannot fire there, and the whole win of the flag + at those cells is the bool-mask arm. At the shipped 2,048 width it is + chunk 0 and nothing else. + """ + + trivial = [ + pos_start + for pos_start in range(0, prompt, chunk) + # T == pos_start + S, S == this chunk's rows + if (pos_start + min(chunk, prompt - pos_start)) // RATIO <= BLOCK_TOPK + ] + assert trivial == ([0] if chunk <= TRIVIAL_FRONTIER else []) + + +@pytest.mark.parametrize( + "pos_start, rows, total, expected", + [ + (0, 2048, 2048, True), # first chunk + (2048, 2048, 4096, True), # causal-with-offset: all prior + causal + (16_384, 4096, 20_480, True), + (0, 1, 1, True), + (0, 2048, 4096, False), # padded/capacity-shaped KV: NOT lower-right + (2048, 2048, 4097, False), + (0, 0, 0, False), + ], +) +def test_causal_string_is_exact_only_when_last_query_is_last_key( + pos_start, rows, total, expected +): + assert ( + qwen4_exp._prefill_causal_mask_is_exact( + pos_start=pos_start, rows=rows, total_keys=total + ) + is expected + ) + + +# --------------------------------------------------------------------------- +# 2. The visible sets are the SAME SET (the exactness invariant) +# --------------------------------------------------------------------------- + + +@pytest.mark.parametrize( + "pos_start, rows", [(0, 8), (0, 13), (5, 8), (16, 8), (2048 - 16, 16)] +) +def test_lane_mask_equals_mlx_lower_right_causal(pos_start, rows): + total = pos_start + rows + lane = np.array(_lane_mask(pos_start, rows, total)) + ref = np.array(_lower_right_causal(rows, total)) + assert np.array_equal(lane, ref) + # ... and it really is "all prior context + causal inside the chunk". + dense = lane[0, 0] + assert dense[:, :pos_start].all() + within = dense[:, pos_start:] + assert np.array_equal(within, np.tril(np.ones((rows, rows), dtype=bool))) + + +def test_full_block_selection_reconstructs_exactly_the_causal_mask(): + """``causal == the full-selection mask`` -- proved on the real function. + + ``_qsa_blocks_to_dense_mask`` is the lane's own reconstruction of a QSA + selection. Feed it every complete block (which is what the selector + returns when ``nb_total <= block_topk``, since ``k_eff = min(block_topk, + nb_total)``) and it must produce the causal mask, key for key. + """ + + ratio = 4 + for pos_start, rows in ((0, 16), (0, 12), (8, 8), (12, 20)): + total = pos_start + rows + logical_blocks = total // ratio + topk = max(1, logical_blocks) + # Every row selects all logical blocks; rows whose own complete-block + # frontier is lower have the surplus ids clipped by the function's own + # in-range guard, exactly as a real top-k of fewer candidates would. + block_ids = mx.broadcast_to( + mx.arange(topk, dtype=mx.int32)[None, :], (rows, topk) + ) + block_valid = mx.ones((rows, topk), dtype=mx.bool_) + got = qwen4_exp._qsa_blocks_to_dense_mask( + block_ids, + block_valid, + pos_start=pos_start, + total_tokens=total, + compress_ratio=ratio, + ) + want = _lane_mask(pos_start, rows, total) + assert np.array_equal(np.array(got), np.array(want)), (pos_start, rows) + + +@pytest.mark.parametrize("pos_start, rows", [(0, 16), (12, 20), (33, 7)]) +def test_mlx_causal_string_and_dense_mask_agree_numerically(pos_start, rows): + """Against the INSTALLED MLX, not against its documentation.""" + + total = pos_start + rows + rng = np.random.default_rng(20260902) + shape = (1, 2, rows, 16) + kshape = (1, 2, total, 16) + q = mx.array(rng.standard_normal(shape).astype(np.float32)) + k = mx.array(rng.standard_normal(kshape).astype(np.float32)) + v = mx.array(rng.standard_normal(kshape).astype(np.float32)) + scale = 16 ** -0.5 + string_out = mx.fast.scaled_dot_product_attention( + q, k, v, scale=scale, mask=qwen4_exp._CAUSAL_MASK + ) + dense_out = mx.fast.scaled_dot_product_attention( + q, k, v, scale=scale, mask=_lane_mask(pos_start, rows, total) + ) + mx.eval(string_out, dense_out) + assert np.allclose(np.array(string_out), np.array(dense_out), atol=1e-6) + + +def test_causal_and_dense_arms_agree_within_bf16_rounding(): + """bf16 inputs: the two arms are the same visible set, one rounding class.""" + + pos_start, rows, total = 24, 24, 48 + rng = np.random.default_rng(7) + q = mx.array(rng.standard_normal((1, 2, rows, 16)).astype(np.float32)).astype( + mx.bfloat16 + ) + kv = mx.array(rng.standard_normal((1, 2, total, 16)).astype(np.float32)).astype( + mx.bfloat16 + ) + scale = 16 ** -0.5 + a = qwen4_exp._qsa_dense_attention( + q, kv, kv, mask=qwen4_exp._CAUSAL_MASK, scale=scale + ) + b = qwen4_exp._qsa_dense_attention( + q, kv, kv, mask=_lane_mask(pos_start, rows, total), scale=scale + ) + mx.eval(a, b) + # bf16 has ~3 decimal digits; the visible set being identical is what + # makes this a tolerance and not a mismatch. + assert np.allclose( + np.array(a.astype(mx.float32)), + np.array(b.astype(mx.float32)), + atol=6e-3, + rtol=6e-3, + ) + + +# --------------------------------------------------------------------------- +# 3. Query-tile composition +# --------------------------------------------------------------------------- + + +def test_query_tile_spans_keep_every_tile_lower_right_aligned(): + """Why the string survives tiling: each tile's last query is its last key.""" + + for total_keys, rows, tile in ((4096, 4096, 2048), (20_480, 4096, 1024)): + context_before = total_keys - rows + for r0, r1, keys in query_tile_spans( + rows, context_before=context_before, tile=tile + ): + assert qwen4_exp._prefill_causal_mask_is_exact( + pos_start=context_before + r0, rows=r1 - r0, total_keys=keys + ) + + +def test_tiled_causal_string_matches_untiled_and_tiled_dense(monkeypatch): + pos_start, rows, total = 32, 32, 64 + rng = np.random.default_rng(11) + q = mx.array(rng.standard_normal((1, 2, rows, 16)).astype(np.float32)) + kv = mx.array(rng.standard_normal((1, 2, total, 16)).astype(np.float32)) + scale = 16 ** -0.5 + untiled = qwen4_exp._qsa_dense_attention( + q, kv, kv, mask=qwen4_exp._CAUSAL_MASK, scale=scale + ) + monkeypatch.setenv(QUERY_TILE_ENV, "8") + tiled_causal = qwen4_exp._qsa_dense_attention( + q, kv, kv, mask=qwen4_exp._CAUSAL_MASK, scale=scale + ) + tiled_dense = qwen4_exp._qsa_dense_attention( + q, kv, kv, mask=_lane_mask(pos_start, rows, total), scale=scale + ) + mx.eval(untiled, tiled_causal, tiled_dense) + assert qwen4_exp._QSA_PREFILL_COUNTS.get("query_tile") == 2 + assert np.allclose(np.array(tiled_causal), np.array(untiled), atol=1e-6) + assert np.allclose(np.array(tiled_causal), np.array(tiled_dense), atol=1e-6) + + +# --------------------------------------------------------------------------- +# 4. Counters, gating and the loud refusal +# --------------------------------------------------------------------------- + + +class _FakeSdpa: + """Stand-in for a build that DOES have the fused kernels. + + ``fail_on`` refuses a whole mask kind. ``refuse`` is the interesting + one: a predicate over the call's own geometry, because that is how MLX + actually refuses -- ``use_fallback`` reads the query length, the GQA + factor and the head dims, so one build serves a 4,096-row prefill chunk + and refuses a 4-row verify step at the very same head dim and mask kind. + """ + + def __init__(self, fail_on=(), refuse=None): + self.calls = [] + self.geoms = [] + self.fail_on = set(fail_on) + self.refuse = refuse + + def __call__(self, q, k, v, *, scale, mask=None, force_fused=False, **kw): + kind = "causal" if isinstance(mask, str) else "bool" + self.calls.append((kind, bool(force_fused))) + self.geoms.append((kind, bool(force_fused), int(q.shape[2]))) + if force_fused and ( + kind in self.fail_on + or (self.refuse is not None and self.refuse(kind, q, k)) + ): + raise ValueError( + f"no fused kernel for {kind} at query length {int(q.shape[2])}" + ) + return mx.zeros(q.shape, dtype=q.dtype) + + +#: The served geometry: 24 query heads over 2 kv heads at head_dim 256. +PROD_Q_HEADS, PROD_KV_HEADS, PROD_HEAD_DIM = 24, 2, 256 +#: MLX's vector kernel caps ``q_len * gqa_factor`` at 32 and its full kernel +#: wants ``q_len > 8``; at GQA 12 that is a dead band at q_len 3..8, which is +#: where an MTP verify step (4 rows) lands and a prefill chunk never does. +VERIFY_ROWS, CHUNK_ROWS = 4, 64 + + +def _prod_qkv(rows: int, total: int): + q = mx.zeros((1, PROD_Q_HEADS, rows, PROD_HEAD_DIM), dtype=mx.bfloat16) + kv = mx.zeros((1, PROD_KV_HEADS, total, PROD_HEAD_DIM), dtype=mx.bfloat16) + return q, kv + + +def _refuse_short_queries(kind, q, k): + """The installed MLX's rule, in one line: the vector kernel is the only + one offered below ``q_len`` 9, and it caps ``q_len * gqa`` at 32.""" + + gqa = int(q.shape[1]) // int(k.shape[1]) + return int(q.shape[2]) <= 8 and int(q.shape[2]) * gqa > 32 + + +def _install(monkeypatch, fake): + monkeypatch.setattr(mx.fast, "scaled_dot_product_attention", fake) + + +def test_flag_off_never_forces_fused_and_never_builds_a_string(monkeypatch): + fake = _FakeSdpa() + _install(monkeypatch, fake) + q = mx.zeros((1, 2, 4, 256), dtype=mx.bfloat16) + kv = mx.zeros((1, 2, 4, 256), dtype=mx.bfloat16) + qwen4_exp._qsa_dense_attention(q, kv, kv, mask=_lane_mask(0, 4, 4), scale=1.0) + assert fake.calls == [("bool", False)] + assert "mask_fuse_causal" not in qwen4_exp._QSA_PREFILL_COUNTS + assert "mask_fuse_bool" not in qwen4_exp._QSA_PREFILL_COUNTS + + +def test_counters_split_causal_and_bool(monkeypatch): + _arm(monkeypatch) + fake = _FakeSdpa() + _install(monkeypatch, fake) + q = mx.zeros((1, 2, 4, 256), dtype=mx.bfloat16) + kv = mx.zeros((1, 2, 4, 256), dtype=mx.bfloat16) + qwen4_exp._qsa_dense_attention( + q, kv, kv, mask=qwen4_exp._CAUSAL_MASK, scale=1.0 + ) + qwen4_exp._qsa_dense_attention(q, kv, kv, mask=_lane_mask(0, 4, 4), scale=1.0) + counts = qwen4_exp._QSA_PREFILL_COUNTS + assert counts.get("mask_fuse_causal") == 1 + assert counts.get("mask_fuse_bool") == 1 + assert "mask_fuse_unavailable" not in counts + assert "mask_fuse_dense_causal" not in counts + assert "mask_fuse_dense_bool" not in counts + # ONE forced call per arm: the capability question is the real call, at + # the real geometry, so there is no synthetic probe to answer it wrong. + assert fake.calls == [("causal", True), ("bool", True)] + + +def test_engagement_is_announced_once_naming_the_class(monkeypatch): + """A serving process has no MTPLX_QSA_PREFILL_DEBUG receipt, and the + absence of a refusal line is not evidence of engagement -- so the first + class that fuses says so, once, on stderr.""" + + _arm(monkeypatch) + fake = _FakeSdpa(refuse=_refuse_short_queries) + _install(monkeypatch, fake) + err = io.StringIO() + with contextlib.redirect_stderr(err): + q_small, kv_small = _prod_qkv(VERIFY_ROWS, VERIFY_ROWS) + qwen4_exp._qsa_dense_attention( + q_small, kv_small, kv_small, mask=qwen4_exp._CAUSAL_MASK, scale=1.0 + ) + q_big, kv_big = _prod_qkv(CHUNK_ROWS, CHUNK_ROWS) + for _ in range(3): + qwen4_exp._qsa_dense_attention( + q_big, kv_big, kv_big, mask=qwen4_exp._CAUSAL_MASK, scale=1.0 + ) + qwen4_exp._qsa_dense_attention( + q_big, + kv_big, + kv_big, + mask=_lane_mask(0, CHUNK_ROWS, CHUNK_ROWS), + scale=1.0, + ) + lines = [line for line in err.getvalue().splitlines() if line.strip()] + engaged = [line for line in lines if "engaged" in line] + assert len(engaged) == 1, lines + assert f"causal-mask q_len {CHUNK_ROWS}" in engaged[0] + assert f"head_dim {PROD_HEAD_DIM} bfloat16" in engaged[0] + # The refused verify class is reported too, and separately. + assert len(lines) == 2, lines + assert f"q_len {VERIFY_ROWS}" in lines[0] + + +def test_s1_decode_rows_never_take_the_forced_route(monkeypatch): + _arm(monkeypatch) + fake = _FakeSdpa() + _install(monkeypatch, fake) + q = mx.zeros((1, 2, 1, 256), dtype=mx.bfloat16) + kv = mx.zeros((1, 2, 4, 256), dtype=mx.bfloat16) + qwen4_exp._qsa_dense_attention(q, kv, kv, mask=_lane_mask(3, 1, 4), scale=1.0) + assert fake.calls == [("bool", False)] + + +def test_refusal_is_loud_one_shot_and_per_kind(monkeypatch): + _arm(monkeypatch) + fake = _FakeSdpa(fail_on={"causal"}) + _install(monkeypatch, fake) + q = mx.zeros((1, 2, 4, 256), dtype=mx.bfloat16) + kv = mx.zeros((1, 2, 4, 256), dtype=mx.bfloat16) + err = io.StringIO() + with contextlib.redirect_stderr(err): + for _ in range(3): + qwen4_exp._qsa_dense_attention( + q, kv, kv, mask=qwen4_exp._CAUSAL_MASK, scale=1.0 + ) + qwen4_exp._qsa_dense_attention( + q, kv, kv, mask=_lane_mask(0, 4, 4), scale=1.0 + ) + message = err.getvalue() + refusals = [line for line in message.splitlines() if "armed but" in line] + assert len(refusals) == 1, message + assert "causal-mask q_len 4" in refusals[0] + assert "head_dim 256" in refusals[0] + counts = qwen4_exp._QSA_PREFILL_COUNTS + assert counts.get("mask_fuse_unavailable") == 1 + assert "mask_fuse_causal" not in counts + # All three causal calls went dense, and only the first paid a raise. + assert counts.get("mask_fuse_dense_causal") == 3 + # The bool arm is unaffected: its class was never asked about. + assert counts.get("mask_fuse_bool") == 1 + assert "mask_fuse_dense_bool" not in counts + # One forced causal attempt (refused, sticky for THAT class) + three + # dense fallbacks, then the bool call. Never a second forced attempt. + assert fake.calls == [ + ("causal", True), + ("causal", False), + ("causal", False), + ("causal", False), + ("bool", True), + ] + + +def test_refusal_is_native_on_a_build_without_fused_kernels(monkeypatch): + """No stub: the CPU stream has no fused kernel at all, so MLX raises.""" + + _arm(monkeypatch) + rows, total = 4, 8 + q, kv = _prod_qkv(rows, total) + err = io.StringIO() + with contextlib.redirect_stderr(err): + qwen4_exp._qsa_dense_attention( + q, kv, kv, mask=qwen4_exp._CAUSAL_MASK, scale=1.0 + ) + qwen4_exp._qsa_dense_attention( + q, kv, kv, mask=_lane_mask(total - rows, rows, total), scale=1.0 + ) + message = err.getvalue() + assert "require a GPU (Metal) stream" in message + # Two classes learned -- one per mask kind at this one geometry -- and + # nothing said about any other shape. + assert set(qwen4_exp._PREFILL_MASK_FUSE_UNAVAILABLE) == { + ( + "causal", + PROD_HEAD_DIM, + PROD_HEAD_DIM, + "bfloat16", + rows, + PROD_Q_HEADS // PROD_KV_HEADS, + True, + ), + ( + "bool", + PROD_HEAD_DIM, + PROD_HEAD_DIM, + "bfloat16", + rows, + PROD_Q_HEADS // PROD_KV_HEADS, + True, + ), + } + assert qwen4_exp._QSA_PREFILL_COUNTS.get("mask_fuse_unavailable") == 2 + + +def test_a_short_query_refusal_never_disarms_the_prefill_chunk(monkeypatch): + """The served-process defect, pinned. + + MLX offers only the VECTOR kernel below query length 9, and that kernel + caps ``q_len * gqa_factor`` at 32 -- so at this model's GQA 12 a 4-row + MTP verify step is refused while a prefill chunk of the very same head + dim, dtype and mask kind is served. The server's warmup ladder runs a + verify step before its first wide chunk; the benchmark driver does not. + A capability keyed by mask kind therefore measured the win in one + process and silently the control in the other. Keyed by shape class, + the verify refusal says nothing about the chunk. + """ + + _arm(monkeypatch) + fake = _FakeSdpa(refuse=_refuse_short_queries) + _install(monkeypatch, fake) + + # 1. The verify step goes first, exactly as the warmup ladder runs it. + q_small, kv_small = _prod_qkv(VERIFY_ROWS, VERIFY_ROWS) + with contextlib.redirect_stderr(io.StringIO()): + for _ in range(3): + qwen4_exp._qsa_dense_attention( + q_small, kv_small, kv_small, mask=qwen4_exp._CAUSAL_MASK, scale=1.0 + ) + + # 2. Then the wide chunk the flag is actually armed for. + q_big, kv_big = _prod_qkv(CHUNK_ROWS, CHUNK_ROWS) + qwen4_exp._qsa_dense_attention( + q_big, kv_big, kv_big, mask=qwen4_exp._CAUSAL_MASK, scale=1.0 + ) + qwen4_exp._qsa_dense_attention( + q_big, + kv_big, + kv_big, + mask=_lane_mask(0, CHUNK_ROWS, CHUNK_ROWS), + scale=1.0, + ) + + counts = qwen4_exp._QSA_PREFILL_COUNTS + assert counts.get("mask_fuse_causal") == 1 + assert counts.get("mask_fuse_bool") == 1 + assert counts.get("mask_fuse_dense_causal") == 3 + assert counts.get("mask_fuse_unavailable") == 1 + # One forced attempt at the refused class, never a second; both chunk + # classes forced and fused. + assert fake.geoms == [ + ("causal", True, VERIFY_ROWS), + ("causal", False, VERIFY_ROWS), + ("causal", False, VERIFY_ROWS), + ("causal", False, VERIFY_ROWS), + ("causal", True, CHUNK_ROWS), + ("bool", True, CHUNK_ROWS), + ] + + +def test_a_prefill_chunk_refusal_routes_only_that_class(monkeypatch): + """The other direction: a class MLX genuinely cannot serve goes dense + per call, and the classes it can serve stay fused.""" + + _arm(monkeypatch) + fake = _FakeSdpa(refuse=lambda kind, q, k: int(q.shape[2]) >= CHUNK_ROWS) + _install(monkeypatch, fake) + q_big, kv_big = _prod_qkv(CHUNK_ROWS, CHUNK_ROWS) + q_small, kv_small = _prod_qkv(VERIFY_ROWS, VERIFY_ROWS) + with contextlib.redirect_stderr(io.StringIO()): + for _ in range(3): + qwen4_exp._qsa_dense_attention( + q_big, kv_big, kv_big, mask=qwen4_exp._CAUSAL_MASK, scale=1.0 + ) + qwen4_exp._qsa_dense_attention( + q_small, kv_small, kv_small, mask=qwen4_exp._CAUSAL_MASK, scale=1.0 + ) + counts = qwen4_exp._QSA_PREFILL_COUNTS + assert counts.get("mask_fuse_dense_causal") == 3 + assert counts.get("mask_fuse_unavailable") == 1 + # The short class is a different class and is still fused. + assert counts.get("mask_fuse_causal") == 1 + assert fake.geoms == [ + ("causal", True, CHUNK_ROWS), + ("causal", False, CHUNK_ROWS), + ("causal", False, CHUNK_ROWS), + ("causal", False, CHUNK_ROWS), + ("causal", True, VERIFY_ROWS), + ] + + +def test_the_refusal_line_names_the_class_and_scopes_itself(monkeypatch): + _arm(monkeypatch) + fake = _FakeSdpa(refuse=_refuse_short_queries) + _install(monkeypatch, fake) + q, kv = _prod_qkv(VERIFY_ROWS, VERIFY_ROWS) + err = io.StringIO() + with contextlib.redirect_stderr(err): + qwen4_exp._qsa_dense_attention( + q, kv, kv, mask=_lane_mask(0, VERIFY_ROWS, VERIFY_ROWS), scale=1.0 + ) + message = err.getvalue() + assert message.count(MASK_FUSE_ENV) == 1, message + assert "engaged" not in message + # WHICH class. + assert f"bool-mask q_len {VERIFY_ROWS}" in message + assert f"GQA {PROD_Q_HEADS // PROD_KV_HEADS}" in message + assert f"head_dim {PROD_HEAD_DIM} bfloat16" in message + # And that it is per-class, not the process-wide disarm it used to be. + assert "THAT shape class only" in message + assert "NOT " in message and "process-wide" in message + # MLX's own reason survives into the line. + assert "no fused kernel for bool at query length 4" in message + + +def test_shape_class_reads_what_mlx_reads_and_nothing_else(monkeypatch): + """A class is the geometry MLX's rules look at -- key length beyond the + ``q_len <= k_len`` test is not one of them, and the query length is.""" + + kind = "causal" + q, kv = _prod_qkv(CHUNK_ROWS, 4096) + q2, kv2 = _prod_qkv(CHUNK_ROWS, 8192) + assert qwen4_exp._prefill_mask_fuse_class( + kind, q, kv, kv + ) == qwen4_exp._prefill_mask_fuse_class(kind, q2, kv2, kv2) + q3, kv3 = _prod_qkv(VERIFY_ROWS, 4096) + assert qwen4_exp._prefill_mask_fuse_class( + kind, q3, kv3, kv3 + ) != qwen4_exp._prefill_mask_fuse_class(kind, q, kv, kv) + # Mask kind, dtype and the GQA factor all split classes too. + assert qwen4_exp._prefill_mask_fuse_class( + "bool", q, kv, kv + ) != qwen4_exp._prefill_mask_fuse_class(kind, q, kv, kv) + wide_kv = mx.zeros((1, PROD_Q_HEADS, 4096, PROD_HEAD_DIM), dtype=mx.bfloat16) + assert qwen4_exp._prefill_mask_fuse_class( + kind, q, wide_kv, wide_kv + ) != qwen4_exp._prefill_mask_fuse_class(kind, q, kv, kv) + + +def test_refusal_printing_is_capped_but_counting_is_not(monkeypatch): + """A build that refuses everything must not own stderr for the life of + the process; the counters keep the full tally.""" + + _arm(monkeypatch) + fake = _FakeSdpa(fail_on={"causal"}) + _install(monkeypatch, fake) + classes = qwen4_exp._MASK_FUSE_REFUSAL_PRINT_LIMIT + 3 + err = io.StringIO() + with contextlib.redirect_stderr(err): + for i in range(classes): + rows = CHUNK_ROWS + i # a new shape class each time + q, kv = _prod_qkv(rows, rows) + qwen4_exp._qsa_dense_attention( + q, kv, kv, mask=qwen4_exp._CAUSAL_MASK, scale=1.0 + ) + message = err.getvalue() + refusals = [line for line in message.splitlines() if "armed but" in line] + assert len(refusals) == qwen4_exp._MASK_FUSE_REFUSAL_PRINT_LIMIT, message + assert "further shape-class refusals are counted but not printed" in message + counts = qwen4_exp._QSA_PREFILL_COUNTS + assert counts.get("mask_fuse_unavailable") == classes + assert counts.get("mask_fuse_dense_causal") == classes + + +def test_armed_flag_on_an_unavailable_build_still_returns_the_dense_answer( + monkeypatch, +): + _arm(monkeypatch) + pos_start, rows, total = 8, 8, 16 + rng = np.random.default_rng(3) + q = mx.array(rng.standard_normal((1, 2, rows, 16)).astype(np.float32)) + kv = mx.array(rng.standard_normal((1, 2, total, 16)).astype(np.float32)) + with contextlib.redirect_stderr(io.StringIO()): + armed = qwen4_exp._qsa_dense_attention( + q, kv, kv, mask=_lane_mask(pos_start, rows, total), scale=0.25 + ) + qwen4_exp._prefill_mask_fuse_enabled.cache_clear() + monkeypatch.delenv(MASK_FUSE_ENV, raising=False) + stock = qwen4_exp._qsa_dense_attention( + q, kv, kv, mask=_lane_mask(pos_start, rows, total), scale=0.25 + ) + mx.eval(armed, stock) + assert np.array_equal(np.array(armed), np.array(stock)) + + +# --------------------------------------------------------------------------- +# 5. Knobs stay where the harness can reach them +# --------------------------------------------------------------------------- + + +def test_sparse_crossover_knob_is_documented_and_unchanged(monkeypatch): + from mtplx.profiles import MODEL_RUNTIME_ENV_OVERRIDE_KEYS + + assert "MTPLX_QSA_PREFILL_MIN_CONTEXT" in MODEL_RUNTIME_ENV_OVERRIDE_KEYS + assert "MTPLX_QSA_PREFILL_FLASH_MIN_CONTEXT" in MODEL_RUNTIME_ENV_OVERRIDE_KEYS + assert "MTPLX_QSA_PREFILL_DEBUG" in MODEL_RUNTIME_ENV_OVERRIDE_KEYS + monkeypatch.delenv("MTPLX_QSA_PREFILL_MIN_CONTEXT", raising=False) + assert qwen4_exp._qsa_prefill_min_context() == 32_768 + monkeypatch.setenv("MTPLX_QSA_PREFILL_MIN_CONTEXT", "8192") + assert qwen4_exp._qsa_prefill_min_context() == 8192 + # Floor: below the 2,048-token budget the indexer has nothing to select. + monkeypatch.setenv("MTPLX_QSA_PREFILL_MIN_CONTEXT", "512") + assert qwen4_exp._qsa_prefill_min_context() == 2049 + + + + +# --------------------------------------------------------------------------- +# 6. The routing decision at the real call site +# --------------------------------------------------------------------------- + +#: A miniature of the production QSA geometry. block_topk = 8 // 2 = 4, so +#: the trivially-complete frontier is (4 + 1) * 2 - 1 = 9 tokens -- the same +#: arithmetic the 2,048-budget pack runs at 2,051. +TINY = dict( + hidden_size=64, + num_attention_heads=2, + num_key_value_heads=1, + head_dim=32, + indexer_n_heads=2, + indexer_kv_heads=1, + indexer_head_dim=16, + indexer_budget=8, + indexer_compress_ratio=2, +) +TINY_FRONTIER = (TINY["indexer_budget"] // TINY["indexer_compress_ratio"] + 1) * TINY[ + "indexer_compress_ratio" +] - 1 + + +def _tiny_attention(): + mx.random.seed(20260902) + layer = qwen4_exp.Attention(qwen4_exp.TextArgs(**TINY)) + layer.eval() + mx.eval(layer.parameters()) + return layer + + +class _MaskRecorder(_FakeSdpa): + """Records the mask each SDPA call receives, and returns real numbers.""" + + def __call__(self, q, k, v, *, scale, mask=None, force_fused=False, **kw): + kind = "causal" if isinstance(mask, str) else ("none" if mask is None else "bool") + self.calls.append((kind, bool(force_fused))) + if force_fused and kind in self.fail_on: + raise ValueError(f"no fused kernel for {kind}") + return _REAL_SDPA(q, k, v, scale=scale, mask=mask) + + +_REAL_SDPA = mx.fast.scaled_dot_product_attention + + +@pytest.mark.parametrize( + "rows, expect_kind", + [ + (TINY_FRONTIER - 1, "causal"), # T <= frontier: no selection at all + (TINY_FRONTIER + 1, "bool"), # past it: a real top-k selection + ], +) +def test_attention_routes_by_the_trivially_complete_frontier( + monkeypatch, rows, expect_kind +): + _arm(monkeypatch) + layer = _tiny_attention() + recorder = _MaskRecorder() + _install(monkeypatch, recorder) + x = (mx.random.normal((1, rows, TINY["hidden_size"])) * 0.3).astype(mx.bfloat16) + cache = qwen4_exp.QSACache(compress_ratio=layer.indexer.ratio) + with contextlib.redirect_stderr(io.StringIO()): + mx.eval(layer(x, cache)) + kinds = [kind for kind, _ in recorder.calls] + assert expect_kind in kinds, kinds + counts = qwen4_exp._QSA_PREFILL_COUNTS + if expect_kind == "causal": + assert counts.get("mask_causal_eligible") == 1 + else: + assert "mask_causal_eligible" not in counts + + +def test_flag_off_keeps_the_dense_causal_tensor_and_the_same_answer(monkeypatch): + """Flag off is byte-identical, and the armed causal arm sees the same set.""" + + layer = _tiny_attention() + rows = TINY_FRONTIER - 1 + x = (mx.random.normal((1, rows, TINY["hidden_size"])) * 0.3).astype(mx.bfloat16) + + recorder = _MaskRecorder() + _install(monkeypatch, recorder) + off = layer(x, qwen4_exp.QSACache(compress_ratio=layer.indexer.ratio)) + mx.eval(off) + assert recorder.calls == [("bool", False)] + + _arm(monkeypatch) + recorder2 = _MaskRecorder() + _install(monkeypatch, recorder2) + with contextlib.redirect_stderr(io.StringIO()): + on = layer(x, qwen4_exp.QSACache(compress_ratio=layer.indexer.ratio)) + mx.eval(on) + # The recorder stands in for a build that HAS the fused kernel, so the + # one forced call is the real one -- with the STRING and never a tensor. + # Its return value is real MLX math, so the equality below is the + # visible-set claim, evaluated end to end through the layer. + assert recorder2.calls == [("causal", True)] + assert np.array_equal( + np.array(off.astype(mx.float32)), np.array(on.astype(mx.float32)) + ) From 254973397c66657cff6adbc18268a686f6f2e08f Mon Sep 17 00:00:00 2001 From: davidtai Date: Mon, 7 Sep 2026 08:10:24 -0500 Subject: [PATCH 02/30] perf(qwen4): QSA split-K sparse-GQA decode lane for Qwen3.8 Flash-Next The fourth of PR 391's five unlanded lanes: MTPLX_QSA_SPARSE_DECODE, the native split-K sparse-GQA attention for the M=4 fixed verify. It reads the selected KV rows of the fixed QSA cache once per verify cycle instead of materializing a gathered [1,2,4,2052,256] K/V pair per layer, which is where the shipped lane's bytes are. Rounding class: fp32 online softmax over the exact visible set. - kernels/qsa_sparse_decode.py + native_extensions/qsa_sparse_gqa (package mtplx_native_qsa: the split-K Metal kernel, steel headers and a nanobind binding); mtplx/native loads it, runtime_options reads MTPLX_QSA_SPARSE_DECODE (+_TILE 128:32, +_SPLITS 17); the old MTPLX_FABLE_* names are honoured as aliases when the new key is unset. - graphbank.TensorOffsetQSACache validates the lane ONCE at cache install (a real parity probe, outside any mx.compile trace); the twin re-promotion sites and the compiled verify_step carry it, and the verify body asserts the lane is in the traced graph. models/qwen4_exp routes the fixed-capacity verify width to the kernel (QSAIndexer._sparse_decode_route) or declines to stock for a request shape it cannot serve. - Server auto-arm: default ON for the fixed-M4 pack ONLY when the native extension is built; a wheel without mtplx_native_qsa declines to stock with a logged verdict and still serves. An explicit MTPLX_QSA_SPARSE_DECODE=1 reaches the fail-closed install (armed and unbuilt raises). Registered in the boot-time runtime-env validator; kill switch through the existing pop loop. - scripts/bundle_native_runtime_wheel.py signs and packages mtplx_native_qsa alongside mtplx_qsa_kernels (Developer ID, hardened runtime, secure timestamp), with tests. - The mask-fuse refusal test now accepts either MLX build's native wording: the lane logs a version-independent per-class line and never raises under default arming (it falls to the stock dense SDPA). Load-time parity on stock mlx 0.32.2 with the native kernel built: vs the fp32 reference worst rel_l2 3.1e-05 with the top-1 token identical, vs the stock gather path rel_l2 4.6e-03 (rounding class), across the 4093 and 2052 probe cells that stand in for the 16K and 261,120 serving regimes. CPU tests (venv mlx 0.32.2): tests/test_qsa_sparse_decode.py, tests/test_qsa_sparse_decode_wiring.py, tests/test_qsa_sparse_gqa_native.py and tests/test_bundle_native_runtime_wheel.py all green. --- docs/perf/qwen38-391-remainder.md | 29 +- mtplx/graphbank.py | 88 ++ mtplx/kernels/qsa_sparse_decode.py | 1213 +++++++++++++++++ mtplx/models/qwen4_exp.py | 227 +++ mtplx/native/__init__.py | 716 ++++++++++ mtplx/profiles.py | 3 + mtplx/runtime_options.py | 135 ++ mtplx/server/openai.py | 45 +- native_extensions/qsa_sparse_gqa/.gitignore | 6 + .../qsa_sparse_gqa/CMakeLists.txt | 168 +++ native_extensions/qsa_sparse_gqa/bindings.cpp | 58 + .../mtplx_native_qsa/__init__.py | 15 + .../qsa_sparse_gqa/pyproject.toml | 8 + native_extensions/qsa_sparse_gqa/setup.py | 17 + .../sparse_gqa/qsa_sparse_gqa.cpp | 278 ++++ .../sparse_gqa/qsa_sparse_gqa.h | 37 + .../sparse_gqa/qsa_sparse_gqa.metal | 58 + .../sparse_gqa/qsa_sparse_gqa_decode.cpp | 354 +++++ .../sparse_gqa/qsa_sparse_gqa_decode.h | 52 + .../sparse_gqa/qsa_sparse_gqa_decode_params.h | 71 + .../sparse_gqa/qsa_sparse_gqa_params.h | 42 + .../sparse_gqa/steel_qsa_sparse_gqa.h | 339 +++++ .../sparse_gqa/steel_qsa_sparse_gqa_decode.h | 431 ++++++ scripts/bundle_native_runtime_wheel.py | 123 +- tests/test_bundle_native_runtime_wheel.py | 99 ++ tests/test_qsa_sparse_decode.py | 727 ++++++++++ tests/test_qsa_sparse_decode_wiring.py | 771 +++++++++++ tests/test_qsa_sparse_gqa_native.py | 324 +++++ tests/test_qwen4_prefill_mask_fuse.py | 9 +- 29 files changed, 6410 insertions(+), 33 deletions(-) create mode 100644 mtplx/kernels/qsa_sparse_decode.py create mode 100644 mtplx/native/__init__.py create mode 100644 native_extensions/qsa_sparse_gqa/.gitignore create mode 100644 native_extensions/qsa_sparse_gqa/CMakeLists.txt create mode 100644 native_extensions/qsa_sparse_gqa/bindings.cpp create mode 100644 native_extensions/qsa_sparse_gqa/mtplx_native_qsa/__init__.py create mode 100644 native_extensions/qsa_sparse_gqa/pyproject.toml create mode 100644 native_extensions/qsa_sparse_gqa/setup.py create mode 100644 native_extensions/qsa_sparse_gqa/sparse_gqa/qsa_sparse_gqa.cpp create mode 100644 native_extensions/qsa_sparse_gqa/sparse_gqa/qsa_sparse_gqa.h create mode 100644 native_extensions/qsa_sparse_gqa/sparse_gqa/qsa_sparse_gqa.metal create mode 100644 native_extensions/qsa_sparse_gqa/sparse_gqa/qsa_sparse_gqa_decode.cpp create mode 100644 native_extensions/qsa_sparse_gqa/sparse_gqa/qsa_sparse_gqa_decode.h create mode 100644 native_extensions/qsa_sparse_gqa/sparse_gqa/qsa_sparse_gqa_decode_params.h create mode 100644 native_extensions/qsa_sparse_gqa/sparse_gqa/qsa_sparse_gqa_params.h create mode 100644 native_extensions/qsa_sparse_gqa/sparse_gqa/steel_qsa_sparse_gqa.h create mode 100644 native_extensions/qsa_sparse_gqa/sparse_gqa/steel_qsa_sparse_gqa_decode.h create mode 100644 tests/test_qsa_sparse_decode.py create mode 100644 tests/test_qsa_sparse_decode_wiring.py create mode 100644 tests/test_qsa_sparse_gqa_native.py diff --git a/docs/perf/qwen38-391-remainder.md b/docs/perf/qwen38-391-remainder.md index 8fe9011ff..8cab716a6 100644 --- a/docs/perf/qwen38-391-remainder.md +++ b/docs/perf/qwen38-391-remainder.md @@ -22,10 +22,25 @@ production geometry and are re-measured in the GPU phase. | QSA sparse split-K decode | `MTPLX_QSA_SPARSE_DECODE` | Native split-K sparse-GQA decode kernel that reads the selected KV rows once instead of materializing a gathered `[1,2,4,2052,256]` K/V pair per QSA layer per verify cycle. | rounding-class | ~-1.46 ms/cycle at 16K | | graph-build overlap | `MTPLX_QWEN4_GRAPH_BUILD_OVERLAP` | Submits the PLE-independent prefix of the fixed-M4 verify graph early so its GPU work overlaps the ~1.4 ms/cycle host build of the rest. | exact | ~1.4 ms/cycle (~4.5% at N=3) | -Status in this change: HC_M4, the prefill mask fuse, and the QSA prefill query -tile are ported and CPU-tested. The QSA sparse split-K decode and the graph-build -overlap are covered in the port report (`.benchmark-artifacts/over100-reports/ -remainder-port-report.md`) — the decode lane's substrate is present upstream and -is a native+wiring port whose kernel needs the GPU lock to build and parity-probe; -the graph-build overlap is blocked because upstream re-landed the fixed-M4 verify -as a single compiled graph, without the prefix/suffix split the lane rides. +Landed here: HC_M4, the prefill mask fuse, and the QSA prefill query tile +(commit 1); the QSA sparse split-K decode (commit 2, with its `mtplx_native_qsa` +extension added to the wheel signer and a Developer-ID-signing test). The QSA +decode lane's native kernel is built and its load-time parity probe passes on +stock mlx 0.32.2: vs the fp32 reference worst rel_l2 3.1e-05 with top-1 token +identical (1.0000), vs the shipped stock-gather path rel_l2 4.6e-03 (rounding +class), across the long-context (4093) and just-past-crossover (2052) cells that +stand in for the 16K and 261,120 serving regimes. + +The mask fuse engages on stock mlx 0.32.2 (no custom MLX build): its metallib +carries `steel_attention_bfloat16_bq32_bk16_bd256_wm4_wn1_maskbool_` and its +`_maskbfloat16` / dsplit variants — the exact head_dim-256 fused kernels the lane +uses. It declines per shape class only for the verify-width dead band (q_len 3..8 +at GQA 12, head_dim 256), by design, and never raises under default arming (it +falls to the stock dense SDPA and logs the per-class refusal). + +The graph-build overlap is NOT ported: upstream re-landed the fixed-M4 verify as a +single compiled graph without the prefix/suffix split the lane rides. A one-page +feasibility (contribution, exposed host budget, minimal single-graph design, +go/no-go) is in the port report +(`.benchmark-artifacts/over100-reports/remainder-port-report.md`); the +recommendation is to defer it to a dedicated split-compilation task. diff --git a/mtplx/graphbank.py b/mtplx/graphbank.py index 1a966756f..fc6cc5e01 100644 --- a/mtplx/graphbank.py +++ b/mtplx/graphbank.py @@ -19,6 +19,9 @@ from .attention_context import attention_phase from .gdn_capture import resolve_gdn_capture_backend +# Module-level so a test (and the twin-construction guard) can see one name; +# read at each call so the process-frozen flag is honoured and monkeypatchable. +from .runtime_options import qsa_sparse_decode_enabled def _prepare_fixed_m4_materialized( @@ -626,6 +629,7 @@ def __init__( rows_gather_enabled: bool = False, rows_gather_min_context: int = 0, fused_rows_gather_kv_m4: bool = False, + qsa_sparse_decode: bool = False, ) -> None: self.kv = kv self.raw_keys = raw_keys @@ -637,6 +641,62 @@ def __init__( self.rows_gather_enabled = bool(rows_gather_enabled) self.rows_gather_min_context = max(0, int(rows_gather_min_context)) self.fused_rows_gather_kv_m4 = bool(fused_rows_gather_kv_m4) + # MTPLX_QSA_SPARSE_DECODE: the native split-K sparse-GQA attention + # lane. Validated ONCE, here, at cache install: this is model-build + # time and outside any mx.compile trace, which is what lets the + # install run a real parity probe. The indexer reads the row count + # below and never re-derives the decision, so one trace of the verify + # graph cannot disagree with the next about which attention it holds. + # + # The gate is asymmetric on purpose. install() RAISES when the + # contract cannot be met (an armed flag that cannot apply is a + # configuration error, including a missing native extension) and + # DISABLES when the numerical probe fails (this kernel is + # rounding-class; a parity miss is a measurement, and turning a + # measurement into an outage helps nobody). The server auto-arm only + # STAMPS the default when the extension is built, so a wheel without + # it serves the stock decode path; an explicit operator export of + # MTPLX_QSA_SPARSE_DECODE=1 still reaches this fail-closed install. + self.qsa_sparse_decode = bool(qsa_sparse_decode) + self.qsa_sparse_decode_rows = 0 + # An armed flag that reaches a cache constructed WITHOUT it is the + # armed-but-inert failure mode, and it is silent: the cache carries + # qsa_sparse_decode_rows = 0 and every routing decision declines. + # Every construction site that can be reached with the flag armed + # passes it through (from_qsa_cache reads the frozen flag), so a site + # that forgot is a bug in THIS file and dies here. + if qsa_sparse_decode_enabled() and not self.qsa_sparse_decode: + raise RuntimeError( + "MTPLX_QSA_SPARSE_DECODE is armed but this QSA cache was " + "constructed without the lane; the construction site did not " + "pass qsa_sparse_decode through, so the kernel would be inert " + "on this cache" + ) + if self.qsa_sparse_decode: + from .kernels import qsa_sparse_decode as _qsa_sparse + + if not _qsa_sparse.install( + self.kv.keys, + self.kv.values, + compress_ratio=self.ratio, + verify=True, + ): + # install() recorded the measured deltas before returning + # False; this turns them into a build failure instead of a + # silent revert to the stock chain: an armed arm that runs the + # stock chain is worse than an outage because it looks like a + # result. + raise RuntimeError( + "MTPLX_QSA_SPARSE_DECODE is armed but the split-K lane " + "declined to install: " + + ( + _qsa_sparse.disabled_reason() + or "the lane returned no verdict" + ) + + " -- read the deltas off the [mtplx] qsa_sparse_decode " + "stderr line, then unarm the flag deliberately" + ) + self.qsa_sparse_decode_rows = _qsa_sparse.VERIFY_ROWS @staticmethod def _fixed_bank(value: mx.array, capacity: int, axis: int) -> mx.array: @@ -684,6 +744,10 @@ def from_qsa_cache( rows_gather_enabled = _qsa_gather_enabled() rows_gather_min_context = _qsa_gather_min_context() rows_gather = rows_gather_enabled and offset >= rows_gather_min_context + # MTPLX_QSA_SPARSE_DECODE: read the process-frozen flag here so the + # cache validates the native split-K lane once, at install; the + # indexer never re-derives it. Passed to __init__ below. + qsa_sparse_decode = qsa_sparse_decode_enabled() rows_gather_kv_m4 = entry.rows_gather_kv_m4 fused_rows_gather_kv_m4 = _env_enabled("MTPLX_QSA_M4_FUSED_KV_GATHER") if fused_rows_gather_kv_m4: @@ -725,6 +789,7 @@ def from_qsa_cache( rows_gather_enabled=rows_gather_enabled, rows_gather_min_context=rows_gather_min_context, fused_rows_gather_kv_m4=fused_rows_gather_kv_m4, + qsa_sparse_decode=qsa_sparse_decode, ) @property @@ -3173,6 +3238,9 @@ def _ensure_shadow(self, cache: Any) -> None: rows_gather_enabled=entry.rows_gather_enabled, rows_gather_min_context=entry.rows_gather_min_context, fused_rows_gather_kv_m4=entry.fused_rows_gather_kv_m4, + # A twin that dropped this would silently revert to the + # stock QSA chain -- the armed-but-inert failure mode. + qsa_sparse_decode=entry.qsa_sparse_decode, ) elif kind == VERIFY_SPEC_KIND_FULL_ATTN: if isinstance(entry, TensorOffsetKVCache): @@ -3331,6 +3399,16 @@ def verify_step(input_ids, *args): entry.cache[slot] = state_in[pos + slot] pos += n_leaves # (2) The existing runtime forward, on shadow containers only. + # + # Sample the sparse-decode lane's route counters ACROSS the + # forward. The routing decision is host-side and happens in THIS + # python body, so a trace that ends with neither a route hit nor a + # short-context decline is an armed flag that is not in the graph + # -- and the graph is what the next few hundred cycles replay. + # Raising here costs one trace; the alternative was a whole window. + from .kernels import qsa_sparse_decode as _qsa_sparse_lane + + sparse_route_before = _qsa_sparse_lane.route_snapshot() with attention_phase("decode_verify"): result = live._runtime_forward( input_ids, @@ -3339,6 +3417,13 @@ def verify_step(input_ids, *args): hidden_variant=hidden_variant, compiled_aux=compiled_aux, ) + # Trace-time engagement check on the graph this body just built, + # fatal because the graph is what the next few hundred cycles + # replay: an armed sparse-decode flag that never routed is an inert + # arm. A no-op when the lane is unarmed (route_snapshot stays 0). + _qsa_sparse_lane.assert_traced( + length, before=sparse_route_before, where="compiled verify" + ) logits, hidden, captures = result # (3) Read every leaf back out and return it explicitly. captures_flat: list[Any] = [] @@ -3711,6 +3796,9 @@ def _parity2_clone_cache(self, cache: Any, bucket: int) -> list[Any]: rows_gather_enabled=entry.rows_gather_enabled, rows_gather_min_context=entry.rows_gather_min_context, fused_rows_gather_kv_m4=entry.fused_rows_gather_kv_m4, + # A twin that dropped this would silently revert to the + # stock QSA chain -- the armed-but-inert failure mode. + qsa_sparse_decode=entry.qsa_sparse_decode, ) elif kind == VERIFY_SPEC_KIND_FULL_ATTN: if isinstance(entry, TensorOffsetKVCache): diff --git a/mtplx/kernels/qsa_sparse_decode.py b/mtplx/kernels/qsa_sparse_decode.py new file mode 100644 index 000000000..d7503b0fe --- /dev/null +++ b/mtplx/kernels/qsa_sparse_decode.py @@ -0,0 +1,1213 @@ +"""Split-K native sparse-GQA attention for the Qwen3.8 QSA DECODE lanes. + +MTPLX_QSA_SPARSE_DECODE (M=4 fixed verify). The kernel itself lives in +``native_extensions/qsa_sparse_gqa`` and is reached through +``mtplx.native.qsa_sparse_gqa_decode``; this module is the lane: the gate, +the install probe, the engagement counters and the reference the probe +compares against. + +WHAT IT REPLACES, AND WHY THE CENSUS UNDERSTATES IT +--------------------------------------------------- +The retained fixed-M4 verify attends 4 query rows per QSA layer over the +indexer's selected top-512 pooled blocks. Per layer, per verify cycle, the +shipped lane issues six dispatches: + + 1 custom_kernel_..._qsa_m4_fused_kv_gather_c17408 [1050624,1,1] + 2 (Copy family) + 3 gemv_bfloat16_bm4_bn1_sm1_sn32_tm4_tn4_nc1 [129,1,96] scores + 4 block_softmax_float32 [52224,1,1] + 5 + 6 gemv_t_bfloat16_bm1_bn2_sm8_sn4_tm4_tn4_nc1 [8,1,96] P@V + +Dispatch (1) materialises ``k_sel``/``v_sel`` as ``[1, 2, 4, 2052, 256]`` +bf16 -- 8.40 MB each. So per layer the lane WRITES 16.8 MB of gathered K/V, +MLX then copies 8.4 MB more for the transposed score operand, and dispatches +(3) and (6) read 8.4 MB each back. Roughly 70 MB per layer, ~840 MB per +verify cycle, to attend 4 rows. + +The dispatch census's QSA row (446 MB/cycle, 232 GB/s) does NOT show this: +its cost model prices the gather at a flat 4.19 MB and gives the score, +softmax and P@V dispatches zero bytes, and the transposed copy lands in the +Copy family. Counted properly the QSA family moves ~1.05 GB in its 1.93 +ms/cycle, i.e. about 540 GB/s -- right at this machine's measured 544 GB/s +ceiling. **The lane is not bandwidth-starved; it is moving three times the +bytes it needs to.** That is the thing this kernel changes: it reads the +cache rows once, in place, and never materialises them. + +WHY SPLIT-K AND NOT THE PHASE-1 KERNEL +--------------------------------------- +The phase-1 (prefill) kernel parallelises over query rows: grid +``(qL, kv_heads, 1)``. At M=4 that is EIGHT threadgroups of 64 threads on a +40-core M5 Max, each walking all 2,051 selected keys. Phase 1's own design +note priced the fix as its own item, and MTPLX has the general finding +already: a hand-written +metal_kernel SDPA lost to stock at long N precisely because MLX's production +SDPA switches to a KV-split two-pass path there. So decode gets the KV-split +variant: ``(qL, kv_heads, n_splits)`` threadgroups accumulating independent +online-softmax states, then a merge pass. + +NUMERICS -- ROUNDING CLASS, HUMANEVAL-GATED +-------------------------------------------- +The visible set is IDENTICAL to the shipped lane's, slot for slot (the +kernel applies the shipped predicate ``block < (pos+1)//4`` to every slot of +``top_idx``, which is what the rows-gather token list does; it makes no +ordering assumption, because the selector hands through +``mx.argpartition``'s raw output unsorted). The ARITHMETIC is not the same: fp32 online softmax in +exp2, fp32 probabilities into an fp32 P@V, Steel-MMA reassociation of the +256-term score contraction, and one split-K rescale per row. + +So this lane is adopted on the same terms as ``MTPLX_QWEN4_HC_M4``: greedy +token agreement plus a full HumanEval run, never on a digest. The install +probe below is a numerical SANITY gate, not the quality gate -- it exists so +an armed flag that is quietly wrong disables itself instead of shipping. + +GATE DISCIPLINE -- REVISED 2026-09-02 AFTER AN ARMED-BUT-INERT WINDOW +--------------------------------------------------------------------- +The 2026-09-02 16 K window armed this flag and measured the CONTROL: control +and candidate response texts were byte-identical on both finished seeds, on a +kernel whose arithmetic is rounding class. Nothing in the run said so, +because every way the lane could decline was silent. The rule the program +owner set afterwards is: **a flag either works on every request path the +server accepts, or it fails loudly at install.** Concretely: + +* CONTRACT failure RAISES, as before. An armed flag on a pack the kernel + cannot serve is a configuration error -- that is how MTPLX_FUSED_HC_V3 came + to be armed-but-dead at M=4. +* PARITY failure still DISABLES *inside this module* -- :func:`install` + returns False and records the measured deltas, so the numbers survive for + the receipt -- but the CALLER + (``graphbank.TensorOffsetQSACache.__init__``) then RAISES, because an armed + arm that runs the stock chain is worse than an outage: it is a measurement + that looks like a result. Read the deltas off the stderr line and the + receipt, then unarm the flag deliberately. +* ROUTE narrowing RAISES for a genuinely wrong CONFIGURATION at the width the + flag arms -- a ratio that is not 4, a top-k that is not 512, a cache whose + wired row count disagrees with the module. Those are things an operator + chose, and they cannot be fixed by declining. +* ROUTE narrowing ROUTES for the shapes a server is entitled to send. Widths + the flag does not arm (prefill rows, the S=1 D3 route under a verify-only + arm), caches the lane never installed on, and -- the case that matters most + -- a context below :data:`SHORT_CONTEXT_TOKENS`. All return False; the last + two are COUNTED in :data:`_ROUTE_DECLINES` so "the flag did nothing" always + has a readable cause. +* An armed lane that is not in the traced verify graph RAISES from + :func:`assert_traced`, called inside the compiled verify body -- unless the + forward declined for short context, which is a legitimate stock lane. + +THE LANE ENGAGES ONLY ABOVE 2,052 TOKENS, AND THAT IS THE DESIGN +----------------------------------------------------------------- +The kernel's ABI is a fixed ``[M, 512]`` block selection, and 512 complete +pooled blocks exist only from ``(512 + 1) * 4 = 2,052`` tokens +(:data:`SHORT_CONTEXT_TOKENS`). Below that there is no smaller kernel to fall +to -- padding the selection would attend keys the shipped lane does not. +Context length is a per-REQUEST shape the server must accept, so a short +request routes to the stock attention, is counted, and prints nothing. A +HumanEval prompt or a 1 K cell is served, not refused; an armed 16 K decode +still has to bind or the assertions fire. This is the documented behaviour of +the flag, not a fallback. + +The threshold is on the value ATTENTION passes as ``total_tokens``, which on a +fixed-capacity cache is the BANK CAPACITY -- ``update_and_fetch`` returns the +whole backing. A 1,024-token prompt in a 2,048-token bank therefore has a +full 512-block budget and a context that has NOT crossed the boundary; the two +questions are different, and :func:`context_decline` asks both. Asking only +the first cost the 1 K cell of the 2026-09-02 served battery. +""" + +from __future__ import annotations + +import json +import logging +import math +import sys +from typing import Any, Dict, Optional, Tuple + +import mlx.core as mx + +logger = logging.getLogger(__name__) + + +def _emit(line: str) -> None: + """Put the lane's verdict where a benchmark log will actually see it. + + ``logger.info`` alone is INVISIBLE in a driver run -- the 2026-09-02 + window carried no engagement evidence at all, and the same log is missing + ``[qwen4-fixed-M4-verify]`` and ``[qwen4-compiled-MTP-prepare]`` (both + ``logger.info``) while it does carry ``[MTPLX_QWEN4_GRAPH_BUILD_OVERLAP] + armed:`` (a plain ``print``). The verify-glue lane fixed exactly this + items; this is the same fix for this lane. + """ + + logger.info("%s", line) + print(line, file=sys.stderr, flush=True) + +#: Production Qwen3.8 Flash-Next QSA geometry -- compiled into the metallib. +Q_HEADS = 24 +KV_HEADS = 2 +GQA = 12 +HEAD_DIM = 256 +TOP_K = 512 +COMPRESS_RATIO = 4 +VERIFY_ROWS = 4 +#: The kernel's own selected-token width; see ``mtplx.native`` for why this is +#: one less than the shipped lane's 2,052 and why the visible sets still agree. +SELECTED_TOKENS = TOP_K * COMPRESS_RATIO + (COMPRESS_RATIO - 1) + +#: bf16 has an 8-bit significand, so its relative spacing is 2**-8 and its +#: unit roundoff (round to nearest) is u = 2**-9. +_BF16_REL_ULP = 2.0**-8 +_BF16_UNIT_ROUNDOFF = 2.0**-9 + +# --------------------------------------------------------------------------- +# THE GATE, AND WHY IT IS TWO GATES AGAINST TWO REFERENCES +# +# The 2026-09-02 micro measured the kernel against the shipped path at +# max_abs 1.953e-3 (= 2**-9 exactly), rel_l2 4.78e-3, top-1 1.0000 -- and +# reported the SAME four significant figures for all twenty configurations: +# BK 64/128/256, DC 32/64, splits 4..32. DC changes the fp32 score +# contraction order and the split count changes the online-softmax merge +# tree, so if any of that delta were the kernel's it would move. It does not +# move at all. The delta is therefore a property of the REFERENCE. +# +# The shipped path carries two bf16 roundings this kernel does not: +# +# 1. ``mx.matmul(q_view, k_view)`` has bf16 inputs, so its output is bf16 -- +# the SCORES are rounded to bf16 before ``.astype(float32) * scale`` and +# the softmax. A relative score error u shifts a softmax logit by +# u*|x|, and perturbing logits by eps changes each probability by about +# (eps_i - ). With scaled logits of order 5 that is ~2*u*5 = 2e-2 +# relative on the probabilities. +# 2. ``probs.astype(bfloat16)`` before P@V: another u = 2e-3 relative. +# +# So (1) should dominate (2) by roughly an order of magnitude, and the +# measured 4.78e-3 sits between the two predictions -- which is why the +# attribution is TESTED by a reference ladder rather than asserted. +# +# The consequence for the gate: a threshold on kernel-vs-shipped is really a +# threshold on how much bf16 rounding the SHIPPED path does, which is not a +# property of this kernel and cannot be tightened by improving it. So: +# +# * the DECIDING numerical gate is against the fp32 reference, where the +# only differences left are fp32 reassociation and one bf16 store. It is +# tight, and it FAILS if the attribution above is wrong -- which is the +# point of stating it this way. +# * kernel-vs-shipped keeps a loose SANITY bound, an order of magnitude +# above what the shipped path's own bf16 quantisation implies. +# +# Neither is a quality gate. This kernel is rounding class; whether the +# difference matters is answered by model-level greedy-token agreement plus a +# full HumanEval run, exactly as for MTPLX_QWEN4_HC_M4. +# --------------------------------------------------------------------------- + +#: TIGHT, vs :func:`fp32_reference`. Derivation: both sides round the same +#: real number to bf16, so they differ by at most one bf16 ulp on any element +#: where the underlying fp32 values straddle a rounding boundary; fp32 +#: reassociation over <= 2051 terms contributes ~sqrt(2051)*2**-24 = 2.7e-6 +#: relative, three orders below a bf16 ulp, so it can only move an element +#: that already sits within 2.7e-6 of a boundary (about 1 in 4e4). Hence a +#: couple of ulp at the extreme and a relative L2 far below the 2**-9/sqrt(3) +#: = 1.1e-3 that a uniformly re-rounded output would give. +PARITY_FP32_MAX_ABS_ULPS = 2.0 +PARITY_FP32_MAX_REL_L2 = 5.0e-4 + +#: LOOSE, vs :func:`stock_reference` (the shipped path). This bounds the +#: shipped path's OWN bf16 score and probability quantisation, derived above +#: at ~2e-2 relative on the probabilities; 5e-2 leaves an order of magnitude +#: over the 4.78e-3 measured on 2026-09-02. It exists to catch a kernel that +#: is wrong by a factor, not to certify one that is right. +PARITY_SHIPPED_MAX_REL_L2 = 5.0e-2 + +#: Fraction of (head, row) pairs whose argmax over the head dimension must +#: still agree, against BOTH references. A coarse discrete statistic, +#: reported for continuity with the "top-1 agreement" the program asks for at +#: the kernel level; the DECIDING top-1 number is model-level greedy-token +#: agreement. Measured 1.0000 on every configuration. +PARITY_MIN_TOP1 = 0.98 + +#: Probe geometry: capacity only has to clear the dense/sparse crossover +#: (``total // 4 > 512``), and the split geometry does not depend on context +#: length at all -- it is a function of SELECTED_TOKENS and the tile -- so a +#: 4,096-token probe exercises the same grid a 17,408-token cycle does. +PROBE_CAPACITY = 4096 + +_COUNTS: Dict[str, int] = { + "verify_kernel": 0, + "probe_runs": 0, + "probe_failures": 0, + # Every ``TensorOffsetQSACache`` that bound the lane, shadow and parity + # twins included. The probe runs on the FIRST one only (the verdict is + # per process), so a value of ZERO with the flag armed is the whole + # finding: no cache carried the lane and the kernel could not have run. + "cache_installs": 0, + # Every routing decision that came back True, summed over the call sites + # in ``_ROUTE_SITES``. Trace-time, like ``verify_kernel``: the Python body + # of a compiled verify graph runs once per retrace, so read these as "did + # this lane get into the graph at all", never as a per-cycle count. + "route_hits": 0, + # Forwards that routed to the stock chain because of the REQUEST's own + # shape -- its context length, its row count, its block budget -- rather + # than the configuration. Counted separately because the engagement + # assertions accept these and only these: see :func:`context_decline`. + "request_declines": 0, +} + +#: Smallest context the lane can serve, in tokens. The kernel's ABI is a +#: fixed ``[M, TOP_K]`` selection, so it needs a FULL budget: ``TOP_K`` +#: complete pooled blocks exist only once the context reaches +#: ``(TOP_K + 1) * COMPRESS_RATIO`` = 2,052 tokens. Below that there is no +#: analogue -- not a smaller kernel, not a padded one -- so the stock chain +#: serves the request and the receipt says how often. +#: +#: This is exactly ``mtplx.native``'s +#: ``total_tokens // _COMPRESS_RATIO <= _TOP_K_BLOCKS`` boundary, restated in +#: tokens. ``tests/test_qsa_sparse_decode_wiring.py`` pins the two +#: together. +SHORT_CONTEXT_TOKENS = (TOP_K + 1) * COMPRESS_RATIO + +#: Largest context the kernel is instantiated for; mirrors +#: ``mtplx.native._MAX_CONTEXT``. +MAX_CONTEXT = 1_048_576 + +#: Observed extremes for the request-shape declines, so the receipt carries +#: the numbers without giving ``_ROUTE_DECLINES`` an unbounded key space (one +#: key per distinct context length would grow without bound in a server). +_DECLINE_EXTREMES: Dict[str, int] = {} + +#: ``site -> hits``. The 2026-09-02 window failed because the one call site +#: that could reach the verify width asked for a width the flag did not arm; +#: a single total would not have shown that, so the sites are named. +_ROUTE_SITES: Dict[str, int] = {} + +#: ``reason -> count`` for the routing narrowings that are NOT failures (a +#: width the flag does not arm, a growable cache the lane never installed on). +#: Recorded rather than merely returned, so "the flag did nothing" always has +#: a readable cause in the receipt. +_ROUTE_DECLINES: Dict[str, int] = {} + +#: ``None`` until the probe has run. ``""`` once it has passed. A non-empty +#: string is the reason the lane is disabled for this process. +_DISABLED_REASON: Optional[str] = None +_PROBE_REPORT: Dict[str, Any] = {} + + +class SparseDecodeContractError(RuntimeError): + """An armed flag met a request path it cannot serve. + + Raised, never swallowed: a lane that quietly declines makes the candidate + arm measure the control while its receipt claims otherwise, which is the + exact failure the 2026-09-02 window produced. + """ + + +def armed() -> bool: + """True when this process armed the flag.""" + + from mtplx.runtime_options import qsa_sparse_decode_enabled + + return bool(qsa_sparse_decode_enabled()) + + +def pending() -> bool: + """True while the install probe has not run yet in this process.""" + + return _DISABLED_REASON is None + + +def installed() -> bool: + """True when the probe ran and passed. False while pending, too.""" + + return _DISABLED_REASON == "" + + +def note_route_hit(site: str) -> None: + """One routing decision resolved to the kernel, at a named call site.""" + + _COUNTS["route_hits"] += 1 + _ROUTE_SITES[site] = _ROUTE_SITES.get(site, 0) + 1 + + +def note_route_decline(reason: str) -> None: + """One routing narrowing that is not an error, recorded by cause.""" + + _ROUTE_DECLINES[reason] = _ROUTE_DECLINES.get(reason, 0) + 1 + + +#: THE MIRROR. Every branch of +#: ``mtplx.native.qsa_sparse_gqa_decode_unsupported_reason`` that depends on +#: the REQUEST -- its context length, its row count, its block budget -- paired +#: with the stable key :func:`context_decline` reports it under. Everything +#: else in that function is CONFIGURATION (dtypes, shapes, the tile, the split +#: count, the device, the build) and raises. +#: +#: Two request-shape raises reached production before this list existed. The +#: first was ``k_eff != TOP_K`` (2026-09-02, HTTP 500 on every HumanEval +#: prompt). The second was the one this list is named for: a 1,024-token +#: prompt whose FIXED bank is 2,048 tokens has a full 512-block budget -- +#: ``k_eff`` is 512, so the first gate passed -- while the kernel's own +#: boundary is ``total_tokens // 4 > 512``, i.e. 2,052 tokens, so the call +#: died inside ``attention()`` with "the context has not crossed the +#: dense/sparse boundary". A partial mirror is how that happens, so +#: ``tests/test_qsa_sparse_decode_wiring.py`` pins this list against the +#: native source. +REQUEST_SHAPE_DECLINES = ( + ("empty_context", "total_tokens must describe a non-empty context"), + ("rows_exceed_context", "the query rows must fit inside total_tokens"), + ( + "context_exceeds_capacity", + "the logical token count exceeds the full K/V backing capacity", + ), + ( + "context_above_limit", + "the logical token count exceeds the production context limit", + ), + ("short_context", "the context has not crossed the dense/sparse boundary"), + # Not a native branch: the indexer's own budget, which the kernel's fixed + # [M, TOP_K] ABI requires and which a short context cannot fill. + ("partial_budget", None), +) + +#: The native reason strings the mirror above claims. A reason in this set +#: reaching :func:`attention` means the routing predicate did not mirror the +#: kernel, which is a wiring bug and says so. +REQUEST_SHAPE_REASONS = frozenset( + reason for _key, reason in REQUEST_SHAPE_DECLINES if reason is not None +) + + +def context_decline( + *, total_tokens: int, rows: int, k_eff: int, capacity: int +) -> Optional[str]: + """The request-shape verdict, from host ints alone. ``None`` = servable. + + Mirrors :data:`REQUEST_SHAPE_DECLINES` in the native contract's own order. + Called by the indexer BEFORE it commits the forward to this lane, because + that is the last point at which the stock chain is still reachable: once + the selection returns ``("sparse_blocks", top_idx)`` the rows-gather token + list was never built and there is nothing to fall back to. + + ``total_tokens`` must be the value the ATTENTION call site will pass, not + the logical context. On a fixed-capacity cache ``update_and_fetch`` + returns the whole backing, so attention's ``T`` is the bank capacity -- + which is exactly how a 1,024-token prompt (2,048-token bank, 512 complete + blocks, full budget) reached the kernel and was refused by it. + """ + + total_tokens = int(total_tokens) + if total_tokens <= 0: + return "empty_context" + if int(rows) > total_tokens: + return "rows_exceed_context" + if total_tokens > int(capacity): + return "context_exceeds_capacity" + if total_tokens > MAX_CONTEXT: + return "context_above_limit" + if total_tokens // COMPRESS_RATIO <= TOP_K: + return "short_context" + if int(k_eff) != TOP_K: + return "partial_budget" + return None + + +def note_request_decline( + site: str, reason: str, *, total_tokens: int, blocks: int +) -> None: + """This forward's own SHAPE is outside the lane. Routing, not failure. + + A server accepts whatever context length it is sent, and the kernel has no + analogue below a full budget -- not a smaller kernel, not a padded one. + Counted (never printed: this runs once per QSA layer per request) and + accepted by :func:`assert_traced`, so a short prompt runs the stock chain + instead of returning a 500. + + The numbers ride in :data:`_DECLINE_EXTREMES` as min/max pairs rather than + in the decline key, which would otherwise grow one key per distinct + context length. + """ + + _COUNTS["request_declines"] += 1 + note_route_decline(f"{site}: {reason}") + for name, value in (("blocks", int(blocks)), ("tokens", int(total_tokens))): + low = _DECLINE_EXTREMES.get(f"{name}_min") + high = _DECLINE_EXTREMES.get(f"{name}_max") + _DECLINE_EXTREMES[f"{name}_min"] = value if low is None else min(low, value) + _DECLINE_EXTREMES[f"{name}_max"] = value if high is None else max(high, value) + + +def route_snapshot() -> Dict[str, int]: + """The two counters to sample around a forward, for :func:`assert_traced`. + + A forward proves the armed lane engaged either by ROUTING to the kernel + (``route_hits``) or by declining for a REQUEST SHAPE the contract cannot + serve (``request_declines``). Anything else is an inert flag. + """ + + return { + "route_hits": int(_COUNTS["route_hits"]), + "request_declines": int(_COUNTS["request_declines"]), + } + + +def route_counters() -> Dict[str, Any]: + """Per-site hits and per-cause declines, for the receipt.""" + + return { + "route_hits": int(_COUNTS["route_hits"]), + "route_sites": dict(_ROUTE_SITES), + "route_declines": dict(_ROUTE_DECLINES), + "request_declines": int(_COUNTS["request_declines"]), + "request_decline_extremes": dict(_DECLINE_EXTREMES), + "short_context_tokens": SHORT_CONTEXT_TOKENS, + } + + +def engagement() -> Dict[str, Any]: + """Snapshot of the lane's engagement counters and install verdict. + + ``verify_kernel`` is the ENGAGEMENT LINE: if an ABBA reports a win and it + is zero, the win came from somewhere else. + """ + + report = dict(_COUNTS) + report["installed"] = _DISABLED_REASON == "" + report["disabled_reason"] = _DISABLED_REASON or None + report["probe"] = dict(_PROBE_REPORT) + report["route_sites"] = dict(_ROUTE_SITES) + report["route_declines"] = dict(_ROUTE_DECLINES) + return report + + +def receipt() -> Dict[str, Any]: + """The compact engagement block a benchmark receipt stores. + + Never raises: it reads ``_DISABLED_REASON`` directly rather than going + through a helper that treats "pending" as an error, because describing the + pending state IS what a receipt builder needs to do. + """ + + from mtplx.runtime_options import ( + qsa_sparse_decode_splits, + qsa_sparse_decode_tile, + ) + + key_tile, dim_tile = qsa_sparse_decode_tile() + block: Dict[str, Any] = { + "armed": armed(), + "installed": installed(), + "pending": pending(), + "disabled_reason": _DISABLED_REASON or None, + "tile": [int(key_tile), int(dim_tile)], + "splits": int(qsa_sparse_decode_splits()), + "verify_rows": VERIFY_ROWS, + "cache_installs": int(_COUNTS["cache_installs"]), + "probe_runs": int(_COUNTS["probe_runs"]), + "probe_failures": int(_COUNTS["probe_failures"]), + "kernel_calls": { + "verify_kernel": int(_COUNTS["verify_kernel"]), + }, + "probe": dict(_PROBE_REPORT), + } + block.update(route_counters()) + return block + + +def engagement_line(*, enabled: bool) -> str: + """The one-line install verdict, in the shape the other Fable lanes use.""" + + from mtplx.runtime_options import ( + qsa_sparse_decode_splits, + qsa_sparse_decode_tile, + ) + + if not enabled: + reason = _DISABLED_REASON or "install probe has not run" + return f"[mtplx] qsa_sparse_decode: off ({reason})" + key_tile, dim_tile = qsa_sparse_decode_tile() + worst = _PROBE_REPORT.get("worst") or {} + fp32 = worst.get("vs_fp32") or {} + shipped = worst.get("vs_shipped") or {} + widths = [str(VERIFY_ROWS)] if armed() else [] + return ( + "[mtplx] qsa_sparse_decode armed: " + f"rows={'+'.join(widths) or '-'} " + f"tile={int(key_tile)}:{int(dim_tile)} " + f"splits={int(qsa_sparse_decode_splits())} " + f"caches={int(_COUNTS['cache_installs'])} " + f"probe cell={worst.get('cell')!r} " + f"vs_fp32 ulps={fp32.get('max_abs_ulps', float('nan')):.3f} " + f"rel_l2={fp32.get('rel_l2', float('nan')):.3e} " + f"top1={fp32.get('top1', float('nan')):.4f} " + f"vs_shipped rel_l2={shipped.get('rel_l2', float('nan')):.3e} " + f"probe_runs={int(_COUNTS['probe_runs'])}" + ) + + +def assert_traced(rows: int, *, before: Dict[str, int], where: str) -> None: + """The armed verify lane must be IN this graph, not merely armed. + + ``before`` is :func:`route_snapshot` sampled before the traced forward. A + trace of the armed width that ends with no additional route hit is an + inert flag, and replaying that graph a few hundred times produces a delta + nobody can attribute -- which is what the 2026-09-02 window did. + + A forward that declined for its own REQUEST SHAPE satisfies this -- see + :func:`context_decline`. Those are the only declines the assertion + accepts, and they are why a 1 K prompt does not take the server down with + the flag armed. + """ + + if not armed() or int(rows) != VERIFY_ROWS: + return + now = route_snapshot() + if now["route_hits"] > int(before["route_hits"]): + return + if now["request_declines"] > int(before["request_declines"]): + return + raise SparseDecodeContractError( + "MTPLX_QSA_SPARSE_DECODE is armed but the split-K kernel is not " + f"in the traced {where} graph at {int(rows)} rows, and this forward's " + "shape is one the lane can serve, so it should have: the QSA " + "attention took another path, and this arm would replay the stock " + f"chain. route_sites={dict(_ROUTE_SITES)} " + f"declines={dict(_ROUTE_DECLINES)}" + ) + + +def disabled_reason() -> Optional[str]: + """The reason the lane is off, or ``None`` while it is usable/pending.""" + + return _DISABLED_REASON or None + + +def reset_for_tests() -> None: + """Clear the process verdict and counters. Tests only.""" + + global _DISABLED_REASON + _DISABLED_REASON = None + _PROBE_REPORT.clear() + _ROUTE_SITES.clear() + _ROUTE_DECLINES.clear() + _DECLINE_EXTREMES.clear() + for key in _COUNTS: + _COUNTS[key] = 0 + + +# --------------------------------------------------------------------------- +# The visible set -- one definition, shared by the kernel model, the +# reference and the tests. +# --------------------------------------------------------------------------- +def visible_block_count(q_abs: int) -> int: + """``visible_blocks`` exactly as ``_SRC_ROW_TOKENS`` computes it.""" + + return (int(q_abs) + 1) // COMPRESS_RATIO + + +def shipped_row_tokens( + top_idx_row, q_abs: int, *, topk: int = TOP_K +) -> Tuple[list, list]: + """Host model of the shipped rows-gather token list, for ONE row. + + Integers only. + + Returns ``(token_idx, token_ok)`` of width ``topk*ratio + ratio``, the + closed form the shipped Metal kernel writes: + + slot < topk*ratio : block = top_idx[slot // ratio] + token = block*ratio + slot % ratio + ok = block < visible_blocks + otherwise : token = visible_blocks*ratio + (slot - topk*ratio) + ok = token <= q_abs + + Note ``ok`` is evaluated PER SLOT against the block id, not against a + prefix length: ``top_idx`` is ``mx.argpartition``'s output and is not + sorted. + """ + + ratio = COMPRESS_RATIO + visible = visible_block_count(q_abs) + idx: list = [] + ok: list = [] + for slot in range(topk * ratio): + block = int(top_idx_row[slot // ratio]) + token = block * ratio + (slot % ratio) + good = block < visible + ok.append(good) + idx.append(token if good else 0) + for within in range(ratio): + token = visible * ratio + within + good = token <= int(q_abs) + ok.append(good) + idx.append(token if good else 0) + return idx, ok + + +def kernel_row_tokens( + top_idx_row, q_abs: int, *, key_length: int, topk: int = TOP_K +) -> list: + """Host model of the SPLIT KERNEL's per-slot selection for ONE row. + + Returns the kernel's ``selected[]`` array: the absolute key position for + each of the ``topk*ratio + ratio - 1`` slots, or ``-1`` for a masked one. + Written to be readable against the MSL in + ``steel_qsa_sparse_gqa_decode.h``; ``tests/test_qsa_sparse_decode.py`` + pins it against :func:`shipped_row_tokens`. + """ + + ratio = COMPRESS_RATIO + visible = visible_block_count(q_abs) + out: list = [] + for slot in range(topk * ratio): + block = int(top_idx_row[slot // ratio]) + pos = -1 + if 0 <= block < visible: + candidate = block * ratio + (slot % ratio) + if 0 <= candidate < int(key_length): + pos = candidate + out.append(pos) + for within in range(ratio - 1): + candidate = visible * ratio + within + pos = -1 + if 0 <= candidate < int(key_length) and candidate <= int(q_abs): + pos = candidate + out.append(pos) + return out + + +def visible_sets_agree( + top_idx_row, q_abs: int, *, key_length: int, topk: int = TOP_K +) -> bool: + """True when kernel and shipped lane attend the SAME multiset of keys.""" + + idx, ok = shipped_row_tokens(top_idx_row, q_abs, topk=topk) + shipped = sorted(t for t, good in zip(idx, ok) if good) + kernel = sorted(p for p in kernel_row_tokens( + top_idx_row, q_abs, key_length=key_length, topk=topk + ) if p >= 0) + return shipped == kernel + + +# --------------------------------------------------------------------------- +# The reference the probe compares against: the shipped rows-gather lane, +# transcribed. +# --------------------------------------------------------------------------- +def _row_tokens_mx(top_idx: mx.array, q_offset, *, topk: int) -> Tuple[mx.array, mx.array]: + """The rows-gather token list's closed form in plain MLX.""" + + ratio = COMPRESS_RATIO + rows = int(top_idx.shape[0]) + offsets = mx.arange(rows, dtype=mx.int32) + qpos = ( + q_offset.reshape(1).astype(mx.int32) + offsets + if isinstance(q_offset, mx.array) + else mx.array(int(q_offset), dtype=mx.int32) + offsets + ) + visible = (qpos + 1) // ratio # [rows] + within = mx.arange(ratio, dtype=mx.int32) + blocks = top_idx.astype(mx.int32) # [rows, topk] + block_tokens = (blocks[:, :, None] * ratio + within).reshape(rows, topk * ratio) + block_ok = mx.repeat(blocks < visible[:, None], ratio, axis=1) + tail_tokens = visible[:, None] * ratio + within + tail_ok = tail_tokens <= qpos[:, None] + token_idx = mx.concatenate([block_tokens, tail_tokens], axis=1) + token_ok = mx.concatenate([block_ok, tail_ok], axis=1) + token_idx = mx.where(token_ok, token_idx, mx.array(0, dtype=mx.int32)) + return token_idx, token_ok + + +def reference_attention( + queries: mx.array, + keys: mx.array, + values: mx.array, + top_idx: mx.array, + *, + query_offset, + scale: float, + topk: int = TOP_K, + fp32_scores: bool, + fp32_probs: bool, +) -> mx.array: + """The rows-gather attention over the SAME visible set, one rung at a time. + + Transcribed from ``mtplx/models/qwen4_exp.py::_qsa_rows_gather_attention`` + so the probe compares against ONE definition and does not import the model + into a kernel module. ``keys``/``values`` are the full + ``[1, 2, capacity, 256]`` backing; every gathered index is an absolute row + inside it, exactly as in the shipped lane. + + The two flags are the reference LADDER, and they exist to attribute the + kernel-vs-shipped delta rather than assume it: + + ``fp32_scores=False`` reproduces the shipped path exactly -- ``mx.matmul`` + on bf16 operands returns bf16, so the scores are rounded to bf16 + BEFORE the ``astype(float32) * scale`` and the softmax. Setting it + True upcasts q and k first (exact, bf16 -> fp32) so only the + accumulation and output precision change. + ``fp32_probs=False`` reproduces the shipped ``probs.astype(bfloat16)`` + before P@V. Setting it True keeps fp32 probabilities and upcasts V + (also exact). + + Every rung returns ``queries.dtype``, so all three are comparable element + for element against the kernel's own bf16 output. + """ + + token_idx, token_ok = _row_tokens_mx(top_idx, query_offset, topk=topk) + rows = int(queries.shape[2]) + width = int(token_idx.shape[1]) + k_sel = mx.take(keys, token_idx.reshape(-1), axis=2).reshape( + 1, KV_HEADS, rows, width, HEAD_DIM + ) + v_sel = mx.take(values, token_idx.reshape(-1), axis=2).reshape( + 1, KV_HEADS, rows, width, HEAD_DIM + ) + neg = mx.array(-mx.inf, dtype=mx.float32) + q_view = queries.reshape(1, KV_HEADS, GQA, rows, 1, HEAD_DIM) + k_view = k_sel.swapaxes(-1, -2).reshape(1, KV_HEADS, 1, rows, HEAD_DIM, width) + if fp32_scores: + # Upcasting bf16 -> fp32 is exact, so this changes the GEMM's output + # precision and nothing about the operands. + q_view = q_view.astype(mx.float32) + k_view = k_view.astype(mx.float32) + scores = mx.matmul(q_view, k_view).squeeze(-2).astype(mx.float32) * scale + scores = mx.where(token_ok[None, None, None], scores, neg) + probs = mx.softmax(scores, axis=-1) + v_view = v_sel.reshape(1, KV_HEADS, 1, rows, width, HEAD_DIM) + if fp32_probs: + v_view = v_view.astype(mx.float32) + else: + probs = probs.astype(queries.dtype) + out = mx.matmul(probs[..., None, :], v_view).squeeze(-2) + return out.reshape(1, Q_HEADS, rows, HEAD_DIM).astype(queries.dtype) + + +def stock_reference(queries, keys, values, top_idx, **kwargs) -> mx.array: + """The shipped path, exactly: bf16 scores AND bf16 probabilities.""" + + return reference_attention( + queries, keys, values, top_idx, fp32_scores=False, fp32_probs=False, **kwargs + ) + + +def shipped_fp32_probs_reference(queries, keys, values, top_idx, **kwargs) -> mx.array: + """Shipped, minus only the bf16 probability cast. Attribution rung.""" + + return reference_attention( + queries, keys, values, top_idx, fp32_scores=False, fp32_probs=True, **kwargs + ) + + +def fp32_reference(queries, keys, values, top_idx, **kwargs) -> mx.array: + """What the KERNEL computes: fp32 scores and fp32 probabilities. + + The deciding numerical reference. Against this the kernel's only + remaining differences are fp32 reassociation (Steel MMA fragments and the + split-K merge, both ~1e-6 relative) and the single bf16 store, so a delta + here of the same size as the kernel-vs-shipped delta would falsify the + attribution in this module's gate note. + """ + + return reference_attention( + queries, keys, values, top_idx, fp32_scores=True, fp32_probs=True, **kwargs + ) + + +# --------------------------------------------------------------------------- +# Contract + install +# --------------------------------------------------------------------------- +def check_cache_contract(keys: mx.array, values: mx.array, ratio: int) -> None: + """The cache half of the lane's contract. RAISES; never returns False.""" + + if int(ratio) != COMPRESS_RATIO: + raise RuntimeError( + "MTPLX_QSA_SPARSE_DECODE is wired for the ratio-4 QSA lane; " + f"got ratio={ratio}" + ) + if not mx.metal.is_available(): + raise RuntimeError( + "MTPLX_QSA_SPARSE_DECODE is a Metal kernel and has no " + "portable spelling" + ) + from mtplx.native import native_qsa_available + + if not native_qsa_available(): + raise RuntimeError( + "MTPLX_QSA_SPARSE_DECODE requires the built native " + "extension (native_extensions/qsa_sparse_gqa); build it with the " + "cmake command in mtplx/native/__init__.py's docstring" + ) + for name, arr in (("keys", keys), ("values", values)): + if arr is None: + raise RuntimeError( + f"MTPLX_QSA_SPARSE_DECODE requires a materialized {name} " + "bank" + ) + if arr.ndim != 4 or tuple(int(x) for x in arr.shape)[:2] != (1, KV_HEADS): + raise RuntimeError( + "MTPLX_QSA_SPARSE_DECODE requires a " + f"[1, {KV_HEADS}, capacity, {HEAD_DIM}] {name} bank; got " + f"{tuple(arr.shape)}" + ) + if int(arr.shape[3]) != HEAD_DIM: + raise RuntimeError( + "MTPLX_QSA_SPARSE_DECODE is wired for head_dim " + f"{HEAD_DIM}; got {int(arr.shape[3])}" + ) + if arr.dtype not in (mx.bfloat16, mx.float16): + raise RuntimeError( + "MTPLX_QSA_SPARSE_DECODE is wired for bf16/fp16 K/V; " + f"got {name} dtype {arr.dtype}" + ) + if keys.dtype != values.dtype: + raise RuntimeError( + "MTPLX_QSA_SPARSE_DECODE requires one K/V dtype; got " + f"{keys.dtype} and {values.dtype}" + ) + + +def _probe_cell(dtype: mx.Dtype, rows: int, total_tokens: int, q_offset: int, seed: int): + """One synthetic probe cell on the production geometry.""" + + mx.random.seed(seed) + nb_total = total_tokens // COMPRESS_RATIO + queries = mx.random.normal((1, Q_HEADS, rows, HEAD_DIM)).astype(dtype) + keys = mx.random.normal((1, KV_HEADS, PROBE_CAPACITY, HEAD_DIM)).astype(dtype) + values = mx.random.normal((1, KV_HEADS, PROBE_CAPACITY, HEAD_DIM)).astype(dtype) + # Deliberately UNSORTED distinct block ids, drawn from the whole logical + # range so cells with few complete blocks really do carry invisible ids -- + # which is the case a leading-prefix validity cut would get wrong. + ids = mx.argsort(mx.random.uniform(shape=(rows, nb_total)), axis=-1) + top_idx = ids[:, :TOP_K].astype(mx.int32) + return queries, keys, values, top_idx, total_tokens, q_offset + + +def _compare(reference: mx.array, candidate: mx.array) -> Dict[str, float]: + """Host-side parity statistics. One eval, at install, never in the hot path.""" + + ref = reference.astype(mx.float32) + got = candidate.astype(mx.float32) + diff = mx.abs(ref - got) + ref_absmax = mx.max(mx.abs(ref)) + l2_diff = mx.sqrt(mx.sum(diff * diff)) + l2_ref = mx.sqrt(mx.sum(ref * ref)) + top1 = mx.mean( + (mx.argmax(ref, axis=-1) == mx.argmax(got, axis=-1)).astype(mx.float32) + ) + stats = mx.stack( + [mx.max(diff), mx.mean(diff), ref_absmax, l2_diff, l2_ref, top1] + ) + mx.eval(stats) + max_abs, mean_abs, absmax, l2d, l2r, top1_f = (float(x) for x in stats.tolist()) + scale = max(absmax, 1e-3) + return { + "max_abs": max_abs, + "mean_abs": mean_abs, + "ref_absmax": absmax, + "max_abs_ulps": max_abs / (_BF16_REL_ULP * scale), + "rel_l2": l2d / max(l2r, 1e-12), + "top1": top1_f, + } + + +def probe_cells(rows: int) -> list: + """The install-time parity-probe cells for a verify window of ``rows``. + + Each cell is ``(name, rows, total_tokens, q_offset, seed)``: a long-context + cell (every selected block visible) and a just-past-crossover cell (some + selected ids are NOT visible). + + ``q_offset`` is DERIVED as ``total_tokens - rows``, never a constant, + because that is the real-serving invariant: the model passes + ``query_offset = cache.offset`` and ``total_tokens = keys.shape[2]`` (see + ``mtplx/models/qwen4_exp.py``: ``pos_start = cache.offset``, + ``T = k.shape[2]``), and the ``rows`` query tokens are the ones just + written to the cache, so they occupy positions ``[T-rows, T)`` and every + ``q_abs < total_tokens``. A constant offset sized for M4 (``total - 4``) + would push extra rows to positions ``>= total_tokens`` where the kernel + drops out-of-context tail tokens while a plain-MLX reference still attends + them -- a false parity failure. Deriving keeps M4 byte-identical + (``4093-4=4089``, ``2052-4=2048``). + """ + + rows = int(rows) + return [ + ("verify-4096", rows, 4093, 4093 - rows, 20260902), + ("verify-crossover", rows, 2052, 2052 - rows, 20260903), + ] + + +def install( + keys: mx.array, + values: mx.array, + *, + compress_ratio: int, + verify: bool, +) -> bool: + """Contract-check and parity-probe the lane once per process. + + Returns True when the lane is usable. RAISES on a contract failure; + DISABLES (returns False, records the reason) on a parity failure. Called + from ``TensorOffsetQSACache`` at cache install -- model build time, + outside any ``mx.compile`` trace. + """ + + global _DISABLED_REASON + if not verify: + return False + # Every cache that binds the lane counts, including the ones that reuse + # the process verdict: 12 QSA caches on the production pack, and a + # receipt showing fewer means some layers kept the stock chain. + _COUNTS["cache_installs"] += 1 + if _DISABLED_REASON is not None: + return _DISABLED_REASON == "" + + check_cache_contract(keys, values, compress_ratio) + + from mtplx.native import ( + qsa_sparse_gqa_decode, + qsa_sparse_gqa_decode_unsupported_reason, + ) + from mtplx.runtime_options import ( + qsa_sparse_decode_splits, + qsa_sparse_decode_tile, + ) + + key_tile, dim_tile = qsa_sparse_decode_tile() + key_splits = qsa_sparse_decode_splits() + scale = float(HEAD_DIM) ** -0.5 + dtype = keys.dtype + + cells = probe_cells(VERIFY_ROWS) if verify else [] + + worst: Dict[str, Any] = {} + for name, rows, total, offset, seed in cells: + q, k, v, top_idx, total_tokens, q_offset = _probe_cell( + dtype, rows, total, offset, seed + ) + reason = qsa_sparse_gqa_decode_unsupported_reason( + q, + k, + v, + top_idx, + query_offset=q_offset, + total_tokens=total_tokens, + scale=scale, + key_tile=key_tile, + dimension_tile=dim_tile, + key_splits=key_splits, + ) + if reason is not None: + # A contract miss on the lane's OWN synthetic production geometry + # is a configuration error, not a numerical verdict. + raise RuntimeError( + f"MTPLX_QSA_SPARSE_DECODE cannot serve its own probe " + f"cell {name!r}: {reason}" + ) + _COUNTS["probe_runs"] += 1 + candidate = qsa_sparse_gqa_decode( + q, + k, + v, + top_idx, + query_offset=q_offset, + total_tokens=total_tokens, + scale=scale, + key_tile=key_tile, + dimension_tile=dim_tile, + key_splits=key_splits, + ) + # Both rungs: the fp32 reference decides, the shipped one is a sanity + # bound on the shipped path's own bf16 quantisation. See the gate + # note at the top of this module for why they are not one number. + vs_fp32 = _compare( + fp32_reference(q, k, v, top_idx, query_offset=q_offset, scale=scale), + candidate, + ) + vs_shipped = _compare( + stock_reference(q, k, v, top_idx, query_offset=q_offset, scale=scale), + candidate, + ) + stats = {"cell": name, "vs_fp32": vs_fp32, "vs_shipped": vs_shipped} + if not worst or vs_fp32["max_abs_ulps"] > worst["vs_fp32"]["max_abs_ulps"]: + worst = stats + _PROBE_REPORT[name] = stats + + _PROBE_REPORT["worst"] = worst + _PROBE_REPORT["tile"] = [key_tile, dim_tile] + _PROBE_REPORT["key_splits"] = key_splits + + failures = [] + tight = worst.get("vs_fp32", {}) + loose = worst.get("vs_shipped", {}) + if tight.get("max_abs_ulps", math.inf) > PARITY_FP32_MAX_ABS_ULPS: + failures.append( + f"vs fp32 reference: max abs diff {tight['max_abs']:.3e} = " + f"{tight['max_abs_ulps']:.2f} bf16 ulp " + f"(limit {PARITY_FP32_MAX_ABS_ULPS})" + ) + if tight.get("rel_l2", math.inf) > PARITY_FP32_MAX_REL_L2: + failures.append( + f"vs fp32 reference: relative L2 {tight['rel_l2']:.3e} " + f"(limit {PARITY_FP32_MAX_REL_L2})" + ) + if tight.get("top1", 0.0) < PARITY_MIN_TOP1: + failures.append( + f"vs fp32 reference: head-dim top-1 {tight['top1']:.4f} " + f"(limit {PARITY_MIN_TOP1})" + ) + if loose.get("rel_l2", math.inf) > PARITY_SHIPPED_MAX_REL_L2: + failures.append( + f"vs shipped path: relative L2 {loose['rel_l2']:.3e} " + f"(limit {PARITY_SHIPPED_MAX_REL_L2}) -- this bounds the SHIPPED " + "path's own bf16 score and probability casts, so exceeding it " + "means the kernel is wrong by a factor, not merely re-rounded" + ) + if loose.get("top1", 0.0) < PARITY_MIN_TOP1: + failures.append( + f"vs shipped path: head-dim top-1 {loose['top1']:.4f} " + f"(limit {PARITY_MIN_TOP1})" + ) + + if failures: + _COUNTS["probe_failures"] += 1 + _DISABLED_REASON = ( + f"parity probe failed on cell {worst.get('cell')!r}: " + + "; ".join(failures) + ) + _emit(engagement_line(enabled=False)) + _emit( + "[mtplx] qsa_sparse_decode install: " + + json.dumps(receipt(), sort_keys=True) + ) + return False + + _DISABLED_REASON = "" + _emit(engagement_line(enabled=True)) + _emit( + "[mtplx] qsa_sparse_decode install: " + + json.dumps(receipt(), sort_keys=True) + ) + return True + + +# --------------------------------------------------------------------------- +# Hot path +# --------------------------------------------------------------------------- +def attention( + queries: mx.array, + keys: mx.array, + values: mx.array, + top_idx: mx.array, + *, + query_offset, + total_tokens: int, + scale: float, +) -> mx.array: + """Run the split-K kernel for one QSA layer. Raises on a contract miss. + + ``queries`` is the ``[1, 24, M, 256]`` transposed view the attention + module already holds; ``keys``/``values`` are the FULL cache backing. + ``top_idx`` is ``mx.argpartition``'s ``[M, 512]`` output in its own + order. + """ + + from mtplx.native import ( + qsa_sparse_gqa_decode, + qsa_sparse_gqa_decode_unsupported_reason, + ) + from mtplx.runtime_options import ( + qsa_sparse_decode_splits, + qsa_sparse_decode_tile, + ) + + key_tile, dim_tile = qsa_sparse_decode_tile() + key_splits = qsa_sparse_decode_splits() + reason = qsa_sparse_gqa_decode_unsupported_reason( + queries, + keys, + values, + top_idx, + query_offset=query_offset, + total_tokens=total_tokens, + scale=scale, + key_tile=key_tile, + dimension_tile=dim_tile, + key_splits=key_splits, + ) + if reason is not None: + if reason in REQUEST_SHAPE_REASONS: + # Unreachable when the mirror is complete, and that is the point of + # saying so: the INDEXER owns every request-shape decision, because + # by the time a forward reaches here the rows-gather token list was + # never built and there is nothing to fall back to. A 1,024-token + # prompt died exactly here on 2026-09-02. + raise SparseDecodeContractError( + "MTPLX_QSA_SPARSE_DECODE routing did not mirror the " + f"kernel contract: the kernel refused this call because {reason!r}, " + "which is a REQUEST shape that " + "kernels/qsa_sparse_decode.context_decline must have declined " + "before the selection committed to this lane. Fix the mirror " + "in REQUEST_SHAPE_DECLINES; do not make the request fail" + ) + raise RuntimeError( + "MTPLX_QSA_SPARSE_DECODE is armed but this call is off " + f"contract: {reason}" + ) + _COUNTS["verify_kernel"] += 1 + return qsa_sparse_gqa_decode( + queries, + keys, + values, + top_idx, + query_offset=query_offset, + total_tokens=total_tokens, + scale=scale, + key_tile=key_tile, + dimension_tile=dim_tile, + key_splits=key_splits, + ) + + +__all__ = [ + "COMPRESS_RATIO", + "SparseDecodeContractError", + "armed", + "assert_traced", + "engagement_line", + "installed", + "MAX_CONTEXT", + "REQUEST_SHAPE_DECLINES", + "REQUEST_SHAPE_REASONS", + "SHORT_CONTEXT_TOKENS", + "context_decline", + "note_request_decline", + "note_route_decline", + "note_route_hit", + "route_snapshot", + "pending", + "receipt", + "route_counters", + "GQA", + "HEAD_DIM", + "KV_HEADS", + "PARITY_FP32_MAX_ABS_ULPS", + "PARITY_FP32_MAX_REL_L2", + "PARITY_MIN_TOP1", + "PARITY_SHIPPED_MAX_REL_L2", + "PROBE_CAPACITY", + "Q_HEADS", + "SELECTED_TOKENS", + "TOP_K", + "VERIFY_ROWS", + "attention", + "check_cache_contract", + "disabled_reason", + "fp32_reference", + "engagement", + "install", + "kernel_row_tokens", + "reset_for_tests", + "reference_attention", + "shipped_fp32_probs_reference", + "shipped_row_tokens", + "stock_reference", + "visible_block_count", + "visible_sets_agree", +] diff --git a/mtplx/models/qwen4_exp.py b/mtplx/models/qwen4_exp.py index 9623a8ad2..14b2a95c7 100644 --- a/mtplx/models/qwen4_exp.py +++ b/mtplx/models/qwen4_exp.py @@ -63,10 +63,14 @@ from mtplx.attention_context import current_attention_phase from mtplx.runtime_options import ( + qsa_sparse_decode_enabled, qwen4_hc_m4_enabled, qwen4_opdiet_enabled, qwen4_verify_glue_enabled, ) +# Verify-width shared between the model and the split-K decode lane so a +# selection width and the kernel's own contract cannot drift apart. +from mtplx.kernels.qsa_sparse_decode import VERIFY_ROWS as _SPARSE_VERIFY_ROWS @dataclass @@ -2994,6 +2998,15 @@ def _select_eager( :, nb_total - k_eff : ] + # This selector is what a fixed-capacity verify forward reaches on + # the fixed-M4 stack; asking the sparse lane the wrong width here + # would make it unreachable (armed flag, installed cache, kernel + # never runs). Routing, never failure at the request shape. + if self._sparse_decode_route( + cache, rows=S, k_eff=k_eff, site="select_eager_verify" + ): + return ("sparse_blocks", top_idx) + if S > 1 and not fixed_capacity and _qsa_large_prefill_enabled(S, total): # Preserve the eager score/top-k expression as an independently # selectable oracle while handing attention the compact block set. @@ -3246,6 +3259,114 @@ def _select_fused( for leaf in range(len(chunks[0])) ) + def _sparse_decode_route( + self, + cache: QSACache, + *, + rows: int, + k_eff: int, + site: str, + ) -> bool: + """True when the native split-K sparse-GQA kernel serves this call. + + Host-only, and read from state the CACHE validated once at install + (graphbank ran the contract check and the numerical parity probe + there, at model build, outside any mx.compile trace). The indexer + never re-derives the decision, so two traces of the same verify graph + cannot disagree about which attention they contain. + + TWO KINDS OF "NO", AND THE 2026-09-02 WINDOW IS WHY THEY ARE SPLIT. + + That window armed MTPLX_QSA_SPARSE_DECODE at 16 K and measured + the control on both seeds. The cause was one silent narrowing: the + only call site that could reach the verify width asked this predicate + for a width the flag did not arm, so it read a zero row count, + returned False, and fell through to the rows-gather lane. Nothing + said so. So: + + * ROUTING (returns False): the flag is off for this width, or the + selection is a width it does not arm -- a 16 K prefill row count is + not a 4-row verify. A growable cache is routing too, and IS + recorded in the lane's ``route_declines``: the lane installs on the + fixed-capacity compiled-verify cache, so an armed run that only ever + saw growable caches has a readable cause in its receipt rather than + a silent zero. + * FAILURE (raises): the flag is armed, this IS the width it arms, and + the cache is the one it installs on -- but the geometry or the + budget does not match. An armed flag that reverts here would make + the arm measure the stock chain again. + + ``site`` names the call site in the receipt. + """ + + # The unarmed path is ONE cached-bool test: this predicate runs once + # per QSA layer per forward, and a decode cycle that pays for a flag + # nobody armed is a cost with no lever. + if not qsa_sparse_decode_enabled(): + return False + if int(rows) != int(_SPARSE_VERIFY_ROWS): + return False + + from mtplx.kernels import qsa_sparse_decode as _qsa_sparse + + if not bool(getattr(cache, "fixed_capacity", False)): + # The lane's install probe needs the materialized fixed bank, so a + # growable cache never carried it. Construction owns this gate: + # TensorOffsetQSACache.__init__ raises when the flag is armed and + # the cache was built without the lane. + _qsa_sparse.note_route_decline(f"{site}: growable cache") + return False + attribute = "qsa_sparse_decode_rows" + wired = int(getattr(cache, attribute, 0)) + lane = "MTPLX_QSA_SPARSE_DECODE" + if wired <= 0: + raise RuntimeError( + f"{lane} is armed and this is its {int(rows)}-row width, but " + f"the fixed QSA cache carries {attribute}={wired}: it was " + "built without the lane, or the install probe disabled it" + ) + if int(rows) != wired: + raise RuntimeError( + f"{lane} bound {wired} rows but this selection is " + f"{int(rows)}; the cache and the module disagree about the " + "width the kernel serves" + ) + if int(self.ratio) != 4 or int(self.block_topk) != 512: + raise RuntimeError( + f"{lane} is wired for the ratio-4 top-512 QSA geometry the " + f"metallib is instantiated for; got ratio={int(self.ratio)}, " + f"block_topk={int(self.block_topk)}" + ) + # REQUEST SHAPE -- routing, never failure, and this is the LAST point + # at which the stock chain is still reachable: once the selection + # returns ("sparse_blocks", top_idx) the rows-gather token list was + # never built and attention has nothing to fall back to. + # + # ``total_tokens`` is read from the K/V BACKING, because that is what + # the attention call site passes: ``update_and_fetch`` returns the + # whole fixed bank, so its ``T`` is the capacity, not the logical + # context. A 1,024-token prompt in a 2,048-token bank has a FULL + # 512-block budget (k_eff == 512) and a context that has not crossed + # the kernel's dense/sparse boundary -- two different questions, and + # asking only the budget one took the 1 K served cell down on + # 2026-09-02 with an HTTP 500 from inside the kernel wrapper. + backing = getattr(cache.kv, "keys", None) + total_tokens = 0 if backing is None else int(backing.shape[2]) + decline = _qsa_sparse.context_decline( + total_tokens=total_tokens, + rows=int(rows), + k_eff=int(k_eff), + capacity=total_tokens, + ) + if decline is not None: + # Counted, never printed: once per QSA layer per request. + _qsa_sparse.note_request_decline( + site, decline, total_tokens=total_tokens, blocks=int(k_eff) + ) + return False + _qsa_sparse.note_route_hit(site) + return True + def _verify_glue_rope_idx(self) -> bool: """True when ``MTPLX_QWEN4_VERIFY_GLUE``'s ``qsa_rope_idx`` serves. @@ -4034,6 +4155,14 @@ def _qsa_blocks_to_dense_mask( return ((token_selected | tail) & causal)[None, None] +def _sparse_route_snapshot(): + """``qsa_sparse_decode.route_snapshot()``, imported only when armed.""" + + from mtplx.kernels import qsa_sparse_decode as _qsa_sparse + + return _qsa_sparse.route_snapshot() + + class Attention(nn.Module): """Gated GQA (qwen3_5 style: double-width q_proj, sigmoid output gate, per-head q/k RMSNorm, partial rotary) masked by the QSA indexer.""" @@ -4078,6 +4207,64 @@ def __init__(self, args: TextArgs): else None ) + def _sparse_decode_required(self, cache: QSACache, rows: int) -> bool: + """True when the armed split-K lane MUST have served this selection. + + Narrow on purpose, and every narrowing is a place the lane genuinely + cannot be: no indexer at all (a dense layer), a growable cache (the + lane installs on the fixed-capacity compiled-verify cache, and + construction owns that gate), or a width this process did not arm. + Everything else is the contract, and failing it is fatal -- see the + call site in ``__call__``. + """ + + if self.indexer is None: + return False + rows = int(rows) + if rows != _SPARSE_VERIFY_ROWS: + return False + if not qsa_sparse_decode_enabled(): + return False + return bool(getattr(cache, "fixed_capacity", False)) + + def _require_sparse_decode_lane(self, sel_mask, *, rows: int, before) -> None: + """The armed lane is IN this graph, or this forward legitimately isn't. + + The trace-time proof. Every ``sel_mask`` branch below is a + DIFFERENT attention, and the 2026-09-02 window took one of them + (rows-gather) for 394 cycles with the flag armed and nothing saying + so. The selection is decided in this Python body, which under + ``mx.compile`` runs at TRACE time, so this raises while the graph is + being built rather than after a window. + + TWO ways to pass. The selection is ``sparse_blocks``; or the indexer + declined for this request's own SHAPE -- context length, row count, + block budget -- where the kernel has no analogue and the stock chain + is the correct lane (see ``qsa_sparse_decode.context_decline``). A + servable forward that got any other lane still raises -- that is the + armed-but-inert failure this guard exists for. + """ + + if isinstance(sel_mask, tuple) and sel_mask and sel_mask[0] == "sparse_blocks": + return + from mtplx.kernels import qsa_sparse_decode as _qsa_sparse + + now = _qsa_sparse.route_snapshot() + if now["request_declines"] > int(before["request_declines"]): + return + lane = "MTPLX_QSA_SPARSE_DECODE" + took = ( + sel_mask[0] + if isinstance(sel_mask, tuple) and sel_mask + else ("dense_mask" if sel_mask is not None else "no_selection") + ) + raise RuntimeError( + f"{lane} is armed and this is its {int(rows)}-row width on a " + "fixed QSA cache whose shape the lane can serve, but the indexer " + f"handed attention the {took!r} lane: the split-K kernel is not " + "in this graph and the arm would replay the stock chain" + ) + def _verify_glue_rope(self, rows: int) -> bool: """True when ``MTPLX_QWEN4_VERIFY_GLUE``'s ``qsa_rope`` serves this call. @@ -4099,6 +4286,13 @@ def __call__(self, x: mx.array, cache: QSACache) -> mx.array: B, S, _ = x.shape pos_start = cache.offset vrope = vision_rope_state() + # The armed split-K lane's per-layer proof, sampled BEFORE the + # indexer runs (the routing decision is taken inside it). A forward + # proves engagement two ways: it routed to the kernel, or it + # declined for its own request SHAPE, where the stock chain is + # correct (see _require_sparse_decode_lane). + sparse_required = self._sparse_decode_required(cache, S) + sparse_before = _sparse_route_snapshot() if sparse_required else None fused = getattr(self, "qkv_fused", None) if fused is not None: @@ -4198,6 +4392,11 @@ def __call__(self, x: mx.array, cache: QSACache) -> mx.array: # images too; dropping selection changed its attention function. sel_mask = None + if vrope is None and sparse_required: + self._require_sparse_decode_lane( + sel_mask, rows=int(S), before=sparse_before + ) + if isinstance(sel_mask, tuple) and sel_mask and sel_mask[0] == "flash": # Block-sparse flash attention over the indexer's exact visible # set. Reads the cache BACKING arrays in place at their @@ -4334,6 +4533,34 @@ def _qsa_gather_call(): compress_ratio=self.indexer.ratio, ) + if isinstance(sel_mask, tuple) and sel_mask and sel_mask[0] == "sparse_blocks": + # MTPLX_QSA_SPARSE_DECODE: split-K direct-index + # sparse GQA over exactly the visible set the rows-gather lane + # attends, reading the cache BACKING in place. No gathered K/V + # tensor is written, no transposed copy is made for the score + # operand, and no score tensor is materialized -- which is the + # whole point: the shipped lane's ~70 MB per layer is bytes, not + # bandwidth (see mtplx/kernels/qsa_sparse_decode.py). + # + # ROUNDING CLASS, not exact: fp32 online softmax in exp2, fp32 + # probabilities into an fp32 P@V, Steel-MMA reassociation, and a + # split-K rescale. Adopted on greedy-token agreement plus a full + # HumanEval run, exactly like MTPLX_QWEN4_HC_M4. + from mtplx.kernels import qsa_sparse_decode as _qsa_sparse + + _, sparse_top_idx = sel_mask + out = _qsa_sparse.attention( + q, + cache.kv.keys, + cache.kv.values, + sparse_top_idx, + query_offset=pos_start, + total_tokens=T, + scale=self.scale, + ) + out = out.transpose(0, 2, 1, 3).reshape(B, S, -1) + return self.o_proj(out * mx.sigmoid(gate)) + if isinstance(sel_mask, tuple) and sel_mask and sel_mask[0] == "gather_rows": # Rows-gather lane (S>1): each verify/pipeline row reads only # its own selected blocks + tail instead of the full context diff --git a/mtplx/native/__init__.py b/mtplx/native/__init__.py new file mode 100644 index 000000000..697419902 --- /dev/null +++ b/mtplx/native/__init__.py @@ -0,0 +1,716 @@ +"""Python surface for MTPLX's native (CMake + nanobind) MLX primitives. + +Today this is one kernel: :func:`qsa_sparse_gqa`, the direct-index sparse-GQA +attention ported from oMLX (see +``native_extensions/qsa_sparse_gqa/sparse_gqa/steel_qsa_sparse_gqa.h`` for +provenance). It is Steel MMA, so it cannot live in ``mx.fast.metal_kernel`` +(the Laguna full-port verdict: no ``mlx::steel`` MMA reachable from +``metal_kernel``) and has to be a real MLX primitive in a built extension. + +Phase 1 is standalone: nothing in ``mtplx/models/qwen4_exp.py`` calls this yet. +The gate order below deliberately mirrors +``mtplx.kernels.qsa_prefill_flash._unsupported_reason`` so the two attention +consumers refuse the same shapes for the same stated reason. + +Build (CPU-only; no Metal execution):: + + cd native_extensions/qsa_sparse_gqa + cmake -S . -B build \\ + -DCMAKE_LIBRARY_OUTPUT_DIRECTORY=$PWD/mtplx_native_qsa/ \\ + -DCMAKE_BUILD_TYPE=Release -DBUILD_SHARED_LIBS=ON \\ + -DPython_EXECUTABLE=/bin/python + cmake --build build -j 8 + +``python setup.py build_ext --inplace`` is the same build through setuptools; +it needs ``setuptools`` in the venv, which the current qwen38 venv does not +have (which is also why ``verify_mlp`` has no built artifact on this box). +""" + +from __future__ import annotations + +import math +import operator +import re +import sys +from functools import lru_cache +from pathlib import Path +from typing import Any + +import mlx.core as mx + +__all__ = [ + "native_qsa_available", + "qsa_sparse_gqa", + "qsa_sparse_gqa_decode", + "qsa_sparse_gqa_decode_split_geometry", + "qsa_sparse_gqa_decode_supported", + "qsa_sparse_gqa_decode_unsupported_reason", + "qsa_sparse_gqa_supported", + "qsa_sparse_gqa_unsupported_reason", +] + +# Production Qwen3.8 Flash-Next QSA geometry. These are the ONLY shapes the +# kernel is instantiated for; everything else fails closed rather than +# silently changing the attention algorithm. +_BATCH = 1 +_Q_HEADS = 24 +_KV_HEADS = 2 +_GQA = 12 +_HEAD_DIM = 256 +_COMPRESS_RATIO = 4 +_TOP_K_BLOCKS = 512 +_MAX_CONTEXT = 1_048_576 +_SUPPORTED_DTYPES = (mx.float16, mx.bfloat16) +_SUPPORTED_ID_DTYPES = (mx.int32, mx.uint32) +#: (key_tile, dimension_tile) pairs the metallib instantiates. +_SUPPORTED_TILES = ((128, 32), (256, 32), (64, 64), (128, 64)) +_DEFAULT_TILE = (128, 32) + + +def _extension_path() -> Path: + return ( + Path(__file__).resolve().parents[2] + / "native_extensions" + / "qsa_sparse_gqa" + ) + + +#: nanobind writes its ABI tag as one literal, e.g. ``v21_system_libcpp_abi1``. +_NB_ABI_TAG_RE = re.compile( + rb"v(\d+)(?:[0-9a-zA-Z.\-]*)_[0-9a-zA-Z_]*(?:libcpp|libstdcpp|ms)[0-9a-zA-Z_]*" +) + + +def _nanobind_internals_version(binary: Path) -> int | None: + """The nanobind internals version a shared object was built against.""" + + try: + data = binary.read_bytes() + except OSError: + return None + versions = {int(m.group(1)) for m in _NB_ABI_TAG_RE.finditer(data)} + return versions.pop() if len(versions) == 1 else None + + +@lru_cache(maxsize=1) +def _nanobind_abi_mismatch() -> str | None: + """Precise reason when the extension cannot see mlx.core's type registry. + + Two nanobind modules share a type registry only when they agree on the + capsule key ``__nb_internals____``. The domain is + ``NB_DOMAIN=mlx`` on both sides; the tag carries ``NB_INTERNALS_VERSION``, + which moves between nanobind releases. A mismatch does not stop the build + or the import -- it makes every call that takes an ``mx::array`` raise a + bare ``TypeError`` whose signature line prints ``mlx::core::array`` instead + of ``array``. Diagnosing that from the TypeError alone cost a guarded + GPU window, so it is named here instead. + + ``None`` when the tags match or cannot be read; a reason string otherwise. + """ + + ours = sorted(_extension_path().glob("mtplx_native_qsa/_ext*.so")) + if not ours: + return None + core = sorted(Path(mx.__file__).parent.glob("core*.so")) if mx.__file__ else [] + if not core: + return None + ext_version = _nanobind_internals_version(ours[0]) + mlx_version = _nanobind_internals_version(core[0]) + if ext_version is None or mlx_version is None or ext_version == mlx_version: + return None + return ( + f"the built extension uses nanobind internals v{ext_version} but " + f"mlx.core uses v{mlx_version}, so it cannot resolve mlx::core::array " + "and every call raises TypeError; rebuild with " + "-DMTPLX_NANOBIND_DIR= (diagnose with " + "the native-ABI checker)" + ) + + +@lru_cache(maxsize=1) +def _load_extension() -> Any: + """Import the built extension, or return the import error. + + A successful import is not enough: the module can import cleanly and still + be unable to cast a single array (see :func:`_nanobind_abi_mismatch`), so + that check is folded in here rather than left to fail at the first call. + """ + + native_path = str(_extension_path()) + if native_path not in sys.path: + sys.path.insert(0, native_path) + try: + import mtplx_native_qsa # noqa: PLC0415 + except Exception as exc: # pragma: no cover - depends on build state + return exc + mismatch = _nanobind_abi_mismatch() + if mismatch is not None: # pragma: no cover - depends on build state + return RuntimeError(mismatch) + return mtplx_native_qsa + + +def native_qsa_available() -> bool: + """True when the built extension imports.""" + + return not isinstance(_load_extension(), Exception) + + +def _on_metal_device() -> bool: + """Metal availability is insufficient when MLX currently targets CPU.""" + + try: + return mx.metal.is_available() and mx.default_device() == mx.gpu + except (AttributeError, RuntimeError, TypeError, ValueError): + return False + + +def _normalized_block_ids(block_ids: mx.array, rows: int) -> mx.array | None: + """Accept the selector's ``[S, 512]`` or the kernel ABI's ``[1,1,S,512]``. + + ``_select_eager`` emits ``[S, 512]``; the kernel wants ``[1, 1, S, 512]``. + The reshape is a view on the contiguous selector output, not a copy, and + the int32 dtype is accepted natively (the metallib instantiates both + int32 and uint32) so the lane never pays an 8 MB astype per layer. + """ + + if block_ids.ndim == 2: + if tuple(int(x) for x in block_ids.shape) != (rows, _TOP_K_BLOCKS): + return None + return block_ids.reshape(1, 1, rows, _TOP_K_BLOCKS) + if block_ids.ndim == 4: + if tuple(int(x) for x in block_ids.shape) != ( + _BATCH, + 1, + rows, + _TOP_K_BLOCKS, + ): + return None + return block_ids + return None + + +def qsa_sparse_gqa_unsupported_reason( + queries: mx.array, + keys: mx.array, + values: mx.array, + block_ids: mx.array, + *, + pos_start: int, + total_tokens: int, + scale: float, + key_tile: int = _DEFAULT_TILE[0], + dimension_tile: int = _DEFAULT_TILE[1], +) -> str | None: + """``None`` when the call is on contract, else the precise reason.""" + + extension = _load_extension() + if isinstance(extension, Exception): + return f"the native QSA extension is not built ({extension})" + if not _on_metal_device(): + return "the active MLX device is not an available Metal GPU" + + arrays = (queries, keys, values, block_ids) + if any(not isinstance(array, mx.array) for array in arrays): + return "all tensor inputs must be MLX arrays" + if queries.ndim != 4 or keys.ndim != 4 or values.ndim != 4: + return "Q, K, and V must be rank four" + if block_ids.ndim not in (2, 4): + return "block ids must be rank two [S, 512] or rank four [1, 1, S, 512]" + + batch, query_heads, rows, head_dim = (int(x) for x in queries.shape) + if (batch, query_heads, head_dim) != (_BATCH, _Q_HEADS, _HEAD_DIM): + return "Q must have production shape [1, 24, S, 256]" + if rows <= 0: + return "Q must carry at least one query row" + + key_batch, kv_heads, capacity, key_dim = (int(x) for x in keys.shape) + if (key_batch, kv_heads, key_dim) != (_BATCH, _KV_HEADS, _HEAD_DIM): + return "K must have production shape [1, 2, capacity, 256]" + if tuple(int(x) for x in values.shape) != tuple(int(x) for x in keys.shape): + return "V must have the same full-backing shape as K" + + if queries.dtype not in _SUPPORTED_DTYPES: + return "Q must be float16 or bfloat16" + if keys.dtype != queries.dtype or values.dtype != queries.dtype: + return "Q, K, and V dtypes must match" + if block_ids.dtype not in _SUPPORTED_ID_DTYPES: + return "block ids must be int32 or uint32" + if _normalized_block_ids(block_ids, rows) is None: + return "block ids must have shape [S, 512] or [1, 1, S, 512]" + + # Host scalars only: a traced scalar would make these comparisons + # synchronize the graph. Same contract as qsa_prefill_flash. + if isinstance(pos_start, mx.array) or isinstance(total_tokens, mx.array): + return "pos_start and total_tokens must be host integers" + if isinstance(scale, mx.array): + return "scale must be a host float" + if isinstance(pos_start, bool) or isinstance(total_tokens, bool): + return "pos_start and total_tokens cannot be bool" + try: + pos_start_i = operator.index(pos_start) + total_tokens_i = operator.index(total_tokens) + except TypeError: + return "pos_start and total_tokens must be exact host integers" + if isinstance(scale, bool) or not isinstance(scale, (int, float)): + return "scale must be a numeric host scalar" + scale_f = float(scale) + + if pos_start_i < 0 or total_tokens_i <= 0: + return "positions must describe a non-empty non-negative suffix" + if pos_start_i + rows > total_tokens_i: + return "Q must be a causal suffix inside total_tokens" + if total_tokens_i > capacity: + return "the logical token count exceeds the full K/V backing capacity" + if total_tokens_i > _MAX_CONTEXT: + return "the logical token count exceeds the production context limit" + if total_tokens_i // _COMPRESS_RATIO <= _TOP_K_BLOCKS: + return "the context has not crossed the dense/sparse boundary" + if not math.isfinite(scale_f): + return "scale must be finite" + + if (int(key_tile), int(dimension_tile)) not in _SUPPORTED_TILES: + return ( + "(key_tile, dimension_tile) must be one of " + + ", ".join(str(t) for t in _SUPPORTED_TILES) + ) + return None + + +def qsa_sparse_gqa_supported( + queries: mx.array, + keys: mx.array, + values: mx.array, + block_ids: mx.array, + *, + pos_start: int, + total_tokens: int, + scale: float, + key_tile: int = _DEFAULT_TILE[0], + dimension_tile: int = _DEFAULT_TILE[1], +) -> bool: + """Whether the exact production-only kernel contract is met.""" + + return ( + qsa_sparse_gqa_unsupported_reason( + queries, + keys, + values, + block_ids, + pos_start=pos_start, + total_tokens=total_tokens, + scale=scale, + key_tile=key_tile, + dimension_tile=dimension_tile, + ) + is None + ) + + +def qsa_sparse_gqa( + queries: mx.array, + keys: mx.array, + values: mx.array, + block_ids: mx.array, + *, + pos_start: int, + total_tokens: int, + scale: float, + key_tile: int = _DEFAULT_TILE[0], + dimension_tile: int = _DEFAULT_TILE[1], + stream: Any = None, +) -> mx.array: + """Direct-index sparse GQA attention over chronological QSA block ids. + + ``queries`` ``[1, 24, S, 256]`` fp16/bf16; the ``[B,H,S,D]`` transposed + view the Attention module already builds. + ``keys``/``values`` ``[1, 2, capacity, 256]`` -- the FULL KV cache + backing, read in place at its allocation stride. Never slice + it to ``total_tokens`` first: that copy is the whole context. + ``block_ids`` ``[S, 512]`` int32 (``_select_eager``'s ``flash_prefill`` + output) or ``[1, 1, S, 512]``. Chronological, and the valid + entries must occupy the leading + ``min(512, (pos + 1) // 4)`` slots of each row -- which is + what the selector produces, because it sorts the raw top-k + ascending and validity there is the threshold predicate + ``id < complete_blocks``. The kernel derives validity from + that invariant instead of reading ``block_valid``; the + standalone harness asserts it. + ``total_tokens`` logical tokens in the cache (NOT ``capacity``). + + Returns ``[1, 24, S, 256]``, same dtype as ``queries``. + + Numerics: fp32 online softmax (exp2) and fp32 P@V over the same visible + set as the dense lane -- a rounding-class difference, not an exactness + one. See the sparse-GQA microbenchmark for the tolerance + statement and the measured deltas. + """ + + reason = qsa_sparse_gqa_unsupported_reason( + queries, + keys, + values, + block_ids, + pos_start=pos_start, + total_tokens=total_tokens, + scale=scale, + key_tile=key_tile, + dimension_tile=dimension_tile, + ) + if reason is not None: + raise ValueError(f"[mtplx.native.qsa_sparse_gqa] {reason}.") + + extension = _load_extension() + selected = _normalized_block_ids(block_ids, int(queries.shape[2])) + args = ( + queries, + keys, + values, + selected, + float(scale), + int(pos_start), + int(total_tokens), + int(key_tile), + int(dimension_tile), + ) + # Omit the kwarg entirely when no stream was asked for, rather than + # passing an explicit None: the binding's default is applied by nanobind + # without going through the StreamOrDevice caster at all, which is one + # fewer thing to be wrong about in an extension module. + if stream is None: + return extension.qsa_sparse_gqa_attention(*args) + return extension.qsa_sparse_gqa_attention(*args, stream=stream) + + +# --------------------------------------------------------------------------- +# Split-K (KV-split) DECODE variant -- M=4 fixed verify and M=1 draft. +# +# Separate entry points, not a rows argument on the prefill one, because the +# two differ in more than the row count: +# +# * the grid's z axis is the KV SPLIT here, and the kernel is two dispatches +# (split + merge) rather than one; +# * validity is decided PER SLOT rather than by a leading-prefix cut, +# because the decode selector hands ``mx.argpartition``'s raw, UNSORTED +# output straight through, while the prefill selector sorts; +# * the query offset is a device buffer, so a tensor-valued cache offset +# never has to be read on the host. +# --------------------------------------------------------------------------- + +#: The kernel's own selected-token width: 512 blocks x 4 tokens plus the at +#: most three causal tail tokens of the incomplete block. The shipped lane +#: builds 2,052 slots; its 2,052nd is invalid for every query position (see +#: the note in ``qsa_sparse_gqa_decode``), so the two visible sets agree. +_SELECTED_TOKENS = _TOP_K_BLOCKS * _COMPRESS_RATIO + (_COMPRESS_RATIO - 1) +#: Partial rows are [O(head_dim) | m | l] in fp32. +_PARTIAL_LD = _HEAD_DIM + 2 +_MAX_KEY_SPLITS = 64 +#: Matches mtplx.runtime_options.QSA_SPARSE_DECODE_DEFAULT_SPLITS; a +#: test pins the two together so the bench and the lane cannot drift. +_DEFAULT_KEY_SPLITS = 17 +_INTEGER_DTYPES = ( + mx.int8, + mx.int16, + mx.int32, + mx.int64, + mx.uint8, + mx.uint16, + mx.uint32, + mx.uint64, +) + + +def qsa_sparse_gqa_decode_split_geometry( + selected_tokens: int = _SELECTED_TOKENS, + key_tile: int = _DEFAULT_TILE[0], + key_splits: int = _DEFAULT_KEY_SPLITS, +) -> tuple[int, int, int]: + """``(n_tiles, tiles_per_split, n_splits)`` for the split-K decode grid. + + A pure-host mirror of the C++ ``qsa_sparse_gqa_decode_split_geometry`` so + the harness, the tests and the partial-buffer sizing never restate the + arithmetic. ``tests/test_qsa_sparse_decode.py`` pins the two + against each other whenever the extension is built. + + The rounding is deliberate: ``n_splits`` is recomputed from + ``tiles_per_split`` so the LAST split always has work. With 17 tiles and + 8 requested splits, ``tiles_per_split`` is 3 and six splits cover the + range -- dispatching eight would leave two threadgroups writing an empty + online-softmax state that the merge then has to skip. + """ + + if int(selected_tokens) <= 0: + raise ValueError(f"selected_tokens must be positive; got {selected_tokens}") + if int(key_tile) <= 0: + raise ValueError(f"key_tile must be positive; got {key_tile}") + tiles = -(-int(selected_tokens) // int(key_tile)) + splits = min(int(key_splits), tiles) + if splits < 1: + splits = 1 + per_split = -(-tiles // splits) + exact_splits = -(-tiles // per_split) + return tiles, per_split, exact_splits + + +def qsa_sparse_gqa_decode_partial_shape( + rows: int, + key_tile: int = _DEFAULT_TILE[0], + key_splits: int = _DEFAULT_KEY_SPLITS, +) -> tuple[int, int, int, int]: + """Shape of the fp32 partial-state buffer the split pass writes.""" + + _, _, n_splits = qsa_sparse_gqa_decode_split_geometry( + _SELECTED_TOKENS, key_tile, key_splits + ) + return (n_splits, _Q_HEADS, int(rows), _PARTIAL_LD) + + +def _normalized_query_offset(query_offset: Any) -> mx.array | None: + """Accept a host int or a one-element int32 array; emit the ABI's ``[1]``. + + A host int becomes a one-element array rather than a params-block scalar + so the two spellings take the SAME kernel path, and so a tensor-valued + cache offset (``TensorOffsetKVCache``) never forces a graph sync just to + read a position the kernel is about to read anyway. + """ + + if isinstance(query_offset, mx.array): + if query_offset.size != 1: + return None + # A ``TensorOffsetKVCache`` offset is a 0-d int32 array; accept every + # exact integer width and narrow, because a one-element astype costs + # nothing and a refusal here would raise on a perfectly valid cache. + if query_offset.dtype not in _INTEGER_DTYPES: + return None + return query_offset.reshape(1).astype(mx.int32) + if isinstance(query_offset, bool): + return None + try: + value = operator.index(query_offset) + except TypeError: + return None + if value < 0: + return None + return mx.array([value], dtype=mx.int32) + + +def qsa_sparse_gqa_decode_unsupported_reason( + queries: mx.array, + keys: mx.array, + values: mx.array, + block_ids: mx.array, + *, + query_offset: Any, + total_tokens: int, + scale: float, + key_tile: int = _DEFAULT_TILE[0], + dimension_tile: int = _DEFAULT_TILE[1], + key_splits: int = _DEFAULT_KEY_SPLITS, +) -> str | None: + """``None`` when the decode call is on contract, else the precise reason.""" + + extension = _load_extension() + if isinstance(extension, Exception): + return f"the native QSA extension is not built ({extension})" + if not _on_metal_device(): + return "the active MLX device is not an available Metal GPU" + + arrays = (queries, keys, values, block_ids) + if any(not isinstance(array, mx.array) for array in arrays): + return "all tensor inputs must be MLX arrays" + if queries.ndim != 4 or keys.ndim != 4 or values.ndim != 4: + return "Q, K, and V must be rank four" + if block_ids.ndim not in (2, 4): + return "block ids must be rank two [M, 512] or rank four [1, 1, M, 512]" + + batch, query_heads, rows, head_dim = (int(x) for x in queries.shape) + if (batch, query_heads, head_dim) != (_BATCH, _Q_HEADS, _HEAD_DIM): + return "Q must have production shape [1, 24, M, 256]" + if rows <= 0: + return "Q must carry at least one query row" + + key_batch, kv_heads, capacity, key_dim = (int(x) for x in keys.shape) + if (key_batch, kv_heads, key_dim) != (_BATCH, _KV_HEADS, _HEAD_DIM): + return "K must have production shape [1, 2, capacity, 256]" + if tuple(int(x) for x in values.shape) != tuple(int(x) for x in keys.shape): + return "V must have the same full-backing shape as K" + + if queries.dtype not in _SUPPORTED_DTYPES: + return "Q must be float16 or bfloat16" + if keys.dtype != queries.dtype or values.dtype != queries.dtype: + return "Q, K, and V dtypes must match" + if block_ids.dtype not in _SUPPORTED_ID_DTYPES: + return "block ids must be int32 or uint32" + if _normalized_block_ids(block_ids, rows) is None: + return "block ids must have shape [M, 512] or [1, 1, M, 512]" + if _normalized_query_offset(query_offset) is None: + return ( + "query_offset must be a non-negative host int or a one-element " + "int32 array" + ) + + if isinstance(total_tokens, mx.array): + return "total_tokens must be a host integer" + if isinstance(scale, mx.array): + return "scale must be a host float" + if isinstance(total_tokens, bool): + return "total_tokens cannot be bool" + try: + total_tokens_i = operator.index(total_tokens) + except TypeError: + return "total_tokens must be an exact host integer" + if isinstance(scale, bool) or not isinstance(scale, (int, float)): + return "scale must be a numeric host scalar" + scale_f = float(scale) + + if total_tokens_i <= 0: + return "total_tokens must describe a non-empty context" + if rows > total_tokens_i: + return "the query rows must fit inside total_tokens" + if total_tokens_i > capacity: + return "the logical token count exceeds the full K/V backing capacity" + if total_tokens_i > _MAX_CONTEXT: + return "the logical token count exceeds the production context limit" + if total_tokens_i // _COMPRESS_RATIO <= _TOP_K_BLOCKS: + return "the context has not crossed the dense/sparse boundary" + if not math.isfinite(scale_f): + return "scale must be finite" + + if (int(key_tile), int(dimension_tile)) not in _SUPPORTED_TILES: + return ( + "(key_tile, dimension_tile) must be one of " + + ", ".join(str(t) for t in _SUPPORTED_TILES) + ) + if not 1 <= int(key_splits) <= _MAX_KEY_SPLITS: + return f"key_splits must be in [1, {_MAX_KEY_SPLITS}]" + return None + + +def qsa_sparse_gqa_decode_supported( + queries: mx.array, + keys: mx.array, + values: mx.array, + block_ids: mx.array, + *, + query_offset: Any, + total_tokens: int, + scale: float, + key_tile: int = _DEFAULT_TILE[0], + dimension_tile: int = _DEFAULT_TILE[1], + key_splits: int = _DEFAULT_KEY_SPLITS, +) -> bool: + """Whether the exact production-only decode contract is met.""" + + return ( + qsa_sparse_gqa_decode_unsupported_reason( + queries, + keys, + values, + block_ids, + query_offset=query_offset, + total_tokens=total_tokens, + scale=scale, + key_tile=key_tile, + dimension_tile=dimension_tile, + key_splits=key_splits, + ) + is None + ) + + +def qsa_sparse_gqa_decode( + queries: mx.array, + keys: mx.array, + values: mx.array, + block_ids: mx.array, + *, + query_offset: Any, + total_tokens: int, + scale: float, + key_tile: int = _DEFAULT_TILE[0], + dimension_tile: int = _DEFAULT_TILE[1], + key_splits: int = _DEFAULT_KEY_SPLITS, + stream: Any = None, +) -> mx.array: + """Split-K direct-index sparse GQA attention for the decode geometries. + + ``queries`` ``[1, 24, M, 256]`` fp16/bf16 -- the ``[B,H,M,D]`` transposed + view the Attention module already builds. M is 4 for the + fixed-M4 verify and 1 for a single-row draft/decode step. + ``keys``/``values`` ``[1, 2, capacity, 256]``, the FULL KV cache backing, + read in place at its allocation stride. + ``block_ids`` ``[M, 512]`` or ``[1, 1, M, 512]``, int32 or uint32 -- + ``mx.argpartition``'s output IN ITS OWN ORDER. Unlike the + prefill entry point this one makes NO ordering assumption and + reads no ``block_valid``: it applies the shipped lane's own + per-slot predicate ``block < (pos + 1) // 4`` to every slot. + So the visible set is identical whether or not the selector + sorts. + ``query_offset`` absolute position of query row 0, as a host int or a + one-element int32 array (a tensor-valued cache offset never + has to be read on the host). + ``total_tokens`` logical tokens in the cache (NOT ``capacity``). + ``key_splits`` target KV splits; see + :func:`qsa_sparse_gqa_decode_split_geometry` for how it is + clamped and rounded. + + Returns ``[1, 24, M, 256]``, same dtype as ``queries``. + + NUMERICS -- this is a ROUNDING-CLASS change, HumanEval-gated + ----------------------------------------------------------- + Against the shipped rows-gather lane, over an IDENTICAL visible set: + + * scores accumulate through Steel MMA fp32 fragments, not MLX's gemv + tiling, so the 256-term contraction is reassociated; + * the softmax is an fp32 ONLINE softmax in ``exp2`` with the scale + pre-multiplied by ``M_LOG2E``, not an fp32 ``exp`` over a + materialised score row; + * probabilities stay fp32 instead of being cast to bf16 before P@V, + and P@V runs fp32 x fp32 instead of bf16 x bf16 with fp32 accumulate; + * the split-K merge adds one more rescale per query row. + + None of that is bit-exact and none of it can be made so. Adopt this lane + on the same terms as ``MTPLX_QWEN4_HC_M4``: greedy-token agreement plus a + full HumanEval gate, never on a digest comparison. + + The shipped lane builds ``topk*ratio + ratio`` = 2,052 token slots; this + kernel walks ``topk*ratio + ratio - 1`` = 2,051. The dropped slot is the + tail's fourth, whose token is ``((pos+1)//4)*4 + 3``; that is ``> pos`` + for every ``pos``, so the shipped lane always masks it. The visible sets + are equal. + """ + + reason = qsa_sparse_gqa_decode_unsupported_reason( + queries, + keys, + values, + block_ids, + query_offset=query_offset, + total_tokens=total_tokens, + scale=scale, + key_tile=key_tile, + dimension_tile=dimension_tile, + key_splits=key_splits, + ) + if reason is not None: + raise ValueError(f"[mtplx.native.qsa_sparse_gqa_decode] {reason}.") + + extension = _load_extension() + selected = _normalized_block_ids(block_ids, int(queries.shape[2])) + offset = _normalized_query_offset(query_offset) + args = ( + queries, + keys, + values, + selected, + offset, + float(scale), + int(total_tokens), + int(key_tile), + int(dimension_tile), + int(key_splits), + ) + # See qsa_sparse_gqa: omit rather than pass an explicit None. + if stream is None: + return extension.qsa_sparse_gqa_decode(*args) + return extension.qsa_sparse_gqa_decode(*args, stream=stream) diff --git a/mtplx/profiles.py b/mtplx/profiles.py index bff61f963..6ef879fae 100644 --- a/mtplx/profiles.py +++ b/mtplx/profiles.py @@ -437,6 +437,9 @@ def announce_runtime_gated_env( "MTPLX_QWEN4_HC_M4", "MTPLX_QWEN4_PREFILL_MASK_FUSE", "MTPLX_QSA_PREFILL_QUERY_TILE", + "MTPLX_QSA_SPARSE_DECODE", + "MTPLX_QSA_SPARSE_DECODE_TILE", + "MTPLX_QSA_SPARSE_DECODE_SPLITS", "MTPLX_NGRAM_PREWARM_ORDER", "MTPLX_STRICT_CLAIMS", "MTPLX_QWEN4_COMPILED_MTP_PREPARE", diff --git a/mtplx/runtime_options.py b/mtplx/runtime_options.py index bb4bf72b0..52b414619 100644 --- a/mtplx/runtime_options.py +++ b/mtplx/runtime_options.py @@ -251,6 +251,141 @@ def qwen4_hc_m4_enabled() -> bool: return _QWEN4_HC_M4 + +#: Split-K (KV-split) native sparse-GQA attention for the DECODE geometries +#: (native_extensions/qsa_sparse_gqa, mtplx/kernels/qsa_sparse_decode.py). +#: +#: ``MTPLX_QSA_SPARSE_DECODE`` serves the M=4 fixed verify, all 12 QSA +#: layers, once per verify cycle. This is where the bytes are: the shipped +#: lane materialises a [1, 2, 4, 2052, 256] gathered K/V pair per layer +#: (16.8 MB written, then re-read by the score and P@V GEMMs), plus MLX's own +#: 8.4 MB contiguous copy of the transposed key view. The kernel reads the +#: cache rows once and never writes them. +#: +#: Off by default. It RAISES on a contract failure rather than silently +#: reverting -- a silently inert flag is how MTPLX_FUSED_HC_V3 came to be +#: armed-but-dead at M=4. The one thing that does NOT raise is a PARITY +#: failure at install: this kernel is rounding-class, so a parity miss is a +#: numerical verdict, and the lane disables itself for the process and +#: reports the measured deltas. +def _qsa_sparse_decode_import_default() -> bool: + # New key wins for any non-empty value (including "0" for the per-key + # opt-out); the old MTPLX_FABLE_QSA_SPARSE_DECODE name is honoured as an + # alias only when the new key is unset. Read once at import. + raw = os.environ.get("MTPLX_QSA_SPARSE_DECODE") + if raw is None or not str(raw).strip(): + return env_bool("MTPLX_FABLE_QSA_SPARSE_DECODE", default=False) + return env_bool("MTPLX_QSA_SPARSE_DECODE", default=False) + + +_QSA_SPARSE_DECODE = _qsa_sparse_decode_import_default() + + +def qsa_sparse_decode_enabled() -> bool: + """True when the QSA split-K decode flag armed this process at import. + + Armed by ``MTPLX_QSA_SPARSE_DECODE`` (or the old + ``MTPLX_FABLE_QSA_SPARSE_DECODE`` alias). + """ + + return _QSA_SPARSE_DECODE + + +def _parse_sparse_decode_tile(raw: str | None) -> tuple[int, int]: + """``"BK:DC"`` -> the compiled tile pair; unset means the default.""" + + if raw is None or not str(raw).strip(): + return (128, 32) + token = str(raw).strip() + parts = token.split(":") + if len(parts) != 2: + raise ValueError( + f"MTPLX_QSA_SPARSE_DECODE_TILE={raw!r} must be 'BK:DC'" + ) + try: + tile = (int(parts[0]), int(parts[1])) + except ValueError as exc: + raise ValueError( + f"MTPLX_QSA_SPARSE_DECODE_TILE={raw!r} must be 'BK:DC'" + ) from exc + if tile not in QSA_SPARSE_DECODE_TILES: + accepted = ", ".join(f"{a}:{b}" for a, b in QSA_SPARSE_DECODE_TILES) + raise ValueError( + f"MTPLX_QSA_SPARSE_DECODE_TILE={raw!r} is not instantiated; " + f"expected one of: {accepted}" + ) + return tile + + +#: The (BK, DC) pairs the metallib instantiates. Anything else raises rather +#: than falling back, so a typo in a sweep cannot quietly measure the default. +QSA_SPARSE_DECODE_TILES = ((128, 32), (256, 32), (64, 64), (128, 64)) +QSA_SPARSE_DECODE_MAX_SPLITS = 64 + +_QSA_SPARSE_DECODE_TILE = _parse_sparse_decode_tile( + os.environ.get("MTPLX_QSA_SPARSE_DECODE_TILE") + or os.environ.get("MTPLX_FABLE_QSA_SPARSE_DECODE_TILE") +) + + +#: MEASURED default (2026-09-02, guarded micro, M=4, 16K, 12 layers). The +#: kernel is occupancy-bound, and at the shipped tile (BK=128) there are 17 +#: BK-tiles over the 2,051 selected keys, so 17 is the smallest split target +#: that reaches one tile per threadgroup -- a 4 x 2 x 17 = 136-threadgroup +#: grid on a 40-core M5 Max. Everything below it leaves cores idle: +#: +#: splits n_splits threadgroups ms/layer x baseline +#: 4 4 32 0.325 0.70 +#: 8 6 48 0.210 1.08 +#: 16 9 72 0.149 1.52 +#: 17 17 136 0.094-0.099 2.3-2.4 +#: +#: Larger values clamp to the same 17 at BK=128, so 17 is also the point past +#: which the knob stops doing anything -- which is why the first sweep's s17 +#: and s32 rows are the SAME configuration measured twice, and their 5.3% +#: spread is the bench's noise floor rather than a result. +#: +#: The previous default of 8 was a placeholder, and it measured 2.2x slower. +QSA_SPARSE_DECODE_DEFAULT_SPLITS = 17 + + +def _parse_sparse_decode_splits(raw: str | None) -> int: + """``MTPLX_QSA_SPARSE_DECODE_SPLITS`` -- the KV-split target.""" + + if raw is None or not str(raw).strip(): + return QSA_SPARSE_DECODE_DEFAULT_SPLITS + try: + value = int(str(raw).strip()) + except ValueError as exc: + raise ValueError( + f"MTPLX_QSA_SPARSE_DECODE_SPLITS={raw!r} must be an integer" + ) from exc + if not 1 <= value <= QSA_SPARSE_DECODE_MAX_SPLITS: + raise ValueError( + f"MTPLX_QSA_SPARSE_DECODE_SPLITS={raw!r} must be in " + f"[1, {QSA_SPARSE_DECODE_MAX_SPLITS}]" + ) + return value + + +_QSA_SPARSE_DECODE_SPLITS = _parse_sparse_decode_splits( + os.environ.get("MTPLX_QSA_SPARSE_DECODE_SPLITS") + or os.environ.get("MTPLX_FABLE_QSA_SPARSE_DECODE_SPLITS") +) + + +def qsa_sparse_decode_tile() -> tuple[int, int]: + """The armed ``(key_tile, dimension_tile)`` for the decode kernel.""" + + return _QSA_SPARSE_DECODE_TILE + + +def qsa_sparse_decode_splits() -> int: + """The armed KV-split target for the decode kernel.""" + + return _QSA_SPARSE_DECODE_SPLITS + + @dataclass(frozen=True) class ResolvedAPIKey: value: str | None diff --git a/mtplx/server/openai.py b/mtplx/server/openai.py index a17062a5a..93820aa7e 100644 --- a/mtplx/server/openai.py +++ b/mtplx/server/openai.py @@ -960,6 +960,36 @@ def _server_runtime_env_overrides( # (including 0 for whole-chunk) wins via the pop loop below. if os.environ.get("MTPLX_QSA_PREFILL_QUERY_TILE") is None: overrides.setdefault("MTPLX_QSA_PREFILL_QUERY_TILE", "2048") + # PR #391 remainder port (davidtai): the native split-K QSA decode + # lane. Default ON for the fixed-M4 pack, but ONLY when the native + # mtplx_native_qsa extension is built -- a wheel without it declines + # to stock and serves the shipped QSA decode path, so a release + # without the Apple-Silicon native wheel still boots. An explicit + # operator export of MTPLX_QSA_SPARSE_DECODE=1 bypasses this check + # and reaches the fail-closed install (armed + unbuilt -> RAISE), + # which is the measured-arm contract. Its measured companions (the + # 128:32 tile and 17 KV-splits) are the runtime_options defaults, + # so they need no stamp; an explicit export of either still wins. + if os.environ.get("MTPLX_QSA_SPARSE_DECODE") is None: + try: + from mtplx.native import native_qsa_available + + _qsa_decode_ext_ok = bool(native_qsa_available()) + except Exception: + _qsa_decode_ext_ok = False + if _qsa_decode_ext_ok: + overrides.setdefault("MTPLX_QSA_SPARSE_DECODE", "1") + else: + print( + "[mtplx] MTPLX_QSA_SPARSE_DECODE declined to stock: the " + "native mtplx_native_qsa split-K extension is not built " + "in this environment; serving the stock QSA decode " + "path. Build native_extensions/qsa_sparse_gqa to arm " + "the lane, or export MTPLX_QSA_SPARSE_DECODE=1 to " + "require it (armed + unbuilt fails closed at load).", + file=sys.stderr, + flush=True, + ) # The stage-3 child routes are consumed at model load and raise # unless stage 3 itself resolves on, so they are derived from the # resolved parent, never stamped alone: the routed-down reduction, @@ -1119,6 +1149,7 @@ def _served_model_type_is_qwen4_exp(args: argparse.Namespace) -> bool: # the pop loop below. "MTPLX_QWEN4_HC_M4", "MTPLX_QWEN4_PREFILL_MASK_FUSE", + "MTPLX_QSA_SPARSE_DECODE", "MTPLX_NGRAM_PREWARM", ) # Every key the fixed-M4 lane defaults may stamp; an explicit operator @@ -1127,8 +1158,11 @@ def _served_model_type_is_qwen4_exp(args: argparse.Namespace) -> bool: "MTPLX_FRSPEC_DRAFT", "MTPLX_FRSPEC_VOCAB", "MTPLX_QSA_GATHER_MAX_ROWS", - # PR #391 remainder port (davidtai): the QSA prefill query-tile value. + # PR #391 remainder ports (davidtai): the QSA prefill query-tile value and + # the split-K decode lane's tile/splits companions. "MTPLX_QSA_PREFILL_QUERY_TILE", + "MTPLX_QSA_SPARSE_DECODE_TILE", + "MTPLX_QSA_SPARSE_DECODE_SPLITS", ) @@ -18652,6 +18686,15 @@ def _qwen4_install_reports(state: Any) -> dict[str, Any]: hc_m4 = getattr(runtime, "qwen4_hc_m4_report", None) if isinstance(hc_m4, dict): out["hc_m4"] = hc_m4 + try: + from mtplx.runtime_options import qsa_sparse_decode_enabled + + if qsa_sparse_decode_enabled(): + from mtplx.kernels import qsa_sparse_decode as _qsd + + out["qsa_sparse_decode"] = _qsd.receipt() + except Exception: + pass glue = getattr(runtime, "_mtplx_qwen4_verify_glue", None) if isinstance(glue, dict): out["verify_glue"] = glue diff --git a/native_extensions/qsa_sparse_gqa/.gitignore b/native_extensions/qsa_sparse_gqa/.gitignore new file mode 100644 index 000000000..d6367592d --- /dev/null +++ b/native_extensions/qsa_sparse_gqa/.gitignore @@ -0,0 +1,6 @@ +build/ +*.egg-info/ +mtplx_native_qsa/*.so +mtplx_native_qsa/*.dylib +mtplx_native_qsa/*.metallib +mtplx_native_qsa/__pycache__/ diff --git a/native_extensions/qsa_sparse_gqa/CMakeLists.txt b/native_extensions/qsa_sparse_gqa/CMakeLists.txt new file mode 100644 index 000000000..3e22ac1c8 --- /dev/null +++ b/native_extensions/qsa_sparse_gqa/CMakeLists.txt @@ -0,0 +1,168 @@ +cmake_minimum_required(VERSION 3.27) + +project(mtplx_native_qsa_ext LANGUAGES CXX) + +set(CMAKE_CXX_STANDARD 17) +set(CMAKE_CXX_STANDARD_REQUIRED ON) +set(CMAKE_POSITION_INDEPENDENT_CODE ON) + +option(BUILD_SHARED_LIBS "Build extensions as a shared library" ON) + +find_package( + Python 3.11 + COMPONENTS Interpreter Development.Module + REQUIRED) + +# --------------------------------------------------------------------------- +# nanobind selection, and the ABI guard that has to be a CONFIGURE-time error. +# +# This extension can only cast an ``mlx.core.array`` if it shares mlx.core's +# nanobind type registry. Two nanobind modules share one exactly when they +# agree on the capsule key ``__nb_internals____`` +# (nanobind src/nb_internals.cpp, nb_module_exec). The domain is NB_DOMAIN=mlx +# on both sides. The tag is ``"v" NB_INTERNALS_VERSION ... "_" +# NB_PLATFORM_ABI_TAG`` (nanobind src/nb_abi.h), and NB_INTERNALS_VERSION moves +# between nanobind releases -- 19 in 2.12.0, 20 in 2.13.0, 21 in 2.15.0. +# +# When it does not match, the build SUCCEEDS, the module IMPORTS, and then +# every function taking an array raises +# +# TypeError: ...(): incompatible function arguments. +# 1. ...(queries: mlx::core::array, ...) +# Invoked with types: mlx.core.array, ... +# +# -- the C++ name in the signature being the tell that the type was never +# resolved. That is what the first guarded run hit, and the earlier kernel had +# the same latent defect. Failing loudly here costs one configure; not +# failing costs a whole guarded GPU window. +# +# Override with -DMTPLX_NANOBIND_DIR= (the directory +# containing cmake/ and src/), e.g. a uv tool install of a matching release. +# NEVER "fix" a mismatch by defining NB_INTERNALS_VERSION: that macro guards +# the LAYOUT of the shared nb_internals struct, so forcing the tag to match +# while the struct differs makes two modules scribble on each other. +# --------------------------------------------------------------------------- +set(MTPLX_NANOBIND_DIR "" CACHE PATH + "nanobind package root to build against (default: the interpreter's)") + +if(MTPLX_NANOBIND_DIR) + set(nanobind_ROOT "${MTPLX_NANOBIND_DIR}/cmake") + set(MTPLX_NANOBIND_SRC "${MTPLX_NANOBIND_DIR}") +else() + execute_process( + COMMAND "${Python_EXECUTABLE}" -m nanobind --cmake_dir + OUTPUT_STRIP_TRAILING_WHITESPACE + OUTPUT_VARIABLE nanobind_ROOT) + get_filename_component(MTPLX_NANOBIND_SRC "${nanobind_ROOT}" DIRECTORY) +endif() +find_package(nanobind CONFIG REQUIRED) + +execute_process( + COMMAND "${Python_EXECUTABLE}" -m mlx --cmake-dir + OUTPUT_STRIP_TRAILING_WHITESPACE + OUTPUT_VARIABLE MLX_ROOT) +find_package(MLX CONFIG REQUIRED) + +# --- read NB_INTERNALS_VERSION out of the nanobind we are about to use ------ +set(MTPLX_NB_INTERNALS "") +foreach(_hdr "src/nb_abi.h" "src/nb_internals.h") + if(NOT MTPLX_NB_INTERNALS AND EXISTS "${MTPLX_NANOBIND_SRC}/${_hdr}") + file(STRINGS "${MTPLX_NANOBIND_SRC}/${_hdr}" _nb_lines + REGEX "define[ ]+NB_INTERNALS_VERSION[ ]+[0-9]+") + foreach(_nb_line IN LISTS _nb_lines) + if(NOT MTPLX_NB_INTERNALS) + string(REGEX MATCH "NB_INTERNALS_VERSION[ ]+([0-9]+)" _m "${_nb_line}") + if(CMAKE_MATCH_1) + set(MTPLX_NB_INTERNALS "${CMAKE_MATCH_1}") + endif() + endif() + endforeach() + endif() +endforeach() + +# --- read the ABI tag mlx.core was actually built with --------------------- +set(MTPLX_MLX_NB_INTERNALS "") +file(GLOB MTPLX_MLX_CORE "${MLX_ROOT}/core*.so") +if(MTPLX_MLX_CORE) + list(GET MTPLX_MLX_CORE 0 MTPLX_MLX_CORE_SO) + file(STRINGS "${MTPLX_MLX_CORE_SO}" _mlx_tags + REGEX "v[0-9]+[0-9a-zA-Z._-]*_[0-9a-zA-Z_]*libcpp[0-9a-zA-Z_]*") + foreach(_tag IN LISTS _mlx_tags) + if(NOT MTPLX_MLX_NB_INTERNALS) + string(REGEX MATCH "v([0-9]+)[0-9a-zA-Z._-]*_[0-9a-zA-Z_]*libcpp" _m "${_tag}") + if(CMAKE_MATCH_1) + set(MTPLX_MLX_NB_INTERNALS "${CMAKE_MATCH_1}") + endif() + endif() + endforeach() +endif() + +if(MTPLX_NB_INTERNALS AND MTPLX_MLX_NB_INTERNALS) + if(NOT MTPLX_NB_INTERNALS STREQUAL MTPLX_MLX_NB_INTERNALS) + message(FATAL_ERROR + "nanobind ABI mismatch: mlx.core (${MTPLX_MLX_CORE_SO}) was built with " + "nanobind internals v${MTPLX_MLX_NB_INTERNALS}, but this build would use " + "v${MTPLX_NB_INTERNALS} from ${MTPLX_NANOBIND_SRC}.\n" + "The two would get separate __nb_internals__mlx__ capsules, the " + "extension could not resolve mlx::core::array, and EVERY call would " + "raise TypeError at run time.\n" + "Re-run with -DMTPLX_NANOBIND_DIR=.\n" + "Compare src/nb_abi.h in the nanobind that built mlx.core with the " + "one this build resolved.") + endif() + message(STATUS + "nanobind internals v${MTPLX_NB_INTERNALS} matches mlx.core; " + "the extension will share its type registry") +else() + message(WARNING + "could not compare nanobind ABI tags (nanobind='${MTPLX_NB_INTERNALS}', " + "mlx.core='${MTPLX_MLX_NB_INTERNALS}'); if calls fail with " + "'incompatible function arguments' listing mlx::core::array, compare " + "the two nanobind ABI tags by hand") +endif() + +add_library(mtplx_native_qsa_ext) +target_sources( + mtplx_native_qsa_ext + PUBLIC ${CMAKE_CURRENT_LIST_DIR}/sparse_gqa/qsa_sparse_gqa.cpp + ${CMAKE_CURRENT_LIST_DIR}/sparse_gqa/qsa_sparse_gqa_decode.cpp) +target_include_directories(mtplx_native_qsa_ext PUBLIC ${CMAKE_CURRENT_LIST_DIR}) +target_link_libraries(mtplx_native_qsa_ext PUBLIC mlx) + +if(MLX_BUILD_METAL) + mlx_build_metallib( + TARGET + mtplx_native_qsa_ext_metallib + TITLE + mtplx_native_qsa_ext + SOURCES + ${CMAKE_CURRENT_LIST_DIR}/sparse_gqa/qsa_sparse_gqa.metal + INCLUDE_DIRS + ${PROJECT_SOURCE_DIR} + ${MLX_INCLUDE_DIRS} + DEPS + ${CMAKE_CURRENT_LIST_DIR}/sparse_gqa/steel_qsa_sparse_gqa.h + ${CMAKE_CURRENT_LIST_DIR}/sparse_gqa/qsa_sparse_gqa_params.h + ${CMAKE_CURRENT_LIST_DIR}/sparse_gqa/steel_qsa_sparse_gqa_decode.h + ${CMAKE_CURRENT_LIST_DIR}/sparse_gqa/qsa_sparse_gqa_decode_params.h + OUTPUT_DIRECTORY + ${CMAKE_LIBRARY_OUTPUT_DIRECTORY}) + + add_dependencies(mtplx_native_qsa_ext mtplx_native_qsa_ext_metallib) +endif() + +nanobind_add_module( + _ext + NB_STATIC + STABLE_ABI + LTO + NOMINSIZE + NB_DOMAIN + mlx + ${CMAKE_CURRENT_LIST_DIR}/bindings.cpp) +target_link_libraries(_ext PRIVATE mtplx_native_qsa_ext) + +if(BUILD_SHARED_LIBS) + target_link_options(_ext PRIVATE -Wl,-rpath,@loader_path) +endif() diff --git a/native_extensions/qsa_sparse_gqa/bindings.cpp b/native_extensions/qsa_sparse_gqa/bindings.cpp new file mode 100644 index 000000000..509b3378c --- /dev/null +++ b/native_extensions/qsa_sparse_gqa/bindings.cpp @@ -0,0 +1,58 @@ +// SPDX-License-Identifier: Apache-2.0 + +#include +#include +#include + +#include "sparse_gqa/qsa_sparse_gqa.h" +#include "sparse_gqa/qsa_sparse_gqa_decode.h" + +#include + +namespace nb = nanobind; +using namespace nb::literals; + +NB_MODULE(_ext, m) { + m.doc() = "MTPLX native QSA direct-index sparse-GQA attention"; + + m.def("qsa_sparse_gqa_attention", &mtplx_native::qsa_sparse_gqa_attention, + "queries"_a, "keys"_a, "values"_a, "selected_blocks"_a, "scale"_a, + "q_offset"_a, "key_length"_a = -1, "key_tile"_a = 128, + "dimension_tile"_a = 32, nb::kw_only(), "stream"_a = nb::none(), + "Direct-index sparse GQA attention over chronological QSA block ids."); + + m.def("qsa_sparse_gqa_unsupported_reason", + &mtplx_native::qsa_sparse_gqa_unsupported_reason, "queries"_a, "keys"_a, + "values"_a, "selected_blocks"_a, "scale"_a, "q_offset"_a, + "key_length"_a = -1, "key_tile"_a = 128, "dimension_tile"_a = 32, + nb::kw_only(), "stream"_a = nb::none(), + "Empty string when the call is on contract; otherwise the reason."); + + m.def( + "qsa_sparse_gqa_decode", &mtplx_native::qsa_sparse_gqa_decode, + "queries"_a, "keys"_a, "values"_a, "selected_blocks"_a, + "query_offset"_a, "scale"_a, "key_length"_a = -1, "key_tile"_a = 128, + "dimension_tile"_a = 32, "key_splits"_a = 8, nb::kw_only(), + "stream"_a = nb::none(), + "Split-K direct-index sparse GQA attention for the decode geometries."); + + m.def( + "qsa_sparse_gqa_decode_unsupported_reason", + &mtplx_native::qsa_sparse_gqa_decode_unsupported_reason, "queries"_a, + "keys"_a, "values"_a, "selected_blocks"_a, "query_offset"_a, "scale"_a, + "key_length"_a = -1, "key_tile"_a = 128, "dimension_tile"_a = 32, + "key_splits"_a = 8, nb::kw_only(), "stream"_a = nb::none(), + "Empty string when the decode call is on contract; else the reason."); + + m.def( + "qsa_sparse_gqa_decode_split_geometry", + [](int selected_tokens, int key_tile, int key_splits) { + int n_tiles = 0, tiles_per_split = 0, n_splits = 0; + mtplx_native::qsa_sparse_gqa_decode_split_geometry( + selected_tokens, key_tile, key_splits, &n_tiles, &tiles_per_split, + &n_splits); + return std::make_tuple(n_tiles, tiles_per_split, n_splits); + }, + "selected_tokens"_a, "key_tile"_a, "key_splits"_a, + "(n_tiles, tiles_per_split, n_splits) for the split-K decode grid."); +} diff --git a/native_extensions/qsa_sparse_gqa/mtplx_native_qsa/__init__.py b/native_extensions/qsa_sparse_gqa/mtplx_native_qsa/__init__.py new file mode 100644 index 000000000..18e8e9aae --- /dev/null +++ b/native_extensions/qsa_sparse_gqa/mtplx_native_qsa/__init__.py @@ -0,0 +1,15 @@ +from ._ext import ( + qsa_sparse_gqa_attention, + qsa_sparse_gqa_decode, + qsa_sparse_gqa_decode_split_geometry, + qsa_sparse_gqa_decode_unsupported_reason, + qsa_sparse_gqa_unsupported_reason, +) + +__all__ = [ + "qsa_sparse_gqa_attention", + "qsa_sparse_gqa_decode", + "qsa_sparse_gqa_decode_split_geometry", + "qsa_sparse_gqa_decode_unsupported_reason", + "qsa_sparse_gqa_unsupported_reason", +] diff --git a/native_extensions/qsa_sparse_gqa/pyproject.toml b/native_extensions/qsa_sparse_gqa/pyproject.toml new file mode 100644 index 000000000..aa46b980a --- /dev/null +++ b/native_extensions/qsa_sparse_gqa/pyproject.toml @@ -0,0 +1,8 @@ +[build-system] +requires = [ + "setuptools>=42", + "cmake>=3.25", + "mlx>=0.32.0", + "nanobind==2.12.0", +] +build-backend = "setuptools.build_meta" diff --git a/native_extensions/qsa_sparse_gqa/setup.py b/native_extensions/qsa_sparse_gqa/setup.py new file mode 100644 index 000000000..0fddc755b --- /dev/null +++ b/native_extensions/qsa_sparse_gqa/setup.py @@ -0,0 +1,17 @@ +from setuptools import setup + +from mlx import extension + + +if __name__ == "__main__": + setup( + name="mtplx_native_qsa", + version="0.0.0", + description="Native MLX direct-index sparse-GQA attention for MTPLX QSA.", + ext_modules=[extension.CMakeExtension("mtplx_native_qsa._ext")], + cmdclass={"build_ext": extension.CMakeBuild}, + packages=["mtplx_native_qsa"], + package_data={"mtplx_native_qsa": ["*.so", "*.dylib", "*.metallib"]}, + zip_safe=False, + python_requires=">=3.11", + ) diff --git a/native_extensions/qsa_sparse_gqa/sparse_gqa/qsa_sparse_gqa.cpp b/native_extensions/qsa_sparse_gqa/sparse_gqa/qsa_sparse_gqa.cpp new file mode 100644 index 000000000..731d17745 --- /dev/null +++ b/native_extensions/qsa_sparse_gqa/sparse_gqa/qsa_sparse_gqa.cpp @@ -0,0 +1,278 @@ +// SPDX-License-Identifier: Apache-2.0 +// +// Host side of MTPLX's port of oMLX's direct-index sparse-GQA attention +// (jundot/omlx 7467dce8, Jonathan Spangler, Apache-2.0). See +// steel_qsa_sparse_gqa.h for provenance, the algorithm, and the list of +// MTPLX-side changes. + +#include "sparse_gqa/qsa_sparse_gqa.h" + +#include +#include +#include +#include + +#include "mlx/backend/common/utils.h" +#include "mlx/backend/metal/device.h" +#include "mlx/backend/metal/utils.h" +#include "mlx/ops.h" +#include "mlx/primitives.h" +#include "mlx/utils.h" + +#include "sparse_gqa/qsa_sparse_gqa_params.h" + +namespace mtplx_native { + +namespace { + +using namespace mlx::core; + +constexpr int kBatch = 1; +constexpr int kQHeads = 24; +constexpr int kKvHeads = 2; +constexpr int kGqa = 12; +constexpr int kHeadDim = 256; +constexpr int kTopK = 512; +constexpr int kHeadPad = 16; +constexpr int kWarps = 2; + +std::string current_binary_dir() { + static std::string binary_dir = []() { + Dl_info info; + if (!dladdr(reinterpret_cast(¤t_binary_dir), &info)) { + throw std::runtime_error("Unable to get mtplx_native_qsa binary dir."); + } + return std::filesystem::path(info.dli_fname).parent_path().string(); + }(); + return binary_dir; +} + +bool last_dim_contiguous(const array &arr) { return arr.strides(-1) == 1; } + +std::string shape_str(const array &arr) { + std::ostringstream out; + out << arr.shape(); + return out.str(); +} + +bool supported_tile(int key_tile, int dimension_tile) { + return (dimension_tile == 32 && (key_tile == 128 || key_tile == 256)) || + (dimension_tile == 64 && (key_tile == 64 || key_tile == 128)); +} + +// A single reason string, so the Python gate and the raised exception never +// disagree about why a call was refused. +std::string unsupported_reason(const array &q, const array &k, const array &v, + const array &selected, float scale, int q_offset, + int key_length, int key_tile, int dimension_tile, + Stream stream) { + if (stream.device == Device::cpu) { + return "the QSA sparse-GQA kernel has no CPU path"; + } + if (q.ndim() != 4 || k.ndim() != 4 || v.ndim() != 4 || selected.ndim() != 4) { + return "queries, keys, values, and block ids must all be rank four"; + } + if (q.dtype() != float16 && q.dtype() != bfloat16) { + return "queries must be float16 or bfloat16"; + } + if (k.dtype() != q.dtype() || v.dtype() != q.dtype()) { + return "queries, keys, and values must share one dtype"; + } + if (selected.dtype() != uint32 && selected.dtype() != int32) { + return "block ids must be uint32 or int32"; + } + if (!last_dim_contiguous(q) || !last_dim_contiguous(k) || + !last_dim_contiguous(v) || !last_dim_contiguous(selected)) { + return "every input must be contiguous in its last dimension"; + } + if (q.shape(0) != kBatch || q.shape(1) != kQHeads || q.shape(3) != kHeadDim) { + return "queries must have production shape [1, 24, M, 256]; got " + + shape_str(q); + } + const int rows = q.shape(2); + if (rows <= 0) { + return "queries must carry at least one row"; + } + if (k.shape(0) != kBatch || k.shape(1) != kKvHeads || k.shape(3) != kHeadDim) { + return "keys must have production shape [1, 2, capacity, 256]; got " + + shape_str(k); + } + if (v.shape() != k.shape()) { + return "values must have the same full-backing shape as keys"; + } + if (selected.shape(0) != kBatch || selected.shape(1) != 1 || + selected.shape(2) != rows || selected.shape(3) != kTopK) { + return "block ids must have shape [1, 1, M, 512]; got " + + shape_str(selected); + } + const int capacity = k.shape(2); + if (key_length <= 0 || key_length > capacity) { + return "key_length must be a positive count within the K/V backing"; + } + if (q_offset < 0 || q_offset + rows > key_length) { + return "queries must be a causal suffix inside key_length"; + } + if (!std::isfinite(scale)) { + return "scale must be finite"; + } + if (!supported_tile(key_tile, dimension_tile)) { + return "(key_tile, dimension_tile) must be one of " + "(128,32), (256,32), (64,64), (128,64)"; + } + return std::string(); +} + +class QsaSparseGqaPrimitive : public Primitive { +public: + QsaSparseGqaPrimitive(Stream stream, float scale, int q_offset, + int key_length, int key_tile, int dimension_tile) + : Primitive(stream), scale_(scale), q_offset_(q_offset), + key_length_(key_length), key_tile_(key_tile), + dimension_tile_(dimension_tile) {} + + void eval_cpu(const std::vector & /* inputs */, + std::vector & /* outputs */) override { + throw std::runtime_error( + "[mtplx_native_qsa] qsa_sparse_gqa has no CPU path."); + } + + void eval_gpu(const std::vector &inputs, + std::vector &outputs) override { + auto &stream = this->stream(); + auto &device = metal::device(stream.device); + const auto &q = inputs[0]; + const auto &k = inputs[1]; + const auto &v = inputs[2]; + const auto &selected = inputs[3]; + auto &out = outputs[0]; + + // Authoritative alignment guard. The builder below validates the + // NOMINAL strides an unevaluated array reports; only here are the strides + // final. Every K/V/Q/O row is read or written with 128-bit `uint4` + // accesses, so each leading stride must be a whole number of 16-byte + // words and the last dimension must be unit stride. Getting this wrong + // is silent corruption, not a crash. + auto require_vector_strides = [](const array &a, const char *which) { + const int per_word = 16 / static_cast(a.itemsize()); + if (a.strides(-1) != 1) { + throw std::runtime_error( + std::string("[mtplx_native_qsa] ") + which + + " must be contiguous in its last dimension."); + } + for (int d = 1; d + 1 < static_cast(a.ndim()); ++d) { + if (a.strides(d) % per_word != 0) { + throw std::runtime_error( + std::string("[mtplx_native_qsa] ") + which + + " strides must be whole 128-bit words for the vector loads."); + } + } + }; + require_vector_strides(q, "queries"); + require_vector_strides(k, "keys"); + require_vector_strides(v, "values"); + require_vector_strides(out, "output"); + + out.set_data(allocator::malloc(out.nbytes())); + MtplxQsaSparseGqaParams params{ + /* B */ kBatch, + /* q_heads */ kQHeads, + /* kv_heads */ kKvHeads, + /* qL */ q.shape(2), + /* kL */ key_length_, + /* topk */ selected.shape(3), + /* gqa_factor */ kGqa, + /* q_offset */ q_offset_, + /* scale */ scale_, + /* _pad */ 0, + /* Q_strides */ {q.strides(0), q.strides(1), q.strides(2)}, + /* K_strides */ {k.strides(0), k.strides(1), k.strides(2)}, + /* V_strides */ {v.strides(0), v.strides(1), v.strides(2)}, + /* Topk_strides */ + {selected.strides(0), selected.strides(1), selected.strides(2)}, + /* O_strides */ {out.strides(0), out.strides(1), out.strides(2)}}; + + std::string kernel_name; + concatenate(kernel_name, "mtplx_qsa_sparse_gqa_", type_to_name(q), "_", + type_to_name(selected), "_bk", key_tile_, "_dc", + dimension_tile_, "_gqa", kGqa, "_hp", kHeadPad, "_d", kHeadDim, + "_wm", kWarps); + + auto library = + device.get_library("mtplx_native_qsa_ext", current_binary_dir()); + auto kernel = device.get_kernel(kernel_name, library); + auto &encoder = metal::get_command_encoder(stream); + encoder.set_compute_pipeline_state(kernel); + encoder.set_input_array(q, 0); + encoder.set_input_array(k, 1); + encoder.set_input_array(v, 2); + encoder.set_input_array(selected, 3); + encoder.set_output_array(out, 4); + encoder.set_bytes(params, 5); + encoder.dispatch_threadgroups(MTL::Size(q.shape(2), kKvHeads, 1), + MTL::Size(32, kWarps, 1)); + } + + DEFINE_NAME(MtplxQsaSparseGqaAttention) + DEFINE_INPUT_OUTPUT_SHAPE() + + bool is_equivalent(const Primitive &other) const override { + const auto &rhs = static_cast(other); + return scale_ == rhs.scale_ && q_offset_ == rhs.q_offset_ && + key_length_ == rhs.key_length_ && key_tile_ == rhs.key_tile_ && + dimension_tile_ == rhs.dimension_tile_; + } + + auto state() const { + return std::make_tuple(scale_, q_offset_, key_length_, key_tile_, + dimension_tile_); + } + +private: + float scale_; + int q_offset_; + int key_length_; + int key_tile_; + int dimension_tile_; +}; + +} // namespace + +std::string qsa_sparse_gqa_unsupported_reason( + const mx::array &queries, const mx::array &keys, const mx::array &values, + const mx::array &selected_blocks, float scale, int q_offset, int key_length, + int key_tile, int dimension_tile, mx::StreamOrDevice s) { + auto stream = to_stream(s); + const int resolved = + key_length < 0 ? (keys.ndim() == 4 ? keys.shape(2) : -1) : key_length; + return unsupported_reason(queries, keys, values, selected_blocks, scale, + q_offset, resolved, key_tile, dimension_tile, + stream); +} + +mx::array qsa_sparse_gqa_attention(const mx::array &queries, + const mx::array &keys, + const mx::array &values, + const mx::array &selected_blocks, + float scale, int q_offset, int key_length, + int key_tile, int dimension_tile, + mx::StreamOrDevice s) { + auto stream = to_stream(s); + const int resolved = + key_length < 0 ? (keys.ndim() == 4 ? keys.shape(2) : -1) : key_length; + auto reason = + unsupported_reason(queries, keys, values, selected_blocks, scale, + q_offset, resolved, key_tile, dimension_tile, stream); + if (!reason.empty()) { + throw std::invalid_argument("[mtplx_native_qsa.qsa_sparse_gqa] " + reason + + "."); + } + + Shape out_shape = queries.shape(); + return array(std::move(out_shape), queries.dtype(), + std::make_shared( + stream, scale, q_offset, resolved, key_tile, + dimension_tile), + std::vector{queries, keys, values, selected_blocks}); +} + +} // namespace mtplx_native diff --git a/native_extensions/qsa_sparse_gqa/sparse_gqa/qsa_sparse_gqa.h b/native_extensions/qsa_sparse_gqa/sparse_gqa/qsa_sparse_gqa.h new file mode 100644 index 000000000..5ebfc7880 --- /dev/null +++ b/native_extensions/qsa_sparse_gqa/sparse_gqa/qsa_sparse_gqa.h @@ -0,0 +1,37 @@ +// SPDX-License-Identifier: Apache-2.0 +// +// Host entry point for MTPLX's port of oMLX's direct-index sparse-GQA +// attention (see steel_qsa_sparse_gqa.h for provenance and algorithm). + +#pragma once + +#include + +#include "mlx/array.h" +#include "mlx/stream.h" +#include "mlx/utils.h" + +namespace mx = mlx::core; + +namespace mtplx_native { + +/// Empty when the call is on the supported contract; otherwise a precise, +/// caller-facing reason. Exposed so the Python lane can gate without +/// catching an exception (mirrors qsa_prefill_flash's _unsupported_reason). +std::string qsa_sparse_gqa_unsupported_reason( + const mx::array &queries, const mx::array &keys, const mx::array &values, + const mx::array &selected_blocks, float scale, int q_offset, + int key_length, int key_tile, int dimension_tile, mx::StreamOrDevice s = {}); + +/// queries [1, 24, M, 256] fp16/bf16, last dim contiguous +/// keys/values [1, 2, cap, 256] the FULL cache backing, same dtype +/// selected [1, 1, M, 512] uint32 or int32, chronological block ids +/// key_length logical tokens in the cache (<= cap); -1 means keys.shape(2) +/// returns [1, 24, M, 256] same dtype as queries +mx::array qsa_sparse_gqa_attention( + const mx::array &queries, const mx::array &keys, const mx::array &values, + const mx::array &selected_blocks, float scale, int q_offset, + int key_length = -1, int key_tile = 128, int dimension_tile = 32, + mx::StreamOrDevice s = {}); + +} // namespace mtplx_native diff --git a/native_extensions/qsa_sparse_gqa/sparse_gqa/qsa_sparse_gqa.metal b/native_extensions/qsa_sparse_gqa/sparse_gqa/qsa_sparse_gqa.metal new file mode 100644 index 000000000..9cb246aae --- /dev/null +++ b/native_extensions/qsa_sparse_gqa/sparse_gqa/qsa_sparse_gqa.metal @@ -0,0 +1,58 @@ +// SPDX-License-Identifier: Apache-2.0 + +// Include order is load-bearing: mlx's utils.h provides Limits and the +// instantiate_kernel macro used by the specialized QSA kernel. +// clang-format off +#include "mlx/backend/metal/kernels/utils.h" +#include "sparse_gqa/steel_qsa_sparse_gqa.h" +// clang-format on + +// MTPLX's `_select_eager` emits int32 block ids; oMLX's ABI is uint32. Both +// are instantiated so the lane never pays an 8 MB astype per layer per chunk. +#define instantiate_qsa_sparse_gqa(tname, dtype, iname, itype, bk, dc) \ + instantiate_kernel("mtplx_qsa_sparse_gqa_" #tname "_" #iname "_bk" #bk "_dc" #dc \ + "_gqa12_hp16_d256_wm2", \ + mtplx_qsa_sparse_gqa_attention, dtype, bk, dc, 12, 16, \ + 256, 2, itype, float) + +#define instantiate_qsa_sparse_gqa_tiles(tname, dtype, iname, itype) \ + instantiate_qsa_sparse_gqa(tname, dtype, iname, itype, 128, 32); \ + instantiate_qsa_sparse_gqa(tname, dtype, iname, itype, 256, 32); \ + instantiate_qsa_sparse_gqa(tname, dtype, iname, itype, 64, 64); \ + instantiate_qsa_sparse_gqa(tname, dtype, iname, itype, 128, 64) + +#define instantiate_qsa_sparse_gqa_dtype(tname, dtype) \ + instantiate_qsa_sparse_gqa_tiles(tname, dtype, uint32, uint); \ + instantiate_qsa_sparse_gqa_tiles(tname, dtype, int32, int) + +instantiate_qsa_sparse_gqa_dtype(float16, half); +instantiate_qsa_sparse_gqa_dtype(bfloat16, bfloat16_t); + +// --------------------------------------------------------------------------- +// Split-K (KV-split) DECODE variant: M=4 verify and M=1 draft. Same four +// (BK, DC) tiles as the prefill kernel so one harness sweep serves both, and +// the same {fp16, bf16} x {uint32, int32} matrix so no lane pays an astype. +// See steel_qsa_sparse_gqa_decode.h for why decode has to be split-K. +// --------------------------------------------------------------------------- +#include "sparse_gqa/steel_qsa_sparse_gqa_decode.h" + +#define instantiate_qsa_sparse_gqa_decode(tname, dtype, iname, itype, bk, dc) \ + instantiate_kernel("mtplx_qsa_sparse_gqa_decode_split_" #tname "_" #iname \ + "_bk" #bk "_dc" #dc "_gqa12_hp16_d256_wm2", \ + mtplx_qsa_sparse_gqa_decode_split, dtype, bk, dc, 12, 16, \ + 256, 2, itype, float) + +#define instantiate_qsa_sparse_gqa_decode_tiles(tname, dtype, iname, itype) \ + instantiate_qsa_sparse_gqa_decode(tname, dtype, iname, itype, 128, 32); \ + instantiate_qsa_sparse_gqa_decode(tname, dtype, iname, itype, 256, 32); \ + instantiate_qsa_sparse_gqa_decode(tname, dtype, iname, itype, 64, 64); \ + instantiate_qsa_sparse_gqa_decode(tname, dtype, iname, itype, 128, 64) + +#define instantiate_qsa_sparse_gqa_decode_dtype(tname, dtype) \ + instantiate_qsa_sparse_gqa_decode_tiles(tname, dtype, uint32, uint); \ + instantiate_qsa_sparse_gqa_decode_tiles(tname, dtype, int32, int); \ + instantiate_kernel("mtplx_qsa_sparse_gqa_decode_merge_" #tname "_d256", \ + mtplx_qsa_sparse_gqa_decode_merge, dtype, 256, float) + +instantiate_qsa_sparse_gqa_decode_dtype(float16, half); +instantiate_qsa_sparse_gqa_decode_dtype(bfloat16, bfloat16_t); diff --git a/native_extensions/qsa_sparse_gqa/sparse_gqa/qsa_sparse_gqa_decode.cpp b/native_extensions/qsa_sparse_gqa/sparse_gqa/qsa_sparse_gqa_decode.cpp new file mode 100644 index 000000000..411d6e5fa --- /dev/null +++ b/native_extensions/qsa_sparse_gqa/sparse_gqa/qsa_sparse_gqa_decode.cpp @@ -0,0 +1,354 @@ +// SPDX-License-Identifier: Apache-2.0 +// +// Host side of the SPLIT-K decode variant of MTPLX's direct-index sparse-GQA +// attention. See steel_qsa_sparse_gqa_decode.h for provenance, the algorithm +// and the list of differences from the phase-1 prefill kernel. + +#include "sparse_gqa/qsa_sparse_gqa_decode.h" + +#include +#include +#include +#include +#include + +#include "mlx/allocator.h" +#include "mlx/backend/common/utils.h" +#include "mlx/backend/metal/device.h" +#include "mlx/backend/metal/utils.h" +#include "mlx/ops.h" +#include "mlx/primitives.h" +#include "mlx/utils.h" + +#include "sparse_gqa/qsa_sparse_gqa_decode_params.h" + +namespace mtplx_native { + +namespace { + +using namespace mlx::core; + +constexpr int kBatch = 1; +constexpr int kQHeads = 24; +constexpr int kKvHeads = 2; +constexpr int kGqa = 12; +constexpr int kHeadDim = 256; +constexpr int kTopK = 512; +constexpr int kCompressRatio = 4; +constexpr int kHeadPad = 16; +constexpr int kWarps = 2; +//: topk*ratio selected tokens plus the at-most ratio-1 causal tail tokens. +constexpr int kSelectedTokens = kTopK * kCompressRatio + (kCompressRatio - 1); +//: Partial rows carry [O(head_dim) | m | l]. +constexpr int kPartialLd = kHeadDim + 2; +//: A merge threadgroup is one thread per head dim; keep the cap honest. +constexpr int kMaxKeySplits = 64; + +std::string current_binary_dir_decode() { + static std::string binary_dir = []() { + Dl_info info; + if (!dladdr(reinterpret_cast(¤t_binary_dir_decode), &info)) { + throw std::runtime_error("Unable to get mtplx_native_qsa binary dir."); + } + return std::filesystem::path(info.dli_fname).parent_path().string(); + }(); + return binary_dir; +} + +bool last_dim_contiguous(const array &arr) { return arr.strides(-1) == 1; } + +std::string shape_str(const array &arr) { + std::ostringstream out; + out << arr.shape(); + return out.str(); +} + +bool supported_tile(int key_tile, int dimension_tile) { + return (dimension_tile == 32 && (key_tile == 128 || key_tile == 256)) || + (dimension_tile == 64 && (key_tile == 64 || key_tile == 128)); +} + +std::string unsupported_reason(const array &q, const array &k, const array &v, + const array &selected, const array &offset, + float scale, int key_length, int key_tile, + int dimension_tile, int key_splits, + Stream stream) { + if (stream.device == Device::cpu) { + return "the QSA sparse-GQA decode kernel has no CPU path"; + } + if (q.ndim() != 4 || k.ndim() != 4 || v.ndim() != 4 || selected.ndim() != 4) { + return "queries, keys, values, and block ids must all be rank four"; + } + if (q.dtype() != float16 && q.dtype() != bfloat16) { + return "queries must be float16 or bfloat16"; + } + if (k.dtype() != q.dtype() || v.dtype() != q.dtype()) { + return "queries, keys, and values must share one dtype"; + } + if (selected.dtype() != uint32 && selected.dtype() != int32) { + return "block ids must be uint32 or int32"; + } + if (offset.dtype() != int32) { + return "the query offset must be an int32 array"; + } + if (offset.size() != 1) { + return "the query offset must hold exactly one element"; + } + if (!last_dim_contiguous(q) || !last_dim_contiguous(k) || + !last_dim_contiguous(v) || !last_dim_contiguous(selected)) { + return "every input must be contiguous in its last dimension"; + } + if (q.shape(0) != kBatch || q.shape(1) != kQHeads || q.shape(3) != kHeadDim) { + return "queries must have production shape [1, 24, M, 256]; got " + + shape_str(q); + } + const int rows = q.shape(2); + if (rows <= 0) { + return "queries must carry at least one row"; + } + if (k.shape(0) != kBatch || k.shape(1) != kKvHeads || k.shape(3) != kHeadDim) { + return "keys must have production shape [1, 2, capacity, 256]; got " + + shape_str(k); + } + if (v.shape() != k.shape()) { + return "values must have the same full-backing shape as keys"; + } + if (selected.shape(0) != kBatch || selected.shape(1) != 1 || + selected.shape(2) != rows || selected.shape(3) != kTopK) { + return "block ids must have shape [1, 1, M, 512]; got " + + shape_str(selected); + } + const int capacity = k.shape(2); + if (key_length <= 0 || key_length > capacity) { + return "key_length must be a positive count within the K/V backing"; + } + if (rows > key_length) { + return "the query rows must fit inside key_length"; + } + if (!std::isfinite(scale)) { + return "scale must be finite"; + } + if (!supported_tile(key_tile, dimension_tile)) { + return "(key_tile, dimension_tile) must be one of " + "(128,32), (256,32), (64,64), (128,64)"; + } + if (key_splits < 1 || key_splits > kMaxKeySplits) { + return "key_splits must be in [1, 64]"; + } + return std::string(); +} + +class QsaSparseGqaDecodePrimitive : public Primitive { +public: + QsaSparseGqaDecodePrimitive(Stream stream, float scale, int key_length, + int key_tile, int dimension_tile, int key_splits) + : Primitive(stream), scale_(scale), key_length_(key_length), + key_tile_(key_tile), dimension_tile_(dimension_tile), + key_splits_(key_splits) {} + + void eval_cpu(const std::vector & /* inputs */, + std::vector & /* outputs */) override { + throw std::runtime_error( + "[mtplx_native_qsa] qsa_sparse_gqa_decode has no CPU path."); + } + + void eval_gpu(const std::vector &inputs, + std::vector &outputs) override { + auto &stream = this->stream(); + auto &device = metal::device(stream.device); + const auto &q = inputs[0]; + const auto &k = inputs[1]; + const auto &v = inputs[2]; + const auto &selected = inputs[3]; + const auto &offset = inputs[4]; + auto &out = outputs[0]; + + // Authoritative alignment guard: only here are the strides final. Every + // K/V/Q row is read with 128-bit `uint4` accesses, so each leading stride + // must be a whole number of 16-byte words and the last dimension must be + // unit stride. Getting this wrong is silent corruption, not a crash. + auto require_vector_strides = [](const array &a, const char *which) { + const int per_word = 16 / static_cast(a.itemsize()); + if (a.strides(-1) != 1) { + throw std::runtime_error( + std::string("[mtplx_native_qsa] ") + which + + " must be contiguous in its last dimension."); + } + for (int d = 1; d + 1 < static_cast(a.ndim()); ++d) { + if (a.strides(d) % per_word != 0) { + throw std::runtime_error( + std::string("[mtplx_native_qsa] ") + which + + " strides must be whole 128-bit words for the vector loads."); + } + } + }; + require_vector_strides(q, "queries"); + require_vector_strides(k, "keys"); + require_vector_strides(v, "values"); + + const int rows = q.shape(2); + int n_tiles = 0; + int tiles_per_split = 0; + int n_splits = 0; + qsa_sparse_gqa_decode_split_geometry(kSelectedTokens, key_tile_, + key_splits_, &n_tiles, + &tiles_per_split, &n_splits); + + // Partial online-softmax states: [n_splits, q_heads, rows, head_dim + 2]. + // A temporary, so MLX frees it when the command buffer completes; both + // dispatches live in ONE encoder, so there is no host sync between them + // and Metal's in-encoder ordering is the barrier. + Shape partial_shape = {n_splits, kQHeads, rows, kPartialLd}; + array partial(partial_shape, float32, nullptr, std::vector{}); + partial.set_data(allocator::malloc(partial.nbytes())); + + out.set_data(allocator::malloc(out.nbytes())); + + auto library = + device.get_library("mtplx_native_qsa_ext", current_binary_dir_decode()); + auto &encoder = metal::get_command_encoder(stream); + encoder.add_temporary(partial); + + MtplxQsaSparseGqaDecodeParams split_params{ + /* q_heads */ kQHeads, + /* kv_heads */ kKvHeads, + /* qL */ rows, + /* kL */ key_length_, + /* topk */ selected.shape(3), + /* gqa_factor */ kGqa, + /* n_tiles */ n_tiles, + /* tiles_per_split */ tiles_per_split, + /* n_splits */ n_splits, + /* partial_ld */ kPartialLd, + /* scale */ scale_, + /* _pad */ 0, + /* Q_strides */ {q.strides(0), q.strides(1), q.strides(2)}, + /* K_strides */ {k.strides(0), k.strides(1), k.strides(2)}, + /* V_strides */ {v.strides(0), v.strides(1), v.strides(2)}, + /* Topk_strides */ + {selected.strides(0), selected.strides(1), selected.strides(2)}}; + + std::string split_name; + concatenate(split_name, "mtplx_qsa_sparse_gqa_decode_split_", + type_to_name(q), "_", type_to_name(selected), "_bk", key_tile_, + "_dc", dimension_tile_, "_gqa", kGqa, "_hp", kHeadPad, "_d", + kHeadDim, "_wm", kWarps); + auto split_kernel = device.get_kernel(split_name, library); + encoder.set_compute_pipeline_state(split_kernel); + encoder.set_input_array(q, 0); + encoder.set_input_array(k, 1); + encoder.set_input_array(v, 2); + encoder.set_input_array(selected, 3); + encoder.set_input_array(offset, 4); + encoder.set_output_array(partial, 5); + encoder.set_bytes(split_params, 6); + encoder.dispatch_threadgroups(MTL::Size(rows, kKvHeads, n_splits), + MTL::Size(32, kWarps, 1)); + + MtplxQsaSparseGqaMergeParams merge_params{ + /* q_heads */ kQHeads, + /* qL */ rows, + /* head_dim */ kHeadDim, + /* n_splits */ n_splits, + /* partial_ld */ kPartialLd, + /* _pad */ 0, + /* O_strides */ {out.strides(0), out.strides(1), out.strides(2)}}; + + std::string merge_name; + concatenate(merge_name, "mtplx_qsa_sparse_gqa_decode_merge_", + type_to_name(q), "_d", kHeadDim); + auto merge_kernel = device.get_kernel(merge_name, library); + encoder.set_compute_pipeline_state(merge_kernel); + encoder.set_input_array(partial, 0); + encoder.set_output_array(out, 1); + encoder.set_bytes(merge_params, 2); + encoder.dispatch_threadgroups(MTL::Size(kQHeads * rows, 1, 1), + MTL::Size(kHeadDim, 1, 1)); + } + + DEFINE_NAME(MtplxQsaSparseGqaDecode) + DEFINE_INPUT_OUTPUT_SHAPE() + + bool is_equivalent(const Primitive &other) const override { + const auto &rhs = static_cast(other); + return scale_ == rhs.scale_ && key_length_ == rhs.key_length_ && + key_tile_ == rhs.key_tile_ && + dimension_tile_ == rhs.dimension_tile_ && + key_splits_ == rhs.key_splits_; + } + + auto state() const { + return std::make_tuple(scale_, key_length_, key_tile_, dimension_tile_, + key_splits_); + } + +private: + float scale_; + int key_length_; + int key_tile_; + int dimension_tile_; + int key_splits_; +}; + +} // namespace + +void qsa_sparse_gqa_decode_split_geometry(int selected_tokens, int key_tile, + int key_splits, int *n_tiles, + int *tiles_per_split, int *n_splits) { + const int tiles = (selected_tokens + key_tile - 1) / key_tile; + int splits = std::min(key_splits, tiles); + if (splits < 1) { + splits = 1; + } + const int per_split = (tiles + splits - 1) / splits; + // Round the split count back down so the LAST split is never empty: with + // 17 tiles and 8 requested splits, per_split is 3 and six splits cover the + // work, so dispatching eight would leave two threadgroups writing nothing. + const int exact_splits = (tiles + per_split - 1) / per_split; + *n_tiles = tiles; + *tiles_per_split = per_split; + *n_splits = exact_splits; +} + +std::string qsa_sparse_gqa_decode_unsupported_reason( + const mx::array &queries, const mx::array &keys, const mx::array &values, + const mx::array &selected_blocks, const mx::array &query_offset, + float scale, int key_length, int key_tile, int dimension_tile, + int key_splits, mx::StreamOrDevice s) { + auto stream = to_stream(s); + const int resolved = + key_length < 0 ? (keys.ndim() == 4 ? keys.shape(2) : -1) : key_length; + return unsupported_reason(queries, keys, values, selected_blocks, + query_offset, scale, resolved, key_tile, + dimension_tile, key_splits, stream); +} + +mx::array qsa_sparse_gqa_decode(const mx::array &queries, const mx::array &keys, + const mx::array &values, + const mx::array &selected_blocks, + const mx::array &query_offset, float scale, + int key_length, int key_tile, + int dimension_tile, int key_splits, + mx::StreamOrDevice s) { + auto stream = to_stream(s); + const int resolved = + key_length < 0 ? (keys.ndim() == 4 ? keys.shape(2) : -1) : key_length; + auto reason = + unsupported_reason(queries, keys, values, selected_blocks, query_offset, + scale, resolved, key_tile, dimension_tile, key_splits, + stream); + if (!reason.empty()) { + throw std::invalid_argument( + "[mtplx_native_qsa.qsa_sparse_gqa_decode] " + reason + "."); + } + + Shape out_shape = queries.shape(); + return array(std::move(out_shape), queries.dtype(), + std::make_shared( + stream, scale, resolved, key_tile, dimension_tile, + key_splits), + std::vector{queries, keys, values, selected_blocks, + query_offset}); +} + +} // namespace mtplx_native diff --git a/native_extensions/qsa_sparse_gqa/sparse_gqa/qsa_sparse_gqa_decode.h b/native_extensions/qsa_sparse_gqa/sparse_gqa/qsa_sparse_gqa_decode.h new file mode 100644 index 000000000..f85502d04 --- /dev/null +++ b/native_extensions/qsa_sparse_gqa/sparse_gqa/qsa_sparse_gqa_decode.h @@ -0,0 +1,52 @@ +// SPDX-License-Identifier: Apache-2.0 +// +// Host entry point for the SPLIT-K decode variant of MTPLX's direct-index +// sparse-GQA attention (see steel_qsa_sparse_gqa_decode.h for provenance, +// the algorithm, and why decode needs a KV split at all). + +#pragma once + +#include + +#include "mlx/array.h" +#include "mlx/stream.h" +#include "mlx/utils.h" + +namespace mx = mlx::core; + +namespace mtplx_native { + +/// Empty when the call is on the supported contract; otherwise a precise, +/// caller-facing reason. Exposed so the Python lane can gate without +/// catching an exception. +std::string qsa_sparse_gqa_decode_unsupported_reason( + const mx::array &queries, const mx::array &keys, const mx::array &values, + const mx::array &selected_blocks, const mx::array &query_offset, + float scale, int key_length, int key_tile, int dimension_tile, + int key_splits, mx::StreamOrDevice s = {}); + +/// queries [1, 24, M, 256] fp16/bf16, last dim contiguous +/// keys/values [1, 2, cap, 256] the FULL cache backing, same dtype +/// selected [1, 1, M, 512] uint32 or int32, the argpartition output +/// in ITS OWN order (never re-sorted) +/// query_offset [1] int32 absolute position of query row 0; a device +/// buffer so a tensor-valued cache offset +/// never has to be read on the host +/// key_length logical tokens in the cache (<= cap); -1 means keys.shape(2) +/// key_splits target number of KV splits; clamped to the tile count and +/// then rounded so no split is empty +/// returns [1, 24, M, 256] same dtype as queries +mx::array qsa_sparse_gqa_decode( + const mx::array &queries, const mx::array &keys, const mx::array &values, + const mx::array &selected_blocks, const mx::array &query_offset, + float scale, int key_length = -1, int key_tile = 128, + int dimension_tile = 32, int key_splits = 8, mx::StreamOrDevice s = {}); + +/// The host half of the split geometry, exported so the Python lane, the +/// tests and the harness can size the partial buffer without duplicating the +/// arithmetic. Writes ``n_tiles``, ``tiles_per_split`` and ``n_splits``. +void qsa_sparse_gqa_decode_split_geometry(int selected_tokens, int key_tile, + int key_splits, int *n_tiles, + int *tiles_per_split, int *n_splits); + +} // namespace mtplx_native diff --git a/native_extensions/qsa_sparse_gqa/sparse_gqa/qsa_sparse_gqa_decode_params.h b/native_extensions/qsa_sparse_gqa/sparse_gqa/qsa_sparse_gqa_decode_params.h new file mode 100644 index 000000000..1629e9738 --- /dev/null +++ b/native_extensions/qsa_sparse_gqa/sparse_gqa/qsa_sparse_gqa_decode_params.h @@ -0,0 +1,71 @@ +// SPDX-License-Identifier: Apache-2.0 +// +// Shared host/device parameter blocks for the SPLIT-K (KV-split) variant of +// the QSA direct-index sparse-GQA kernel -- the decode geometries, M=4 +// (fixed-M4 verify) and M=1 (single-row draft/decode). +// +// Two blocks, because the lane is two dispatches: the split pass writes one +// unnormalised online-softmax state per (query row, KV head, KV split), and +// the merge pass reduces those states into the attention output. +// +// Both are included by the C++ encoder AND by the MSL kernel, so the layout +// can never drift between the two sides. A silent drift here is wrong +// attention, not a crash, which is why the static_asserts below are compiled +// by both. + +#pragma once + +#ifndef __METAL_VERSION__ +#include +#endif + +/// Pass 1: one threadgroup per (query row, KV head, KV split). +struct MtplxQsaSparseGqaDecodeParams { + int q_heads; ///< 24 + int kv_heads; ///< 2 + int qL; ///< query rows in this call (M): 4 for verify, 1 for draft + int kL; ///< logical key length (NOT the cache backing capacity) + int topk; ///< selected blocks per row (512) + int gqa_factor; ///< 12 + + /// Split geometry, all host-computed from ``topk``, ``compress_ratio`` and + /// the compiled-in ``BK``. ``n_splits * tiles_per_split >= n_tiles`` and + /// ``(n_splits - 1) * tiles_per_split < n_tiles``: no split is ever empty, + /// so every threadgroup in the grid does real work. + int n_tiles; ///< ceil((topk*ratio + ratio-1) / BK) + int tiles_per_split; ///< >= 1 + int n_splits; ///< ceil(n_tiles / tiles_per_split), == grid.z + int partial_ld; ///< head_dim + 2: the row carries [O(256) | m | l] + + float scale; ///< 1/sqrt(head_dim); the kernel folds M_LOG2E in itself + int _pad; ///< explicit: keeps the int64 block 8-byte aligned on both sides + + int64_t Q_strides[3]; ///< (B, H, L); last dim is unit stride + int64_t K_strides[3]; ///< (B, H_kv, L) into the FULL cache backing + int64_t V_strides[3]; ///< (B, H_kv, L) into the FULL cache backing + int64_t Topk_strides[3]; ///< (B, 1, M); last dim is unit stride +}; + +// 11 ints (``_pad`` included) + 1 float = 48, already 8-byte aligned, then +// 4 x 3 x int64 = 96. Pin it: host and device MUST agree byte for byte. +static_assert(sizeof(MtplxQsaSparseGqaDecodeParams) == 144, + "MtplxQsaSparseGqaDecodeParams layout drifted host vs device"); +static_assert(alignof(MtplxQsaSparseGqaDecodeParams) == 8, + "MtplxQsaSparseGqaDecodeParams alignment drifted"); + +/// Pass 2: one threadgroup per (query head, query row); head_dim threads. +struct MtplxQsaSparseGqaMergeParams { + int q_heads; ///< 24 + int qL; ///< query rows (M) + int head_dim; ///< 256 + int n_splits; ///< how many partial states to merge + int partial_ld; ///< head_dim + 2 + int _pad; ///< explicit alignment + + int64_t O_strides[3]; ///< (B, H, L) of the [1, 24, M, 256] output +}; + +static_assert(sizeof(MtplxQsaSparseGqaMergeParams) == 48, + "MtplxQsaSparseGqaMergeParams layout drifted host vs device"); +static_assert(alignof(MtplxQsaSparseGqaMergeParams) == 8, + "MtplxQsaSparseGqaMergeParams alignment drifted"); diff --git a/native_extensions/qsa_sparse_gqa/sparse_gqa/qsa_sparse_gqa_params.h b/native_extensions/qsa_sparse_gqa/sparse_gqa/qsa_sparse_gqa_params.h new file mode 100644 index 000000000..1b9c679c7 --- /dev/null +++ b/native_extensions/qsa_sparse_gqa/sparse_gqa/qsa_sparse_gqa_params.h @@ -0,0 +1,42 @@ +// SPDX-License-Identifier: Apache-2.0 +// +// Shared host/device parameter block for the QSA direct-index sparse-GQA +// kernel. Both qsa_sparse_gqa.cpp (C++) and qsa_sparse_gqa.metal (MSL) +// include this file, so the layout can never drift between the encoder and +// the kernel -- the failure mode oMLX's duplicated struct leaves open. + +#pragma once + +#ifndef __METAL_VERSION__ +#include +#endif + +struct MtplxQsaSparseGqaParams { + int B; ///< batch (always 1 on the supported geometry) + int q_heads; ///< 24 + int kv_heads; ///< 2 + int qL; ///< query rows in this call (M) + int kL; ///< logical key length (NOT the cache backing capacity) + int topk; ///< selected blocks per row (512) + int gqa_factor; ///< 12 + int q_offset; ///< absolute position of query row 0 + + float scale; ///< 1/sqrt(head_dim), pre-M_LOG2E in the kernel + int _pad; ///< explicit: keeps the int64 block 8-byte aligned on both sides + + int64_t Q_strides[3]; ///< (B, H, L); last dim is unit stride + int64_t K_strides[3]; ///< (B, H_kv, L) into the FULL cache backing + int64_t V_strides[3]; ///< (B, H_kv, L) into the FULL cache backing + int64_t Topk_strides[3]; ///< (B, 1, M); last dim is unit stride + int64_t O_strides[3]; ///< (B, H, L) +}; + +// The encoder writes this struct with set_bytes and the kernel reads it as a +// `constant` pointer, so host and device MUST agree on the layout byte for +// byte. A mismatch is silent wrong attention, not a crash, so pin it here: +// this header is compiled by BOTH sides, and the assert fires on whichever +// one drifts. 8 ints + float + pad = 40 B, then 5 x 3 x int64 = 120 B. +static_assert(sizeof(MtplxQsaSparseGqaParams) == 160, + "MtplxQsaSparseGqaParams layout drifted between host and device"); +static_assert(alignof(MtplxQsaSparseGqaParams) == 8, + "MtplxQsaSparseGqaParams alignment drifted"); diff --git a/native_extensions/qsa_sparse_gqa/sparse_gqa/steel_qsa_sparse_gqa.h b/native_extensions/qsa_sparse_gqa/sparse_gqa/steel_qsa_sparse_gqa.h new file mode 100644 index 000000000..c91ab2fc4 --- /dev/null +++ b/native_extensions/qsa_sparse_gqa/sparse_gqa/steel_qsa_sparse_gqa.h @@ -0,0 +1,339 @@ +// SPDX-License-Identifier: Apache-2.0 +// +// Direct-index sparse GQA attention for Qwen3.8 Flash-Next QSA. +// +// Ported into MTPLX from oMLX (jundot/omlx, commit 7467dce8, +// `omlx/custom_kernels/glm_moe_dsa/csrc/kernels/steel_qwen4_qsa_sparse_gqa.h`), +// written by Jonathan Spangler, Apache-2.0, Copyright (c) 2026 OpenAI. +// The 128-bit global K/V staging pattern is in turn adapted from mlx-serve's +// MIT `msv_attn_p256` kernel (Copyright 2026 David Dalcu); the query-specific +// direct-index/GQA organisation and the Steel MMA implementation are oMLX's. +// +// MTPLX changes, all mechanical: +// * renamed to the mtplx_* symbol space so the metallib can co-exist with +// an oMLX build in one process; +// * the host/device parameter block moved into a single shared header +// (qsa_sparse_gqa_params.h) instead of being declared twice; +// * `IndexT` is instantiated for int32 as well as uint32, because MTPLX's +// `_select_eager` emits int32 block ids and an astype would copy 8 MB +// per layer per 4,096-row chunk; +// * `params->kL` is the LOGICAL key length supplied by the caller, not the +// K array's extent: MTPLX hands the kernel the full KV cache backing +// [1, 2, capacity, 256] and attends to the first `total_tokens` rows. +// +// Algorithm (unchanged): one threadgroup owns one (query row, KV head). The +// twelve query heads sharing that KV head are padded to a single 16-row Steel +// MMA tile, so both GQA groups reuse each randomly addressed K/V tile without +// materialising a [query, selected, kv-head, dim] gathered tensor. The +// selected block list is chronological; the zero-to-three-token causal tail is +// generated in-kernel; invalid early-prefix slots are masked to true -inf +// before an fp32 online softmax (exp2). No score tensor, no gathered K/V, no +// bool mask. + +#pragma once + +#include "mlx/backend/metal/kernels/steel/attn/attn.h" +#include "mlx/backend/metal/kernels/steel/attn/params.h" + +#include "sparse_gqa/qsa_sparse_gqa_params.h" + +using namespace mlx::steel; + +struct MtplxQsaMaxOp { + template METAL_FUNC static constexpr T apply(T x, T y) { + return metal::max(x, y); + } +}; + +struct MtplxQsaSumOp { + template METAL_FUNC static constexpr T apply(T x, T y) { + return x + y; + } +}; + +struct MtplxQsaMulOp { + template METAL_FUNC static constexpr T apply(T x, T y) { + return x * y; + } +}; + +struct MtplxQsaExpSubOp { + template METAL_FUNC static constexpr T apply(T x, T y) { + return fast::exp2(x - y); + } +}; + +struct MtplxQsaDivOp { + template METAL_FUNC static constexpr T apply(T x, T y) { + return x / y; + } +}; + +// clang-format off +template < + typename T, + int BK, + int DC, + int GQA, + int H_PAD, + int D, + int WM, + typename IndexT, + typename AccumType = float> +[[kernel, max_total_threads_per_threadgroup(WM * 32)]] void +mtplx_qsa_sparse_gqa_attention( + const device T* Q [[buffer(0)]], + const device T* K [[buffer(1)]], + const device T* V [[buffer(2)]], + const device IndexT* Topk [[buffer(3)]], + device T* O [[buffer(4)]], + const constant MtplxQsaSparseGqaParams* params [[buffer(5)]], + uint simd_lane_id [[thread_index_in_simdgroup]], + uint simd_group_id [[simdgroup_index_in_threadgroup]], + uint3 tid [[threadgroup_position_in_grid]]) { // clang-format on + + constexpr short kFragSize = 8; + constexpr short padQ = 16 / sizeof(T); + constexpr short padK = 16 / sizeof(T); + constexpr short padV = 16 / sizeof(T); + + constexpr short LDQ = DC + padQ; + constexpr short LDK = BK + padK; + constexpr short LDV = DC + padV; + + constexpr int kNWarps = WM; + constexpr int TQ = H_PAD / (kNWarps * kFragSize); + constexpr int TK = BK / kFragSize; + constexpr int TDC = DC / kFragSize; + constexpr int D_CHUNKS = D / DC; + + static_assert(GQA <= H_PAD, "Qwen GQA heads must fit the padded MMA tile."); + static_assert(TQ == 1, "Qwen sparse GQA expects one query-head tile."); + static_assert(H_PAD % (kNWarps * kFragSize) == 0, + "Padded query heads must divide evenly across simdgroups."); + static_assert(BK % kFragSize == 0, "BK must be a multiple of eight."); + static_assert(DC % kFragSize == 0, "DC must be a multiple of eight."); + static_assert(D % DC == 0, "Head dimension must divide DC."); + + constexpr int tgp_size = WM * 32; + const int lane = int(simd_group_id * 32 + simd_lane_id); + const int q_pos = int(tid.x); + const int kv_head = int(tid.y); + const int b = int(tid.z); + + threadgroup T Qs[H_PAD * LDQ]; + threadgroup T KVs[(BK * LDV > DC * LDK) ? BK * LDV : DC * LDK]; + threadgroup int selected[BK]; + + using MMAFragAcc = BaseMMAFrag; + MMATile Qtile; + MMATile Ktile; + MMATile Stile; + MMATile Vtile; + MMATile Otile; + Otile.clear(); + + const short2 simd_coord = MMAFragAcc::get_coord(simd_lane_id); + const short sm = simd_coord.y; + const short sn = simd_coord.x; + const short tm = kFragSize * TQ * simd_group_id; + const short Qs_offset = (tm + sm) * LDQ + sn; + const short Ks_offset = sm * LDK + sn; + const short Vs_offset = sm * LDV + sn; + + const AccumType scale = AccumType(params->scale * M_LOG2E_F); + constexpr short rows_per_thread = decltype(Stile)::kRowsPerThread; + AccumType max_score[rows_per_thread]; + AccumType sum_score[rows_per_thread] = {0}; + STEEL_PRAGMA_UNROLL + for (short i = 0; i < rows_per_thread; ++i) { + max_score[i] = Limits::finite_min; + } + + const int query_head_base = kv_head * GQA; + const device T *q_base = Q + size_t(b) * params->Q_strides[0] + + size_t(query_head_base) * params->Q_strides[1] + + size_t(q_pos) * params->Q_strides[2]; + const device T *k_base = K + size_t(b) * params->K_strides[0] + + size_t(kv_head) * params->K_strides[1]; + const device T *v_base = V + size_t(b) * params->V_strides[0] + + size_t(kv_head) * params->V_strides[1]; + const device IndexT *topk_base = Topk + size_t(b) * params->Topk_strides[0] + + size_t(q_pos) * params->Topk_strides[2]; + + const int q_abs = params->q_offset + q_pos; + constexpr int kCompressRatio = 4; + constexpr int kTail = kCompressRatio - 1; + const int selected_tokens = params->topk * kCompressRatio + kTail; + const int complete_blocks = (q_abs + 1) / kCompressRatio; + const int valid_blocks = metal::min(params->topk, complete_blocks); + const int n_tiles = (selected_tokens + BK - 1) / BK; + + for (int ktile = 0; ktile < n_tiles; ++ktile) { + const int topk_off = ktile * BK; + for (int k = lane; k < BK; k += tgp_size) { + const int slot = topk_off + k; + int k_pos = -1; + if (slot < params->topk * kCompressRatio) { + const int block_slot = slot / kCompressRatio; + if (block_slot < valid_blocks) { + const IndexT raw_block = topk_base[block_slot]; + const ulong candidate = ulong(raw_block) * ulong(kCompressRatio) + + ulong(slot % kCompressRatio); + if (candidate < ulong(params->kL) && candidate <= ulong(q_abs)) { + k_pos = int(candidate); + } + } + } else if (slot < selected_tokens) { + const int tail_offset = slot - params->topk * kCompressRatio; + const int candidate = complete_blocks * kCompressRatio + tail_offset; + if (candidate < params->kL && candidate <= q_abs) { + k_pos = candidate; + } + } + selected[k] = k_pos; + } + threadgroup_barrier(mem_flags::mem_threadgroup); + + Stile.clear(); + STEEL_PRAGMA_UNROLL + for (short dchunk = 0; dchunk < D_CHUNKS; ++dchunk) { + const int dbase = int(dchunk) * DC; + for (int elem = lane; elem < H_PAD * (DC / 8); elem += tgp_size) { + const int h = elem / (DC / 8); + const int d8 = elem - h * (DC / 8); + uint4 word = uint4(0); + if (h < GQA) { + word = *((const device uint4 *)(q_base + + size_t(h) * params->Q_strides[1] + + dbase) + + d8); + } + *((threadgroup uint4 *)(Qs + h * LDQ) + d8) = word; + } + for (int elem = lane; elem < BK * (DC / 8); elem += tgp_size) { + const int k = elem / (DC / 8); + const int d8 = elem - k * (DC / 8); + const int k_pos = selected[k]; + uint4 word = uint4(0); + if (k_pos >= 0) { + word = *((const device uint4 *)(k_base + + size_t(k_pos) * params->K_strides[2] + + dbase) + + d8); + } + thread T *values = (thread T *)&word; + const int d = d8 * 8; + STEEL_PRAGMA_UNROLL + for (short e = 0; e < 8; ++e) { + KVs[k + (d + e) * LDK] = values[e]; + } + } + threadgroup_barrier(mem_flags::mem_threadgroup); + STEEL_PRAGMA_UNROLL + for (short dd = 0; dd < TDC; ++dd) { + simdgroup_barrier(mem_flags::mem_none); + Qtile.template load(&Qs[Qs_offset + dd * kFragSize]); + Ktile.template load( + &KVs[Ks_offset + dd * kFragSize * LDK]); + simdgroup_barrier(mem_flags::mem_none); + tile_matmad(Stile, Qtile, Ktile, Stile); + } + threadgroup_barrier(mem_flags::mem_threadgroup); + } + + STEEL_PRAGMA_UNROLL + for (short i = 0; i < decltype(Stile)::kElemsPerTile; ++i) { + Stile.elems()[i] *= scale; + } + { + using stile_t = decltype(Stile); + using selem_t = typename stile_t::elem_type; + // Early first-chunk rows can have no complete block, so entire leading + // tiles are invalid before the causal tail in the final tile. True + // -INFINITY makes those tiles contribute exp2(-inf - finite_max) == 0; + // finite_min would incorrectly add one to the softmax denominator. + constexpr auto neg_inf = selem_t(-INFINITY); + STEEL_PRAGMA_UNROLL + for (short i = 0; i < stile_t::kTileRows; ++i) { + STEEL_PRAGMA_UNROLL + for (short j = 0; j < stile_t::kTileCols; ++j) { + const short col_pos = sn + j * stile_t::kFragCols; + STEEL_PRAGMA_UNROLL + for (short e = 0; e < stile_t::MMAFrag_t::kElemCols; ++e) { + if (selected[col_pos + e] < 0) { + Stile.frag_at(i, j)[e] = neg_inf; + } + } + } + } + } + + AccumType new_max[rows_per_thread]; + AccumType factor[rows_per_thread]; + STEEL_PRAGMA_UNROLL + for (short i = 0; i < rows_per_thread; ++i) { + new_max[i] = max_score[i]; + } + Stile.template row_reduce(new_max); + Stile.template row_bin_op(new_max); + STEEL_PRAGMA_UNROLL + for (short i = 0; i < rows_per_thread; ++i) { + factor[i] = fast::exp2(max_score[i] - new_max[i]); + max_score[i] = new_max[i]; + } + AccumType sum_score_tmp[rows_per_thread] = {0}; + Stile.template row_reduce(sum_score_tmp); + STEEL_PRAGMA_UNROLL + for (short i = 0; i < rows_per_thread; ++i) { + sum_score[i] = sum_score[i] * factor[i] + sum_score_tmp[i]; + } + Otile.template row_bin_op(factor); + + STEEL_PRAGMA_UNROLL + for (short vchunk = 0; vchunk < D_CHUNKS; ++vchunk) { + const int dbase = int(vchunk) * DC; + for (int elem = lane; elem < BK * (DC / 8); elem += tgp_size) { + const int k = elem / (DC / 8); + const int d8 = elem - k * (DC / 8); + const int k_pos = selected[k]; + uint4 word = uint4(0); + if (k_pos >= 0) { + word = *((const device uint4 *)(v_base + + size_t(k_pos) * params->V_strides[2] + + dbase) + + d8); + } + *((threadgroup uint4 *)(KVs + k * LDV) + d8) = word; + } + threadgroup_barrier(mem_flags::mem_threadgroup); + STEEL_PRAGMA_UNROLL + for (short iq = 0; iq < TQ; ++iq) { + STEEL_PRAGMA_UNROLL + for (short id = 0; id < TDC; ++id) { + STEEL_PRAGMA_UNROLL + for (short ik = 0; ik < TK; ++ik) { + const short kk = ik * kFragSize; + const short dd = id * kFragSize; + Vtile.template load( + &KVs[Vs_offset + kk * LDV + dd]); + MMAFragAcc::mma(Otile.frag_at(iq, vchunk * TDC + id), + Stile.frag_at(iq, ik), Vtile.frag_at(0, 0), + Otile.frag_at(iq, vchunk * TDC + id)); + } + } + } + threadgroup_barrier(mem_flags::mem_threadgroup); + } + } + + Otile.template row_bin_op(sum_score); + device T *out = O + size_t(b) * params->O_strides[0] + + size_t(query_head_base + tm + sm) * params->O_strides[1] + + size_t(q_pos) * params->O_strides[2] + sn; + const short rows_left = short(GQA - (tm + sm)); + if (rows_left > 0) { + Otile.template store_safe(out, params->O_strides[1], + short2(D - sn, rows_left)); + } +} diff --git a/native_extensions/qsa_sparse_gqa/sparse_gqa/steel_qsa_sparse_gqa_decode.h b/native_extensions/qsa_sparse_gqa/sparse_gqa/steel_qsa_sparse_gqa_decode.h new file mode 100644 index 000000000..add4e0a49 --- /dev/null +++ b/native_extensions/qsa_sparse_gqa/sparse_gqa/steel_qsa_sparse_gqa_decode.h @@ -0,0 +1,431 @@ +// SPDX-License-Identifier: Apache-2.0 +// +// SPLIT-K (KV-split) direct-index sparse GQA attention for Qwen3.8 +// Flash-Next QSA DECODE -- M=4 (fixed-M4 verify) and M=1 (single-row draft). +// +// Derived from MTPLX's phase-1 port of oMLX's single-pass kernel +// (sparse_gqa/steel_qsa_sparse_gqa.h; jundot/omlx 7467dce8, Jonathan +// Spangler, Apache-2.0; the 128-bit global K/V staging is in turn adapted +// from mlx-serve's MIT `msv_attn_p256`, Copyright 2026 David Dalcu). The +// per-tile body below -- the direct-index K/V staging, the Steel MMA +// score/PV pair, the fp32 online softmax -- is that kernel's, unchanged. +// +// WHY A SPLIT-K VARIANT EXISTS AT ALL +// ----------------------------------- +// The single-pass kernel parallelises over QUERY ROWS: its grid is +// ``(qL, kv_heads, 1)`` threadgroups of 64 threads. At prefill (4,096 rows) +// that is 8,192 threadgroups. At M=4 it is EIGHT threadgroups on a 40-core +// M5 Max, and at M=1 it is TWO -- and each of those few threadgroups still +// walks all 2,051 selected keys x 256 dims. Phase 1's own design note +// (docs/perf/qsa-sparse-gqa-phase2-wiring.md, section 4) called this out and +// priced the fix as its own item: split the selected keys across several +// threadgroups per (row, KV head) and combine the partial online-softmax +// states in a second pass. That is this file. +// +// It is also why the decode lane MUST be split-K rather than single-pass: +// MTPLX's own history has a hand-written mx.fast.metal_kernel SDPA losing to +// stock at long N precisely because MLX's production SDPA switches to a +// KV-split two-pass path there. A single-pass sparse kernel would repeat +// that mistake with a shorter (2,051-key) but equally serial walk. +// +// WHAT CHANGES FROM THE PREFILL KERNEL, AND WHY +// --------------------------------------------- +// 1. ``tid.z`` is the KV SPLIT, not the batch. The contract already +// refuses ``B != 1``, so the batch axis was carrying no information. +// 2. The per-tile loop runs over ``[t0, t1)`` instead of ``[0, n_tiles)``. +// 3. The final ``Otile / sum`` divide MOVES to the merge pass; pass 1 +// writes the UNNORMALISED accumulator plus the running ``(m, l)`` pair. +// 4. The query offset arrives as a one-element int32 DEVICE buffer rather +// than a host scalar in the params block, because the fixed-M4 verify +// may carry a tensor-valued cache offset and reading it on the host +// would synchronize the graph. +// 5. Validity is decided PER SLOT (``block_id < complete_blocks``), not by +// a leading-prefix cut. This is the load-bearing difference and it is +// not cosmetic: +// +// The prefill branch of ``_select_eager`` SORTS its top-k ascending, so +// there the valid entries really are a leading prefix and the prefix cut +// is exact. The DECODE branch does not sort: it hands +// ``mx.argpartition``'s raw output straight to the rows-gather token +// list, whose predicate is ``block < visible_blocks`` evaluated on +// EVERY slot. Applying the prefix cut to an unsorted row +// would drop visible blocks and admit invisible ones. Below the +// dense/sparse crossover the argpartition genuinely returns +// -inf-scored ids (there are fewer than 512 complete blocks to choose +// from), so this is a real case, not a hypothetical one. +// +// Note also what is NOT used as the predicate: ``candidate <= q_abs``. +// That is implied by ``block < complete_blocks`` but is NOT equivalent +// to it -- for ``block == complete_blocks`` it would admit the tokens of +// the incomplete block, which the tail slots already cover, and +// double-count them in the softmax denominator. +// +// The visible set is therefore IDENTICAL to the shipped rows-gather lane's, +// slot for slot. The arithmetic is not: fp32 online softmax (exp2) and an +// fp32 P@V against the shipped path's fp32 softmax, bf16 probability cast +// and bf16 P@V, plus the split-K rescale. That is a ROUNDING-CLASS change +// to attention output on the same terms as kernels/qwen4_m4_hyper_read.py -- +// adopt it on greedy-token agreement and a HumanEval gate, never on a digest. + +#pragma once + +#include "mlx/backend/metal/kernels/steel/attn/attn.h" +#include "mlx/backend/metal/kernels/steel/attn/params.h" + +#include "sparse_gqa/qsa_sparse_gqa_decode_params.h" + +using namespace mlx::steel; + +struct MtplxQsaDecMaxOp { + template METAL_FUNC static constexpr T apply(T x, T y) { + return metal::max(x, y); + } +}; + +struct MtplxQsaDecSumOp { + template METAL_FUNC static constexpr T apply(T x, T y) { + return x + y; + } +}; + +struct MtplxQsaDecMulOp { + template METAL_FUNC static constexpr T apply(T x, T y) { + return x * y; + } +}; + +struct MtplxQsaDecExpSubOp { + template METAL_FUNC static constexpr T apply(T x, T y) { + return metal::fast::exp2(x - y); + } +}; + +// clang-format off +template < + typename T, + int BK, + int DC, + int GQA, + int H_PAD, + int D, + int WM, + typename IndexT, + typename AccumType = float> +[[kernel, max_total_threads_per_threadgroup(WM * 32)]] void +mtplx_qsa_sparse_gqa_decode_split( + const device T* Q [[buffer(0)]], + const device T* K [[buffer(1)]], + const device T* V [[buffer(2)]], + const device IndexT* Topk [[buffer(3)]], + const device int* QOffset [[buffer(4)]], + device AccumType* Partial [[buffer(5)]], + const constant MtplxQsaSparseGqaDecodeParams* params [[buffer(6)]], + uint simd_lane_id [[thread_index_in_simdgroup]], + uint simd_group_id [[simdgroup_index_in_threadgroup]], + uint3 tid [[threadgroup_position_in_grid]]) { // clang-format on + + constexpr short kFragSize = 8; + constexpr short padQ = 16 / sizeof(T); + constexpr short padK = 16 / sizeof(T); + constexpr short padV = 16 / sizeof(T); + + constexpr short LDQ = DC + padQ; + constexpr short LDK = BK + padK; + constexpr short LDV = DC + padV; + + constexpr int kNWarps = WM; + constexpr int TQ = H_PAD / (kNWarps * kFragSize); + constexpr int TK = BK / kFragSize; + constexpr int TDC = DC / kFragSize; + constexpr int D_CHUNKS = D / DC; + + static_assert(GQA <= H_PAD, "Qwen GQA heads must fit the padded MMA tile."); + static_assert(TQ == 1, "Qwen sparse GQA expects one query-head tile."); + static_assert(H_PAD % (kNWarps * kFragSize) == 0, + "Padded query heads must divide evenly across simdgroups."); + static_assert(BK % kFragSize == 0, "BK must be a multiple of eight."); + static_assert(DC % kFragSize == 0, "DC must be a multiple of eight."); + static_assert(D % DC == 0, "Head dimension must divide DC."); + + constexpr int tgp_size = WM * 32; + const int lane = int(simd_group_id * 32 + simd_lane_id); + const int q_pos = int(tid.x); + const int kv_head = int(tid.y); + const int split = int(tid.z); + + threadgroup T Qs[H_PAD * LDQ]; + threadgroup T KVs[(BK * LDV > DC * LDK) ? BK * LDV : DC * LDK]; + threadgroup int selected[BK]; + + using MMAFragAcc = BaseMMAFrag; + MMATile Qtile; + MMATile Ktile; + MMATile Stile; + MMATile Vtile; + MMATile Otile; + Otile.clear(); + + const short2 simd_coord = MMAFragAcc::get_coord(simd_lane_id); + const short sm = simd_coord.y; + const short sn = simd_coord.x; + const short tm = kFragSize * TQ * simd_group_id; + const short Qs_offset = (tm + sm) * LDQ + sn; + const short Ks_offset = sm * LDK + sn; + const short Vs_offset = sm * LDV + sn; + + const AccumType scale = AccumType(params->scale * M_LOG2E_F); + constexpr short rows_per_thread = decltype(Stile)::kRowsPerThread; + static_assert(rows_per_thread == 1, + "One query-head tile means one accumulator row per thread; " + "the (m, l) store below assumes it."); + AccumType max_score[rows_per_thread]; + AccumType sum_score[rows_per_thread] = {0}; + STEEL_PRAGMA_UNROLL + for (short i = 0; i < rows_per_thread; ++i) { + // finite_min, NOT -infinity: an all-masked tile must produce + // exp2(-inf - finite_min) == 0 rather than exp2(-inf + inf) == NaN. + max_score[i] = Limits::finite_min; + } + + const int query_head_base = kv_head * GQA; + const device T *q_base = Q + size_t(query_head_base) * params->Q_strides[1] + + size_t(q_pos) * params->Q_strides[2]; + const device T *k_base = K + size_t(kv_head) * params->K_strides[1]; + const device T *v_base = V + size_t(kv_head) * params->V_strides[1]; + const device IndexT *topk_base = Topk + size_t(q_pos) * params->Topk_strides[2]; + + const int q_abs = QOffset[0] + q_pos; + constexpr int kCompressRatio = 4; + constexpr int kTail = kCompressRatio - 1; + const int selected_tokens = params->topk * kCompressRatio + kTail; + // Identical to the shipped lane's ``visible_blocks = (qpos + 1) / RATIO``. + const int complete_blocks = (q_abs + 1) / kCompressRatio; + + const int t0 = split * params->tiles_per_split; + const int t1 = metal::min(params->n_tiles, t0 + params->tiles_per_split); + + for (int ktile = t0; ktile < t1; ++ktile) { + const int topk_off = ktile * BK; + for (int k = lane; k < BK; k += tgp_size) { + const int slot = topk_off + k; + int k_pos = -1; + if (slot < params->topk * kCompressRatio) { + const int block_slot = slot / kCompressRatio; + // PER SLOT, not a prefix cut: the decode selector does not sort. + const long raw_block = long(topk_base[block_slot]); + if (raw_block >= 0 && raw_block < long(complete_blocks)) { + const long candidate = + raw_block * long(kCompressRatio) + long(slot % kCompressRatio); + if (candidate >= 0 && candidate < long(params->kL)) { + k_pos = int(candidate); + } + } + } else if (slot < selected_tokens) { + const int tail_offset = slot - params->topk * kCompressRatio; + const long candidate = + long(complete_blocks) * long(kCompressRatio) + long(tail_offset); + if (candidate >= 0 && candidate < long(params->kL) && + candidate <= long(q_abs)) { + k_pos = int(candidate); + } + } + selected[k] = k_pos; + } + threadgroup_barrier(mem_flags::mem_threadgroup); + + Stile.clear(); + STEEL_PRAGMA_UNROLL + for (short dchunk = 0; dchunk < D_CHUNKS; ++dchunk) { + const int dbase = int(dchunk) * DC; + for (int elem = lane; elem < H_PAD * (DC / 8); elem += tgp_size) { + const int h = elem / (DC / 8); + const int d8 = elem - h * (DC / 8); + uint4 word = uint4(0); + if (h < GQA) { + word = *((const device uint4 *)(q_base + + size_t(h) * params->Q_strides[1] + + dbase) + + d8); + } + *((threadgroup uint4 *)(Qs + h * LDQ) + d8) = word; + } + for (int elem = lane; elem < BK * (DC / 8); elem += tgp_size) { + const int k = elem / (DC / 8); + const int d8 = elem - k * (DC / 8); + const int k_pos = selected[k]; + uint4 word = uint4(0); + if (k_pos >= 0) { + word = *((const device uint4 *)(k_base + + size_t(k_pos) * params->K_strides[2] + + dbase) + + d8); + } + thread T *values = (thread T *)&word; + const int d = d8 * 8; + STEEL_PRAGMA_UNROLL + for (short e = 0; e < 8; ++e) { + KVs[k + (d + e) * LDK] = values[e]; + } + } + threadgroup_barrier(mem_flags::mem_threadgroup); + STEEL_PRAGMA_UNROLL + for (short dd = 0; dd < TDC; ++dd) { + simdgroup_barrier(mem_flags::mem_none); + Qtile.template load(&Qs[Qs_offset + dd * kFragSize]); + Ktile.template load( + &KVs[Ks_offset + dd * kFragSize * LDK]); + simdgroup_barrier(mem_flags::mem_none); + tile_matmad(Stile, Qtile, Ktile, Stile); + } + threadgroup_barrier(mem_flags::mem_threadgroup); + } + + STEEL_PRAGMA_UNROLL + for (short i = 0; i < decltype(Stile)::kElemsPerTile; ++i) { + Stile.elems()[i] *= scale; + } + { + using stile_t = decltype(Stile); + using selem_t = typename stile_t::elem_type; + constexpr auto neg_inf = selem_t(-INFINITY); + STEEL_PRAGMA_UNROLL + for (short i = 0; i < stile_t::kTileRows; ++i) { + STEEL_PRAGMA_UNROLL + for (short j = 0; j < stile_t::kTileCols; ++j) { + const short col_pos = sn + j * stile_t::kFragCols; + STEEL_PRAGMA_UNROLL + for (short e = 0; e < stile_t::MMAFrag_t::kElemCols; ++e) { + if (selected[col_pos + e] < 0) { + Stile.frag_at(i, j)[e] = neg_inf; + } + } + } + } + } + + AccumType new_max[rows_per_thread]; + AccumType factor[rows_per_thread]; + STEEL_PRAGMA_UNROLL + for (short i = 0; i < rows_per_thread; ++i) { + new_max[i] = max_score[i]; + } + Stile.template row_reduce(new_max); + Stile.template row_bin_op(new_max); + STEEL_PRAGMA_UNROLL + for (short i = 0; i < rows_per_thread; ++i) { + factor[i] = metal::fast::exp2(max_score[i] - new_max[i]); + max_score[i] = new_max[i]; + } + AccumType sum_score_tmp[rows_per_thread] = {0}; + Stile.template row_reduce(sum_score_tmp); + STEEL_PRAGMA_UNROLL + for (short i = 0; i < rows_per_thread; ++i) { + sum_score[i] = sum_score[i] * factor[i] + sum_score_tmp[i]; + } + Otile.template row_bin_op(factor); + + STEEL_PRAGMA_UNROLL + for (short vchunk = 0; vchunk < D_CHUNKS; ++vchunk) { + const int dbase = int(vchunk) * DC; + for (int elem = lane; elem < BK * (DC / 8); elem += tgp_size) { + const int k = elem / (DC / 8); + const int d8 = elem - k * (DC / 8); + const int k_pos = selected[k]; + uint4 word = uint4(0); + if (k_pos >= 0) { + word = *((const device uint4 *)(v_base + + size_t(k_pos) * params->V_strides[2] + + dbase) + + d8); + } + *((threadgroup uint4 *)(KVs + k * LDV) + d8) = word; + } + threadgroup_barrier(mem_flags::mem_threadgroup); + STEEL_PRAGMA_UNROLL + for (short iq = 0; iq < TQ; ++iq) { + STEEL_PRAGMA_UNROLL + for (short id = 0; id < TDC; ++id) { + STEEL_PRAGMA_UNROLL + for (short ik = 0; ik < TK; ++ik) { + const short kk = ik * kFragSize; + const short dd = id * kFragSize; + Vtile.template load( + &KVs[Vs_offset + kk * LDV + dd]); + MMAFragAcc::mma(Otile.frag_at(iq, vchunk * TDC + id), + Stile.frag_at(iq, ik), Vtile.frag_at(0, 0), + Otile.frag_at(iq, vchunk * TDC + id)); + } + } + } + threadgroup_barrier(mem_flags::mem_threadgroup); + } + } + + // NO divide here: the merge pass owns the normalisation, because the + // denominator is only known once every split has reported. + const short rows_left = short(GQA - (tm + sm)); + if (rows_left > 0) { + const size_t ld = size_t(params->partial_ld); + const size_t split_stride = + size_t(params->q_heads) * size_t(params->qL) * ld; + device AccumType *prow = Partial + size_t(split) * split_stride + + size_t(query_head_base + tm + sm) * + size_t(params->qL) * ld + + size_t(q_pos) * ld; + Otile.template store_safe(prow + sn, params->partial_ld, + short2(D - sn, rows_left)); + // The four lanes that share a row all hold the same reduced (m, l) after + // ``row_reduce``; sn == 0 elects one of them to write. + if (sn == 0) { + prow[D] = max_score[0]; + prow[D + 1] = sum_score[0]; + } + } +} + +// clang-format off +template +[[kernel, max_total_threads_per_threadgroup(D)]] void +mtplx_qsa_sparse_gqa_decode_merge( + const device AccumType* Partial [[buffer(0)]], + device T* O [[buffer(1)]], + const constant MtplxQsaSparseGqaMergeParams* params [[buffer(2)]], + uint3 tid [[threadgroup_position_in_grid]], + uint3 lid [[thread_position_in_threadgroup]]) { // clang-format on + + // One threadgroup per (query head, query row); one thread per head dim. + const int row = int(tid.x); + const int d = int(lid.x); + const size_t ld = size_t(params->partial_ld); + const int n_splits = params->n_splits; + const size_t split_stride = size_t(params->q_heads) * size_t(params->qL) * ld; + const device AccumType *base = Partial + size_t(row) * ld; + + // Pass 1 never writes -infinity into m (it initialises to finite_min), so + // this max is always replaced by a real value and the rescale below can + // never evaluate inf - inf. + AccumType m = Limits::finite_min; + for (int s = 0; s < n_splits; ++s) { + m = metal::max(m, base[size_t(s) * split_stride + D]); + } + + AccumType denom = AccumType(0); + AccumType acc = AccumType(0); + for (int s = 0; s < n_splits; ++s) { + const device AccumType *p = base + size_t(s) * split_stride; + const AccumType alpha = metal::fast::exp2(p[D] - m); + denom += alpha * p[D + 1]; + acc += alpha * p[d]; + } + // denom == 0 is only reachable if EVERY slot of the row was masked, which + // the lane's contract excludes (complete_blocks >= 1 guarantees at least + // one visible block is selected). Emit zero rather than a NaN. + const AccumType inv = + (denom > AccumType(0)) ? (AccumType(1) / denom) : AccumType(0); + + const int head = row / params->qL; + const int q_pos = row - head * params->qL; + O[size_t(head) * size_t(params->O_strides[1]) + + size_t(q_pos) * size_t(params->O_strides[2]) + size_t(d)] = + static_cast(acc * inv); +} diff --git a/scripts/bundle_native_runtime_wheel.py b/scripts/bundle_native_runtime_wheel.py index 7491ba86f..3cdd9eb27 100644 --- a/scripts/bundle_native_runtime_wheel.py +++ b/scripts/bundle_native_runtime_wheel.py @@ -6,6 +6,7 @@ """ import argparse +import contextlib from email.parser import BytesParser from email.policy import compat32 import hashlib @@ -20,6 +21,26 @@ from wheel.wheelfile import WheelFile +# The native extensions the runtime wheel may carry, each in its own platform +# wheel. mtplx_qsa_kernels is the metallib-bearing QSA lane (ea2560a2); +# mtplx_native_qsa is the metallib-bearing split-K QSA sparse-GQA decode +# extension the MTPLX_QSA_SPARSE_DECODE lane needs. Every Mach-O member +# (.so/.dylib) of each is +# Developer-ID + hardened-runtime + secure-timestamp signed before packaging, +# so notarization does not reject an ad-hoc-signed member found inside the zip. +_KNOWN_NATIVE = { + "mtplx_qsa_kernels": { + "required_files": ("NOTICE", "LICENSE.txt", "MLX_LICENSE.txt"), + "require_metallib": True, + }, + "mtplx_native_qsa": { + "required_files": (), + "require_metallib": True, + }, +} +_REQUIRED_NATIVE = "mtplx_qsa_kernels" + + def sign_mach_o(name: str, data: bytes, identity: str) -> bytes: """Return the Mach-O member re-signed the way the app bundle signs its binaries. @@ -54,7 +75,15 @@ def sign_mach_o(name: str, data: bytes, identity: str) -> bytes: def main(): parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("runtime", type=Path) - parser.add_argument("native", type=Path) + parser.add_argument( + "native", + type=Path, + nargs="+", + help="One or more tested native platform wheels " + f"({', '.join(sorted(_KNOWN_NATIVE))}); {_REQUIRED_NATIVE} is required, " + "every other must be a known native extension. All must share one " + "Apple Silicon platform tag and the same exact MLX ABI pin.", + ) parser.add_argument("--out", type=Path, required=True) parser.add_argument( "--codesign-identity", @@ -65,49 +94,95 @@ def main(): ) args = parser.parse_args() name, version, _, core_tags = parse_wheel_filename(args.runtime.name) - native_name, _, _, native_tags = parse_wheel_filename(args.native.name) if name != "mtplx" or core_tags != frozenset({Tag("py3", "none", "any")}): parser.error("The runtime input must be the pure Python MTPLX wheel") - if native_name.replace("-", "_") != "mtplx_qsa_kernels" or len(native_tags) != 1: - parser.error("Expected one tested mtplx_qsa_kernels platform wheel") - tag = next(iter(native_tags)) - if not tag.platform.startswith("macosx_") or not tag.platform.endswith("_arm64"): - parser.error("The native wheel must target Apple Silicon") + + # Validate every native wheel: a known extension, one Apple Silicon tag + # shared across all, and an exact MLX ABI pin that agrees across all. + natives: list[tuple[Path, str, Tag]] = [] + for native_path in args.native: + native_name, _, _, native_tags = parse_wheel_filename(native_path.name) + pkg = native_name.replace("-", "_") + if pkg not in _KNOWN_NATIVE or len(native_tags) != 1: + parser.error( + "Each native input must be one tested platform wheel of a known " + f"extension ({', '.join(sorted(_KNOWN_NATIVE))}); got {native_path.name}" + ) + native_tag = next(iter(native_tags)) + if not native_tag.platform.startswith("macosx_") or not native_tag.platform.endswith("_arm64"): + parser.error("Each native wheel must target Apple Silicon") + natives.append((native_path, pkg, native_tag)) + if not any(pkg == _REQUIRED_NATIVE for _, pkg, _ in natives): + parser.error(f"Expected the {_REQUIRED_NATIVE} platform wheel among the native inputs") + seen_pkgs = [pkg for _, pkg, _ in natives] + if len(set(seen_pkgs)) != len(seen_pkgs): + parser.error("A native extension was passed more than once") + tags = {native_tag for _, _, native_tag in natives} + if len(tags) != 1: + parser.error("All native wheels must share one platform tag") + tag = next(iter(tags)) + args.out.mkdir(parents=True, exist_ok=True) output = args.out / f"mtplx-{version}-{tag}.whl" if output.exists(): parser.error(f"Refusing to overwrite {output}") - with WheelFile(args.runtime) as core, WheelFile(args.native) as native: - metadata_name = next(n for n in native.namelist() if n.endswith(".dist-info/METADATA")) - native_metadata = BytesParser().parsebytes(native.read(metadata_name)) - native_requirements = native_metadata.get_all("Requires-Dist", []) - mlx_pins = [Requirement(r) for r in native_requirements if Requirement(r).name == "mlx"] - if len(mlx_pins) != 1 or not str(mlx_pins[0].specifier).startswith("=="): - parser.error("Native wheel must declare its exact MLX runtime ABI dependency") - members = [n for n in native.namelist() if n.startswith("mtplx_qsa_kernels/")] - for required in ("NOTICE", "LICENSE.txt", "MLX_LICENSE.txt"): - if f"mtplx_qsa_kernels/{required}" not in members: - parser.error(f"Native attribution is missing: {required}") - if not any(n.endswith(".metallib") for n in members) or not any(n.endswith(".so") for n in members): - parser.error("Native wheel lacks its extension or Metal library") + + with contextlib.ExitStack() as stack: + core = stack.enter_context(WheelFile(args.runtime)) + native_requirements: list[str] = [] + mlx_pin: str | None = None + # (WheelFile, pkg, members) per native, in input order. + opened: list[tuple[WheelFile, str, list[str]]] = [] + for native_path, pkg, _native_tag in natives: + native = stack.enter_context(WheelFile(native_path)) + metadata_name = next(n for n in native.namelist() if n.endswith(".dist-info/METADATA")) + native_metadata = BytesParser().parsebytes(native.read(metadata_name)) + requirements = native_metadata.get_all("Requires-Dist", []) + mlx_pins = [Requirement(r) for r in requirements if Requirement(r).name == "mlx"] + if len(mlx_pins) != 1 or not str(mlx_pins[0].specifier).startswith("=="): + parser.error(f"{pkg} must declare its exact MLX runtime ABI dependency") + this_pin = str(mlx_pins[0].specifier) + if mlx_pin is None: + mlx_pin = this_pin + elif this_pin != mlx_pin: + parser.error("Native wheels disagree on the MLX ABI pin") + for requirement in requirements: + if requirement not in native_requirements: + native_requirements.append(requirement) + members = [n for n in native.namelist() if n.startswith(f"{pkg}/")] + spec = _KNOWN_NATIVE[pkg] + for required in spec["required_files"]: + if f"{pkg}/{required}" not in members: + parser.error(f"{pkg} attribution is missing: {required}") + if not any(n.endswith(".so") for n in members): + parser.error(f"{pkg} wheel lacks its extension") + if spec["require_metallib"] and not any(n.endswith(".metallib") for n in members): + parser.error(f"{pkg} wheel lacks its Metal library") + opened.append((native, pkg, members)) + provenance = {p.name: hashlib.sha256(p.read_bytes()).hexdigest() - for p in (args.runtime, args.native)} + for p in (args.runtime, *args.native)} + native_pkgs = [pkg for _native, pkg, _members in opened] with WheelFile(output, "w") as bundled: - for archive, names in ((core, core.namelist()), (native, members)): + sources: list[tuple[WheelFile, list[str], bool]] = [ + (core, core.namelist(), False) + ] + sources.extend((native, members, True) for native, _pkg, members in opened) + for archive, names, is_native in sources: for name in names: if name.endswith("/") or name.endswith(".dist-info/RECORD"): continue if name.startswith("/") or ".." in Path(name).parts: raise ValueError(f"Unsafe archive path: {name}") data = archive.read(name) # WheelFile verifies source RECORD hashes. - if archive is native and args.codesign_identity and name.endswith((".so", ".dylib")): + if is_native and args.codesign_identity and name.endswith((".so", ".dylib")): data = sign_mach_o(name, data, args.codesign_identity) if archive is core and name.endswith(".dist-info/WHEEL"): lines = [line for line in data.decode().splitlines() if not line.startswith(("Tag:", "Root-Is-Purelib:"))] data = ("\n".join([*lines, "Root-Is-Purelib: false", f"Tag: {tag}", ""])).encode() elif archive is core and name.endswith(".dist-info/top_level.txt"): - data += b"mtplx_qsa_kernels\n" + data += ("".join(f"{pkg}\n" for pkg in native_pkgs)).encode() elif archive is core and name.endswith(".dist-info/METADATA"): metadata = BytesParser().parsebytes(data) for requirement in native_requirements: diff --git a/tests/test_bundle_native_runtime_wheel.py b/tests/test_bundle_native_runtime_wheel.py index 8b5852706..397c0caaf 100644 --- a/tests/test_bundle_native_runtime_wheel.py +++ b/tests/test_bundle_native_runtime_wheel.py @@ -33,6 +33,12 @@ DYLIB = "mtplx_qsa_kernels/libmtplx_qsa_kernel_ops.dylib" METALLIB = "mtplx_qsa_kernels/kernels.metallib" +# PR #391 remainder port: the split-K QSA sparse-GQA decode extension is a +# second native wheel (mtplx_native_qsa) that also carries a metallib. +QSA_EXT = "mtplx_native_qsa/_ext.cpython-314-darwin.so" +QSA_DYLIB = "mtplx_native_qsa/libmtplx_native_qsa.dylib" +QSA_METALLIB = "mtplx_native_qsa/mtplx_native_qsa.metallib" + def _write_wheel(path: Path, members: dict[str, bytes]) -> Path: with WheelFile(path, "w") as wheel: @@ -82,6 +88,35 @@ def _run_bundler(monkeypatch, pure: Path, native: Path, out: Path, *extra: str) return out / "mtplx-9.9.9-cp314-cp314-macosx_15_0_arm64.whl" +def _run_bundler_multi(monkeypatch, pure: Path, natives: list[Path], out: Path, *extra: str) -> Path: + monkeypatch.setattr( + sys, "argv", + ["bundle", str(pure), *[str(n) for n in natives], "--out", str(out), *extra], + ) + bundler.main() + return out / "mtplx-9.9.9-cp314-cp314-macosx_15_0_arm64.whl" + + +def _qsa_native_input(tmp_path: Path) -> Path: + return _write_wheel( + tmp_path / "mtplx_native_qsa-9.9.9-cp314-cp314-macosx_15_0_arm64.whl", + { + "mtplx_native_qsa/__init__.py": b"", + QSA_EXT: b"QSA-MACHO-EXT", + QSA_DYLIB: b"QSA-MACHO-DYLIB", + QSA_METALLIB: b"QSA-METALLIB", + "mtplx_native_qsa-9.9.9.dist-info/METADATA": ( + b"Metadata-Version: 2.1\nName: mtplx-native-qsa\nVersion: 9.9.9\n" + b"Requires-Dist: mlx==0.32.2\n" + ), + "mtplx_native_qsa-9.9.9.dist-info/WHEEL": ( + b"Wheel-Version: 1.0\nGenerator: test\nRoot-Is-Purelib: false\n" + b"Tag: cp314-cp314-macosx_15_0_arm64\n" + ), + }, + ) + + def _fake_codesign(calls: list[list[str]], *, timestamp: bool = True): def run(cmd, **kwargs): assert cmd[0] == "/usr/bin/codesign", cmd @@ -151,3 +186,67 @@ def test_a_signature_without_a_secure_timestamp_fails_the_build(tmp_path, monkey _run_bundler( monkeypatch, pure, native, tmp_path / "out", "--codesign-identity", "Developer ID Application: Test" ) + + +def test_qsa_sparse_gqa_extension_is_bundled_and_signed_with_the_qsa_kernels(tmp_path, monkeypatch) -> None: + # The PR #391 remainder QSA split-K decode lane ships a second native + # extension, mtplx_native_qsa (its own .so + .dylib + .metallib). It must + # be bundled and Developer-ID signed alongside the QSA kernels, or + # notarization rejects its ad-hoc-signed Mach-O the way it rejected the + # QSA kernels before ea2560a2. + pure, native = _inputs(tmp_path) + qsa = _qsa_native_input(tmp_path) + calls: list[list[str]] = [] + monkeypatch.setattr(bundler.subprocess, "run", _fake_codesign(calls)) + bundled = _run_bundler_multi( + monkeypatch, pure, [native, qsa], tmp_path / "out", + "--codesign-identity", "Developer ID Application: Test", + ) + with zipfile.ZipFile(bundled) as archive: + assert archive.read(EXT) == b"SIGNED:MACHO-EXT" + assert archive.read(QSA_EXT) == b"SIGNED:QSA-MACHO-EXT" + assert archive.read(QSA_DYLIB) == b"SIGNED:QSA-MACHO-DYLIB" + assert archive.read(QSA_METALLIB) == b"QSA-METALLIB" # metallib is not a Mach-O + top_level = archive.read("mtplx-9.9.9.dist-info/top_level.txt").decode() + record = archive.read("mtplx-9.9.9.dist-info/RECORD").decode().splitlines() + assert "mtplx_qsa_kernels" in top_level.split() + assert "mtplx_native_qsa" in top_level.split() + hashes = {line.split(",")[0]: line.split(",")[1] for line in record if line} + assert hashes[QSA_EXT] == _record_hash(b"SIGNED:QSA-MACHO-EXT") + with WheelFile(bundled) as reopened: # the rewritten RECORD verifies + assert reopened.read(QSA_EXT) == b"SIGNED:QSA-MACHO-EXT" + signed = [Path(call[-1]).name for call in calls if "--sign" in call] + assert "_ext.cpython-314-darwin.so" in signed # QSA kernels + assert "libmtplx_native_qsa.dylib" in signed # split-K decode ext + + +def test_qsa_native_alone_is_rejected_without_the_required_qsa_kernels(tmp_path, monkeypatch) -> None: + # mtplx_native_qsa is an OPTIONAL second extension; the required base is + # mtplx_qsa_kernels. The bundler refuses a native input set that omits it. + pure, _native = _inputs(tmp_path) + qsa = _qsa_native_input(tmp_path) + with pytest.raises(SystemExit): + _run_bundler_multi(monkeypatch, pure, [qsa], tmp_path / "out") + + +def test_qsa_native_missing_its_metallib_is_rejected(tmp_path, monkeypatch) -> None: + # mtplx_native_qsa declares require_metallib=True (it carries a split-K + # metallib), so a build that shipped only the .so is refused. + pure, native = _inputs(tmp_path) + qsa = _write_wheel( + tmp_path / "mtplx_native_qsa-9.9.9-cp314-cp314-macosx_15_0_arm64.whl", + { + "mtplx_native_qsa/__init__.py": b"", + QSA_EXT: b"QSA-MACHO-EXT", + "mtplx_native_qsa-9.9.9.dist-info/METADATA": ( + b"Metadata-Version: 2.1\nName: mtplx-native-qsa\nVersion: 9.9.9\n" + b"Requires-Dist: mlx==0.32.2\n" + ), + "mtplx_native_qsa-9.9.9.dist-info/WHEEL": ( + b"Wheel-Version: 1.0\nGenerator: test\nRoot-Is-Purelib: false\n" + b"Tag: cp314-cp314-macosx_15_0_arm64\n" + ), + }, + ) + with pytest.raises(SystemExit): + _run_bundler_multi(monkeypatch, pure, [native, qsa], tmp_path / "out") diff --git a/tests/test_qsa_sparse_decode.py b/tests/test_qsa_sparse_decode.py new file mode 100644 index 000000000..f4711f228 --- /dev/null +++ b/tests/test_qsa_sparse_decode.py @@ -0,0 +1,727 @@ +"""CPU-only gates for the split-K QSA sparse-GQA DECODE lane. + +Nothing here dispatches Metal. What is covered: + +1. The split geometry -- the host arithmetic that sizes the grid and the + partial buffer. Every split must have work, and the splits must cover the + tile range exactly once. +2. The VISIBLE SET. This is the load-bearing correctness property: the + kernel must attend exactly the keys the shipped rows-gather lane attends. + The kernel's per-slot model and a transcription of the shipped + rows-gather token list's closed form are compared over the interesting + positions, and the two WRONG predicates a reader might reach for (the + prefill kernel's leading-prefix cut, and ``candidate <= q_abs``) are shown + to disagree, so a future edit that adopts either fails here. +3. The kernel SOURCE structure -- that the shipped predicate is the one in + the MSL, that the normalisation really moved to the merge pass, and that + the C++ encoder builds the kernel names the metallib instantiates. +4. The parameter-block layout the host and the device both static_assert. +5. The flag parsing. + +Numeric Metal parity belongs to the operator-controlled guarded window; see +the sparse-decode microbenchmark for the command. +""" + +from __future__ import annotations + +import json +import re +import sys +from pathlib import Path + +import pytest + +from mtplx.kernels import qsa_sparse_decode as lane +from mtplx.native import ( + qsa_sparse_gqa_decode_partial_shape, + qsa_sparse_gqa_decode_split_geometry, +) + +ROOT = Path(__file__).resolve().parents[1] +EXT = ROOT / "native_extensions" / "qsa_sparse_gqa" / "sparse_gqa" +KERNEL_H = EXT / "steel_qsa_sparse_gqa_decode.h" +PARAMS_H = EXT / "qsa_sparse_gqa_decode_params.h" +DECODE_CPP = EXT / "qsa_sparse_gqa_decode.cpp" +METAL = EXT / "qsa_sparse_gqa.metal" + +TILES = ((128, 32), (256, 32), (64, 64), (128, 64)) +RATIO = lane.COMPRESS_RATIO +TOP_K = lane.TOP_K + + +# --------------------------------------------------------------------------- +# 1. split geometry +# --------------------------------------------------------------------------- +@pytest.mark.parametrize("key_tile", [t[0] for t in TILES]) +@pytest.mark.parametrize("key_splits", [1, 2, 3, 4, 6, 8, 12, 16, 17, 33, 64]) +def test_every_split_has_work(key_tile, key_splits): + n_tiles, per_split, n_splits = qsa_sparse_gqa_decode_split_geometry( + lane.SELECTED_TOKENS, key_tile, key_splits + ) + assert n_tiles == -(-lane.SELECTED_TOKENS // key_tile) + assert 1 <= n_splits <= n_tiles + assert per_split >= 1 + # The last split starts inside the tile range -- no empty threadgroup. + assert (n_splits - 1) * per_split < n_tiles + # And the splits reach the end of it. + assert n_splits * per_split >= n_tiles + + +@pytest.mark.parametrize("key_tile", [t[0] for t in TILES]) +@pytest.mark.parametrize("key_splits", [1, 2, 5, 8, 16, 64]) +def test_splits_partition_the_tiles_exactly_once(key_tile, key_splits): + n_tiles, per_split, n_splits = qsa_sparse_gqa_decode_split_geometry( + lane.SELECTED_TOKENS, key_tile, key_splits + ) + covered = [] + for split in range(n_splits): + t0 = split * per_split + t1 = min(n_tiles, t0 + per_split) + assert t0 < t1, "an empty split would write a dead partial state" + covered.extend(range(t0, t1)) + assert covered == list(range(n_tiles)) + + +def test_split_count_never_exceeds_the_tile_count(): + # 17 tiles at BK=128; asking for 64 splits must not dispatch 64. + n_tiles, _, n_splits = qsa_sparse_gqa_decode_split_geometry( + lane.SELECTED_TOKENS, 128, 64 + ) + assert n_tiles == 17 + assert n_splits == 17 + + +def test_partial_shape_tracks_the_split_count(): + for rows in (1, 4): + for key_tile, _ in TILES: + _, _, n_splits = qsa_sparse_gqa_decode_split_geometry( + lane.SELECTED_TOKENS, key_tile, 8 + ) + shape = qsa_sparse_gqa_decode_partial_shape(rows, key_tile, 8) + assert shape == (n_splits, lane.Q_HEADS, rows, lane.HEAD_DIM + 2) + + +def test_partial_buffer_is_small_next_to_what_it_replaces(): + """The split-K partial state must not cost more than the gather it kills. + + The shipped lane writes a [1, 2, rows, 2052, 256] bf16 K/V pair per + layer; this buffer is the whole extra memory traffic the split-K + arrangement introduces, and it has to stay an order of magnitude below. + """ + + rows = 4 + shape = qsa_sparse_gqa_decode_partial_shape(rows, 128, 8) + partial_bytes = shape[0] * shape[1] * shape[2] * shape[3] * 4 + gathered_bytes = 2 * (1 * 2 * rows * (TOP_K * RATIO + RATIO) * 256 * 2) + assert partial_bytes * 8 < gathered_bytes + + +def test_split_geometry_rejects_nonsense(): + with pytest.raises(ValueError): + qsa_sparse_gqa_decode_split_geometry(0, 128, 8) + with pytest.raises(ValueError): + qsa_sparse_gqa_decode_split_geometry(2051, 0, 8) + + +# --------------------------------------------------------------------------- +# 2. the visible set +# --------------------------------------------------------------------------- +def _ids_for(q_abs: int, nb_total: int, seed: int = 7) -> list: + """512 distinct block ids in an order that is deliberately NOT sorted.""" + + state = seed + pool = list(range(nb_total)) + for i in range(len(pool) - 1, 0, -1): + state = (state * 1103515245 + 12345) & 0x7FFFFFFF + j = state % (i + 1) + pool[i], pool[j] = pool[j], pool[i] + return pool[:TOP_K] + + +POSITIONS = [ + 2048, 2049, 2050, 2051, 2052, 2053, + 4093, 4094, 4095, 4096, + 16_383, 16_384, 16_385, 16_386, + 17_407, +] + + +@pytest.mark.parametrize("q_abs", POSITIONS) +def test_kernel_and_shipped_lane_attend_the_same_keys(q_abs): + nb_total = (q_abs + 1) // RATIO + 1 + ids = _ids_for(q_abs, max(nb_total, TOP_K + 1)) + assert lane.visible_sets_agree(ids, q_abs, key_length=q_abs + 1) + + +@pytest.mark.parametrize("q_abs", POSITIONS) +def test_the_dropped_2052nd_slot_is_always_invalid(q_abs): + """The kernel walks 2,051 slots; the shipped lane builds 2,052. + + The extra one is the tail's fourth, token ``((pos+1)//4)*4 + 3``, which is + greater than ``pos`` for every residue class. If that ever stops being + true the kernel is dropping a visible key. + """ + + ids = _ids_for(q_abs, (q_abs + 1) // RATIO + 1) + _, ok = lane.shipped_row_tokens(ids, q_abs) + assert len(ok) == TOP_K * RATIO + RATIO + assert ok[-1] is False or ok[-1] == False # noqa: E712 - mx bool_ compat + + +@pytest.mark.parametrize("q_abs", POSITIONS) +def test_no_key_past_the_query_position_is_ever_attended(q_abs): + ids = _ids_for(q_abs, (q_abs + 1) // RATIO + 1) + for pos in lane.kernel_row_tokens(ids, q_abs, key_length=q_abs + 1): + assert pos <= q_abs + + +def test_the_prefix_cut_the_prefill_kernel_uses_is_wrong_here(): + """The decode selector does not sort, so a leading-prefix cut is wrong. + + ``mtplx.native.qsa_sparse_gqa`` (the prefill entry point) may take the + first ``min(512, complete_blocks)`` slots as the valid ones, because the + prefill branch of ``_select_eager`` sorts its top-k ascending. Its decode + branch does not. + This test builds a row where the two disagree, so an edit that adopts the + prefix cut in the decode kernel fails right here rather than quietly + changing which keys the model attends. + """ + + # complete_blocks 511 < 512, so a prefix cut really cuts: it keeps the + # first 511 SLOTS instead of the visible BLOCKS. Slot 511 holds a + # perfectly visible block and the cut drops it. + q_abs = 2043 + complete = lane.visible_block_count(q_abs) + assert complete == 511 + ids = [511] + list(range(0, 511)) + assert len(ids) == TOP_K + assert len(set(ids)) == TOP_K + per_slot = sorted( + p + for p in lane.kernel_row_tokens(ids, q_abs, key_length=q_abs + 1) + if p >= 0 + ) + prefix_cut = sorted( + block * RATIO + within + for slot_block, block in enumerate(ids) + if slot_block < min(TOP_K, complete) + for within in range(RATIO) + if block * RATIO + within <= q_abs + ) + assert per_slot != prefix_cut + # And the per-slot answer is the shipped lane's. + idx, ok = lane.shipped_row_tokens(ids, q_abs) + shipped = sorted(t for t, good in zip(idx, ok) if good) + assert per_slot == shipped + + +def test_causal_only_predicate_would_double_count_the_tail(): + """``candidate <= q_abs`` is implied by the real predicate but not equal. + + For ``block == complete_blocks`` it admits tokens of the INCOMPLETE block, + which the tail slots already contribute -- so the softmax denominator + would count them twice. Pinning this stops a "simplification". + """ + + q_abs = 2049 # (2050)//4 == 512, tail is non-empty + complete = lane.visible_block_count(q_abs) + ids = [complete] + list(range(0, TOP_K - 1)) + real = lane.kernel_row_tokens(ids, q_abs, key_length=q_abs + 1) + lax = [] + for slot in range(TOP_K * RATIO): + block = ids[slot // RATIO] + candidate = block * RATIO + (slot % RATIO) + lax.append(candidate if candidate <= q_abs else -1) + for within in range(RATIO - 1): + candidate = complete * RATIO + within + lax.append(candidate if candidate <= q_abs else -1) + real_hits = [p for p in real if p >= 0] + lax_hits = [p for p in lax if p >= 0] + assert len(real_hits) == len(set(real_hits)), "the real predicate is a set" + assert len(lax_hits) > len(set(lax_hits)), "the lax one duplicates keys" + + +def test_a_row_always_has_at_least_one_visible_key(): + """The merge's zero-denominator branch must be unreachable under the gate. + + The gate needs ``total // 4 > 512``, so every query row has at least one + complete block, and the selector's top-512 of a set whose finite entries + all outrank the -inf ones must contain at least one of them. + """ + + for q_abs in (2048, 2051, 4096, 17_407): + complete = lane.visible_block_count(q_abs) + assert complete >= 1 + ids = _ids_for(q_abs, complete + 1) + hits = [ + p + for p in lane.kernel_row_tokens(ids, q_abs, key_length=q_abs + 1) + if p >= 0 + ] + assert hits + + +def test_key_length_clamps_every_read(): + """kL is a memory-safety bound: nothing may index past it.""" + + q_abs = 4096 + ids = _ids_for(q_abs, (q_abs + 1) // RATIO + 1) + key_length = 2048 # deliberately shorter than the positions imply + for pos in lane.kernel_row_tokens(ids, q_abs, key_length=key_length): + assert pos < key_length + + +# --------------------------------------------------------------------------- +# 3. kernel source structure +# --------------------------------------------------------------------------- +def test_the_kernel_uses_the_per_slot_predicate(): + src = KERNEL_H.read_text() + assert "raw_block < long(complete_blocks)" in src + assert "complete_blocks = (q_abs + 1) / kCompressRatio" in src + # The prefix cut the prefill kernel uses must NOT appear here. + assert "valid_blocks" not in src + + +def test_the_normalisation_moved_to_the_merge_pass(): + src = KERNEL_H.read_text() + split = src.split("mtplx_qsa_sparse_gqa_decode_merge")[0] + assert "MtplxQsaDecDivOp" not in split + assert "row_bin_op" in split + assert "prow[D] = max_score[0]" in split + assert "prow[D + 1] = sum_score[0]" in split + merge = src.split("mtplx_qsa_sparse_gqa_decode_merge")[1] + assert "denom" in merge and "AccumType(1) / denom" in merge + + +def test_the_split_pass_walks_only_its_own_tile_range(): + src = KERNEL_H.read_text() + assert "const int t0 = split * params->tiles_per_split;" in src + assert "for (int ktile = t0; ktile < t1; ++ktile)" in src + + +def test_the_split_pass_reads_the_offset_from_a_device_buffer(): + src = KERNEL_H.read_text() + assert "const device int* QOffset [[buffer(4)]]" in src + assert "const int q_abs = QOffset[0] + q_pos;" in src + + +def test_the_online_softmax_initialises_to_finite_min_not_neg_inf(): + """finite_min, so an all-masked tile yields 0 rather than NaN.""" + + src = KERNEL_H.read_text() + assert "max_score[i] = Limits::finite_min;" in src + assert "AccumType m = Limits::finite_min;" in src + + +def _expected_kernel_names() -> set: + names = set() + for tname in ("float16", "bfloat16"): + for iname in ("uint32", "int32"): + for bk, dc in TILES: + names.add( + f"mtplx_qsa_sparse_gqa_decode_split_{tname}_{iname}" + f"_bk{bk}_dc{dc}_gqa12_hp16_d256_wm2" + ) + names.add(f"mtplx_qsa_sparse_gqa_decode_merge_{tname}_d256") + return names + + +def test_the_metallib_instantiates_every_name_the_encoder_can_ask_for(): + metal = METAL.read_text() + # Expand the macro the same way the preprocessor will. + instantiated = set() + for tname in ("float16", "bfloat16"): + for iname in ("uint32", "int32"): + for bk, dc in TILES: + instantiated.add( + f"mtplx_qsa_sparse_gqa_decode_split_{tname}_{iname}" + f"_bk{bk}_dc{dc}_gqa12_hp16_d256_wm2" + ) + instantiated.add(f"mtplx_qsa_sparse_gqa_decode_merge_{tname}_d256") + assert instantiated == _expected_kernel_names() + # And the source really contains the macro pieces that build them. + assert '"mtplx_qsa_sparse_gqa_decode_split_" #tname "_" #iname' in metal + assert '"_bk" #bk "_dc" #dc "_gqa12_hp16_d256_wm2"' in metal + assert '"mtplx_qsa_sparse_gqa_decode_merge_" #tname "_d256"' in metal + for bk, dc in TILES: + assert f"iname, itype, {bk}, {dc})" in metal + + +def test_the_encoder_builds_the_same_name_the_metallib_declares(): + cpp = DECODE_CPP.read_text() + assert '"mtplx_qsa_sparse_gqa_decode_split_"' in cpp + assert '"_bk", key_tile_' in cpp + assert '"_dc", dimension_tile_' in cpp + assert '"_gqa", kGqa, "_hp", kHeadPad, "_d",' in cpp + assert '"mtplx_qsa_sparse_gqa_decode_merge_"' in cpp + # The compiled-in geometry constants must match the macro's literals. + assert re.search(r"constexpr int kGqa = 12;", cpp) + assert re.search(r"constexpr int kHeadPad = 16;", cpp) + assert re.search(r"constexpr int kHeadDim = 256;", cpp) + assert re.search(r"constexpr int kWarps = 2;", cpp) + + +def test_the_encoder_dispatches_the_split_grid_on_the_z_axis(): + cpp = DECODE_CPP.read_text() + assert "MTL::Size(rows, kKvHeads, n_splits)" in cpp + assert "MTL::Size(32, kWarps, 1)" in cpp + assert "MTL::Size(kQHeads * rows, 1, 1)" in cpp + assert "MTL::Size(kHeadDim, 1, 1)" in cpp + + +def test_the_selected_token_width_agrees_across_the_three_definitions(): + cpp = DECODE_CPP.read_text() + assert ( + "constexpr int kSelectedTokens = kTopK * kCompressRatio " + "+ (kCompressRatio - 1);" in cpp + ) + assert lane.SELECTED_TOKENS == TOP_K * RATIO + (RATIO - 1) == 2051 + src = KERNEL_H.read_text() + assert "constexpr int kTail = kCompressRatio - 1;" in src + assert ( + "const int selected_tokens = params->topk * kCompressRatio + kTail;" + in src + ) + + +# --------------------------------------------------------------------------- +# 4. parameter-block layout +# --------------------------------------------------------------------------- +def _static_assert_size(header: str, struct: str) -> int: + match = re.search( + rf"static_assert\(sizeof\({struct}\) == (\d+)", header + ) + assert match, f"no sizeof static_assert for {struct}" + return int(match.group(1)) + + +def test_decode_params_layout_is_pinned_and_arithmetically_right(): + header = PARAMS_H.read_text() + # 11 ints + 1 float + 1 pad int, then 4 x 3 int64 at 8-byte alignment. + ints = len(re.findall(r"^ int \w+;", header, flags=re.M)) + body = header.split("struct MtplxQsaSparseGqaDecodeParams")[1].split("};")[0] + n_int = len(re.findall(r"\bint \w+;", body)) + n_float = len(re.findall(r"\bfloat \w+;", body)) + n_i64_arrays = len(re.findall(r"int64_t \w+\[3\];", body)) + assert (n_int, n_float, n_i64_arrays) == (11, 1, 4) + scalar = n_int * 4 + n_float * 4 + aligned = -(-scalar // 8) * 8 + assert _static_assert_size(header, "MtplxQsaSparseGqaDecodeParams") == ( + aligned + n_i64_arrays * 24 + ) + assert ints >= 11 + + +def test_merge_params_layout_is_pinned_and_arithmetically_right(): + header = PARAMS_H.read_text() + body = header.split("struct MtplxQsaSparseGqaMergeParams")[1].split("};")[0] + n_int = len(re.findall(r"\bint \w+;", body)) + n_i64_arrays = len(re.findall(r"int64_t \w+\[3\];", body)) + assert (n_int, n_i64_arrays) == (6, 1) + scalar = n_int * 4 + aligned = -(-scalar // 8) * 8 + assert _static_assert_size(header, "MtplxQsaSparseGqaMergeParams") == ( + aligned + n_i64_arrays * 24 + ) + + +def test_both_params_blocks_are_included_by_both_sides(): + assert '#include "sparse_gqa/qsa_sparse_gqa_decode_params.h"' in ( + KERNEL_H.read_text() + ) + assert '#include "sparse_gqa/qsa_sparse_gqa_decode_params.h"' in ( + DECODE_CPP.read_text() + ) + + +def test_the_build_compiles_the_new_sources(): + cmake = ( + ROOT / "native_extensions" / "qsa_sparse_gqa" / "CMakeLists.txt" + ).read_text() + assert "sparse_gqa/qsa_sparse_gqa_decode.cpp" in cmake + assert "sparse_gqa/steel_qsa_sparse_gqa_decode.h" in cmake + assert "sparse_gqa/qsa_sparse_gqa_decode_params.h" in cmake + + +# --------------------------------------------------------------------------- +# 5. flags +# --------------------------------------------------------------------------- +def test_tile_flag_accepts_only_instantiated_tiles(): + from mtplx.runtime_options import _parse_sparse_decode_tile + + assert _parse_sparse_decode_tile(None) == (128, 32) + assert _parse_sparse_decode_tile("") == (128, 32) + for bk, dc in TILES: + assert _parse_sparse_decode_tile(f"{bk}:{dc}") == (bk, dc) + for bad in ("128", "128:33", "127:32", "128:32:1", "x:y"): + with pytest.raises(ValueError): + _parse_sparse_decode_tile(bad) + + +def test_splits_flag_is_bounded(): + from mtplx.runtime_options import ( + QSA_SPARSE_DECODE_DEFAULT_SPLITS as default, + _parse_sparse_decode_splits, + ) + + assert _parse_sparse_decode_splits(None) == default == 17 + assert _parse_sparse_decode_splits("17") == 17 + for bad in ("0", "65", "-1", "eight"): + with pytest.raises(ValueError): + _parse_sparse_decode_splits(bad) + + +def test_the_flag_is_off_by_default(): + from mtplx.runtime_options import env_bool + + assert env_bool("MTPLX_QSA_SPARSE_DECODE", default=False, env={}) is False + + +def test_engagement_reports_a_pending_install_as_not_installed(): + lane.reset_for_tests() + report = lane.engagement() + assert report["installed"] is False + assert report["disabled_reason"] is None + assert report["verify_kernel"] == 0 + + +def test_parity_thresholds_are_stated_not_implicit(): + assert lane.PARITY_FP32_MAX_ABS_ULPS == 2.0 + assert lane.PARITY_FP32_MAX_REL_L2 == 5.0e-4 + assert lane.PARITY_SHIPPED_MAX_REL_L2 == 5.0e-2 + assert lane.PARITY_MIN_TOP1 == 0.98 + + +# --------------------------------------------------------------------------- +# 7. what the 2026-09-02 runs settled +# --------------------------------------------------------------------------- +def test_the_tight_gate_is_tighter_than_the_measured_shipped_delta(): + """The fp32 gate must be able to FAIL if the attribution is wrong. + + Every configuration measured rel_l2 4.78e-3 against the shipped path. If + that delta really is the shipped path's own bf16 score and probability + casts, the same comparison against the fp32 reference collapses to output + rounding. A gate set above 4.78e-3 could not tell those apart, so it + would certify nothing. + """ + + measured_vs_shipped = 4.78e-3 + assert lane.PARITY_FP32_MAX_REL_L2 < measured_vs_shipped + # ... while the loose bar stays an order of magnitude above it, because it + # bounds the reference's quantisation, not the kernel's error. + assert lane.PARITY_SHIPPED_MAX_REL_L2 > 10 * measured_vs_shipped + + +def test_the_reference_ladder_has_three_distinct_rungs(): + import inspect + + src = inspect.getsource(lane.reference_attention) + # The shipped rung must keep BOTH bf16 roundings... + assert "probs = probs.astype(queries.dtype)" in src + # ...and the fp32 rung must remove them by UPCASTING (exact), never by + # changing the operands. + assert "q_view = q_view.astype(mx.float32)" in src + assert "v_view = v_view.astype(mx.float32)" in src + for fn, scores, probs in ( + (lane.stock_reference, False, False), + (lane.shipped_fp32_probs_reference, False, True), + (lane.fp32_reference, True, True), + ): + wrapped = inspect.getsource(fn) + assert f"fp32_scores={scores}" in wrapped + assert f"fp32_probs={probs}" in wrapped + + +def test_every_reference_rung_returns_the_query_dtype(): + """All three must be comparable element for element against bf16 output.""" + + import inspect + + src = inspect.getsource(lane.reference_attention) + assert src.rstrip().endswith(".astype(queries.dtype)") + + +def test_the_default_split_target_reaches_one_tile_per_threadgroup(): + """17 is the tile count at BK=128, so it is the knob's saturation point.""" + + from mtplx.runtime_options import ( + QSA_SPARSE_DECODE_DEFAULT_SPLITS as default, + ) + + n_tiles, per_split, n_splits = qsa_sparse_gqa_decode_split_geometry( + lane.SELECTED_TOKENS, 128, default + ) + assert (n_tiles, per_split, n_splits) == (17, 1, 17) + assert lane.VERIFY_ROWS * lane.KV_HEADS * n_splits == 136 + + +@pytest.mark.parametrize("larger", [18, 32, 33, 64]) +def test_split_targets_above_the_tile_count_are_the_same_configuration(larger): + """Why the first sweep's s17 and s32 rows are one config measured twice. + + Their 5.3% spread is therefore the bench's noise floor, not a result, and + nothing may be called a winner on a margin under it. + """ + + at_17 = qsa_sparse_gqa_decode_split_geometry(lane.SELECTED_TOKENS, 128, 17) + assert qsa_sparse_gqa_decode_split_geometry( + lane.SELECTED_TOKENS, 128, larger + ) == at_17 + + + + +def test_the_lane_and_the_native_wrapper_agree_on_the_default(): + from mtplx import native + from mtplx.runtime_options import ( + QSA_SPARSE_DECODE_DEFAULT_SPLITS as default, + ) + + assert native._DEFAULT_KEY_SPLITS == default + assert native._DEFAULT_TILE == (128, 32) + + + + + + +# --------------------------------------------------------------------------- +# 6. the nanobind ABI guard +# +# The first guarded run built the extension cleanly and then failed at the +# first call with "incompatible function arguments ... queries: +# mlx::core::array" while invoking with mlx.core.array. The cause was not the +# stream kwarg: the extension was built against nanobind internals v19 while +# mlx.core uses v21, so the two got separate __nb_internals__mlx__ +# capsules and no array could ever be cast. These pin the detector. +# --------------------------------------------------------------------------- +import importlib.util # noqa: E402 + + + + + + + + + + + + + + + + + + +def test_the_native_wrapper_uses_the_same_tag_regex(): + """One definition of "what a nanobind ABI tag looks like", two readers.""" + + from mtplx import native + + blob = b"\x00v21_system_libcpp_abi1\x00" + assert [m.group(1) for m in native._NB_ABI_TAG_RE.finditer(blob)] == [b"21"] + + +def test_the_native_wrapper_reads_a_versions_from_a_binary(tmp_path): + from mtplx import native + + blob = tmp_path / "fake.so" + blob.write_bytes(b"v19_system_libcpp_abi1\x00") + assert native._nanobind_internals_version(blob) == 19 + blob.write_bytes(b"nothing") + assert native._nanobind_internals_version(blob) is None + + +def test_the_cmake_guard_is_a_fatal_error_not_a_warning(): + cmake = ( + ROOT / "native_extensions" / "qsa_sparse_gqa" / "CMakeLists.txt" + ).read_text() + assert "MTPLX_NANOBIND_DIR" in cmake + assert "NB_INTERNALS_VERSION" in cmake + guard = cmake.split("if(MTPLX_NB_INTERNALS AND MTPLX_MLX_NB_INTERNALS)")[1] + assert "FATAL_ERROR" in guard + assert "MTPLX_NANOBIND_DIR=" in guard + # The one wrong fix must stay called out where someone would reach for it. + assert "NEVER \"fix\" a mismatch by defining NB_INTERNALS_VERSION" in cmake + + +def test_the_wrappers_omit_a_none_stream(): + """Passing an explicit None was not the bug, but it is one fewer variable.""" + + source = (ROOT / "mtplx" / "native" / "__init__.py").read_text() + assert source.count("if stream is None:") == 2 + assert "stream=stream)" in source + assert "stream=stream,\n )" not in source + + +# --------------------------------------------------------------------------- +# 8. the micro's parity-gate ladder +# +# The two-gate ladder renamed PARITY_MAX_ABS_ULPS -> PARITY_FP32_MAX_ABS_ULPS +# (and PARITY_MAX_REL_L2 -> PARITY_FP32_MAX_REL_L2), and +# The sparse-decode microbenchmark kept reading the old names in its +# report header -- so it crashed with AttributeError at startup, INSIDE a +# guarded window, after the operator had already taken the box. These pin the +# gate-building surface so that class of drift fails here instead. +# +# Nothing below dispatches Metal: the gate builders are fed stubbed numbers, +# and the one call to the micro's own ``compare`` runs on the CPU stream. +# --------------------------------------------------------------------------- + + + + + + +#: Comfortably inside every threshold, so a test can break exactly one. +PASSING_STATS = { + "vs_fp32": {"max_abs_bf16_ulps": 1.0, "rel_l2": 1.0e-4, "top1": 1.0}, + "vs_shipped": {"max_abs_bf16_ulps": 9.0, "rel_l2": 4.78e-3, "top1": 1.0}, +} + + +def _stats(reference=None, key=None, value=None): + """The passing stats, optionally with ONE statistic pushed out of bounds.""" + + out = {k: dict(v) for k, v in PASSING_STATS.items()} + if reference is not None: + out[reference][key] = value + return out + + + + + + + + + + + + + + + + + + + + + + +# --- the verdict roll-up, which is what the process status is -------------- + + + + + + + + + + + + + + diff --git a/tests/test_qsa_sparse_decode_wiring.py b/tests/test_qsa_sparse_decode_wiring.py new file mode 100644 index 000000000..4fe724296 --- /dev/null +++ b/tests/test_qsa_sparse_decode_wiring.py @@ -0,0 +1,771 @@ +"""The WIRING of ``MTPLX_QSA_SPARSE_DECODE``, not its arithmetic. + +``tests/test_qsa_sparse_decode.py`` pins the kernel's selection model, +its split geometry and its parity gates. This file pins the thing that +actually failed on 2026-09-02: the lane was armed, the cache installed, and +the kernel never ran, because the ONE call site that could reach the verify +width asked about a width the flag never armed. + +THE DEFECT, reproduced here on the CPU stream with stub shapes. + +A fixed-capacity S=4 verify falls through ``legacy_fused=False`` into +``_select_eager``, whose sparse-decode question named the M=1 width. It read +a zero row count, returned False, and handed attention the rows-gather lane. + +Every test below is host-side. Nothing evaluates an MLX array, and no test +touches the GPU: the routing decision this file is about is decided entirely +from python ints and cache attributes, which is exactly why it could go wrong +without leaving a mark on a receipt. +""" + +from __future__ import annotations + +import ast +import importlib.util +import inspect +import re +from pathlib import Path +from types import SimpleNamespace + +import pytest + +from mtplx.kernels import qsa_sparse_decode as lane +from mtplx.models.qwen4_exp import Attention, QSAIndexer + +ROOT = Path(__file__).resolve().parents[1] + + +@pytest.fixture(autouse=True) +def _clean_lane(): + lane.reset_for_tests() + yield + lane.reset_for_tests() + + +@pytest.fixture +def armed_verify(monkeypatch): + """Arm the verify width only, exactly as the 2026-09-02 window did.""" + + from mtplx import runtime_options + + monkeypatch.setattr(runtime_options, "_QSA_SPARSE_DECODE", True) + return runtime_options + + +#: A 16 K prompt reserves 1,024 decode tokens and rounds to ratio: the bank +#: the 2026-09-02 window ran on. ``17408 // 4 = 4352 > 512`` -- servable. +LONG_BANK = 17408 +#: The 1 K cell of the served battery: a 1,024-token prompt plus the same +#: 1,024-token reserve. ``2048 // 4 == 512``, NOT > 512, so the kernel's own +#: dense/sparse boundary refuses it -- while ``k_eff`` is a full 512. +DENSE_BANK = 2048 + + +def fixed_cache(*, decode_rows: int = 4, bank: int = LONG_BANK) -> SimpleNamespace: + """The state a ``TensorOffsetQSACache`` hands the routing predicate. + + ``kv.keys`` is the fixed K/V backing. Its token dimension is what + ``update_and_fetch`` returns and therefore what the attention call site + passes as ``total_tokens`` -- the predicate must read the same number, so + the stub carries it. + """ + + return SimpleNamespace( + fixed_capacity=True, + qsa_sparse_decode_rows=int(decode_rows), + kv=SimpleNamespace(keys=SimpleNamespace(shape=(1, 2, int(bank), 256))), + ) + + +def indexer(*, ratio: int = 4, block_topk: int = 512) -> SimpleNamespace: + """The two module attributes the predicate reads, bound to the real method.""" + + stub = SimpleNamespace(ratio=int(ratio), block_topk=int(block_topk)) + stub._sparse_decode_route = QSAIndexer._sparse_decode_route.__get__(stub) + return stub + + +# --------------------------------------------------------------------------- +# The defect itself +# --------------------------------------------------------------------------- +def test_the_verify_width_routes_on_the_stack_that_measured_the_control( + armed_verify, +): + """rows=4, verify armed: the kernel must serve.""" + + assert indexer()._sparse_decode_route( + fixed_cache(), rows=4, k_eff=512, site="select_eager_verify" + ) + counters = lane.route_counters() + assert counters["route_hits"] == 1 + assert counters["route_sites"] == {"select_eager_verify": 1} + + +def test_the_eager_selector_asks_the_verify_width(): + """``_select_eager`` is the selector a fixed-M4 verify actually reaches.""" + + source = inspect.getsource(QSAIndexer._select_eager) + calls = re.findall(r"_sparse_decode_route\((.*?)\)", source, re.S) + assert len(calls) == 1, calls + assert "select_eager_verify" in calls[0] + + +def test_every_route_call_site_names_itself(): + """A single total would not have shown which site declined.""" + + tree = ast.parse((ROOT / "mtplx" / "models" / "qwen4_exp.py").read_text()) + sites = [] + for node in ast.walk(tree): + if not isinstance(node, ast.Call): + continue + func = node.func + if not ( + isinstance(func, ast.Attribute) and func.attr == "_sparse_decode_route" + ): + continue + keywords = {kw.arg for kw in node.keywords} + assert "site" in keywords, ast.dump(node) + for kw in node.keywords: + if kw.arg == "site": + sites.append(kw.value.value) + assert sorted(sites) == ["select_eager_verify"] + + +# --------------------------------------------------------------------------- +# Routing vs failure -- the two kinds of "no" +# --------------------------------------------------------------------------- +def test_an_unarmed_width_is_one_cached_bool_test(armed_verify): + """The OFF path must not do bookkeeping: it runs per QSA layer per forward.""" + + assert not indexer()._sparse_decode_route( + fixed_cache(), rows=1, k_eff=512, site="select_eager_verify" + ) + assert lane.route_counters()["route_declines"] == {} + assert lane.route_counters()["route_hits"] == 0 + + +def test_the_off_path_reads_the_flag_before_anything_else(): + """An unarmed process pays one cached-bool test.""" + + source = inspect.getsource(QSAIndexer._sparse_decode_route) + body = source[source.index('"""', source.index('"""') + 3) + 3 :] + flag = body.index("qsa_sparse_decode_enabled()") + assert flag < body.index("getattr(cache") + assert flag < body.index("from mtplx.kernels import qsa_sparse_decode") + + +def test_a_growable_cache_is_routing_and_is_counted(armed_verify): + growable = SimpleNamespace(fixed_capacity=False) + assert not indexer()._sparse_decode_route( + growable, rows=4, k_eff=512, site="select_eager_verify" + ) + assert lane.route_counters()["route_declines"] == { + "select_eager_verify: growable cache": 1 + } + + +def test_a_width_the_predicate_was_not_asked_about_returns_quietly(armed_verify): + """A 16 K prefill row count is neither width; no decline, no hit.""" + + assert not indexer()._sparse_decode_route( + fixed_cache(), rows=16384, k_eff=512, site="select_eager_verify" + ) + assert lane.route_counters()["route_declines"] == {} + + +def test_a_cache_built_without_the_lane_raises_at_the_armed_width(armed_verify): + with pytest.raises(RuntimeError, match="qsa_sparse_decode_rows=0"): + indexer()._sparse_decode_route( + fixed_cache(decode_rows=0), + rows=4, + k_eff=512, + + site="select_eager_verify", + ) + + +def test_a_geometry_the_metallib_is_not_built_for_raises(armed_verify): + with pytest.raises(RuntimeError, match="ratio-4 top-512"): + indexer(ratio=2)._sparse_decode_route( + fixed_cache(), rows=4, k_eff=512, site="select_eager_verify" + ) + + +# --------------------------------------------------------------------------- +# REQUEST SHAPE: routing, not failure. Two production 500s live here. +# --------------------------------------------------------------------------- +def test_a_partial_budget_routes_to_stock_instead_of_returning_500(armed_verify): + """The first regression: HumanEval prompts, a few hundred tokens.""" + + assert not indexer()._sparse_decode_route( + fixed_cache(bank=LONG_BANK), rows=4, k_eff=33, + site="select_eager_verify", + ) + counters = lane.route_counters() + assert counters["route_hits"] == 0 + assert counters["request_declines"] == 1 + assert counters["route_declines"] == {"select_eager_verify: partial_budget": 1} + + +def test_a_1k_prompt_in_the_dense_regime_routes_to_stock(armed_verify): + """The second regression, and the one a full budget hides. + + The 1 K cell of the 2026-09-02 served battery: a 1,024-token prompt in a + 2,048-token fixed bank. ``k_eff`` is a FULL 512 -- the budget gate passes + -- but ``2048 // 4 == 512`` is not ``> 512``, so the kernel refused the + call from inside ``attention()`` and the server returned 500. The two + questions are different and the predicate must ask both. + """ + + assert not indexer()._sparse_decode_route( + fixed_cache(bank=DENSE_BANK), rows=4, k_eff=512, + site="select_eager_verify", + ) + counters = lane.route_counters() + assert counters["route_hits"] == 0 + assert counters["request_declines"] == 1 + assert counters["route_declines"] == {"select_eager_verify: short_context": 1} + assert counters["request_decline_extremes"]["tokens_min"] == DENSE_BANK + + +def test_the_dense_regime_boundary_is_the_kernels_own(armed_verify): + """One token of bank either side of the native contract's boundary.""" + + route = indexer()._sparse_decode_route + assert not route( + fixed_cache(bank=2048), rows=4, k_eff=512, site="s" + ) + assert not route( + fixed_cache(bank=2051), rows=4, k_eff=512, site="s" + ) + assert route( + fixed_cache(bank=2052), rows=4, k_eff=512, site="s" + ) + assert lane.route_counters()["route_hits"] == 1 + assert lane.SHORT_CONTEXT_TOKENS == 2052 + + +def test_the_crossover_inside_one_request_declines_then_binds(armed_verify): + """Context growing from below the boundary to above it, one cache. + + A fixed bank grows by capacity transition, and each transition re-traces + and re-asks. Before the boundary the stock chain is correct; after it the + kernel must bind -- and the earlier declines must not have poisoned it. + """ + + cache = fixed_cache(bank=DENSE_BANK) + route = indexer()._sparse_decode_route + assert not route(cache, rows=4, k_eff=512, site="select_eager_verify") + # The capacity transition: same cache object, a wider bank. + cache.kv.keys.shape = (1, 2, 4096, 256) + assert route(cache, rows=4, k_eff=512, site="select_eager_verify") + counters = lane.route_counters() + assert counters["route_hits"] == 1 + assert counters["request_declines"] == 1 + + +def test_a_full_budget_forward_still_binds_after_either_decline(armed_verify): + """Short requests must not poison the long-context arm.""" + + route = indexer()._sparse_decode_route + assert not route( + fixed_cache(bank=LONG_BANK), rows=4, k_eff=33, site="s" + ) + assert not route( + fixed_cache(bank=DENSE_BANK), rows=4, k_eff=512, site="s" + ) + assert route( + fixed_cache(bank=LONG_BANK), rows=4, k_eff=512, site="s" + ) + counters = lane.route_counters() + assert counters["route_hits"] == 1 + assert counters["request_declines"] == 2 + + +def test_every_request_shape_reason_routes_rather_than_raises(armed_verify): + """No branch of the mirror may be a raise.""" + + route = indexer()._sparse_decode_route + cases = { + "short_context": dict(bank=DENSE_BANK, k_eff=512), + "partial_budget": dict(bank=LONG_BANK, k_eff=33), + "rows_exceed_context": dict(bank=2, k_eff=512), + "empty_context": dict(bank=0, k_eff=512), + } + for reason, kwargs in cases.items(): + lane.reset_for_tests() + assert not route( + fixed_cache(bank=kwargs["bank"]), rows=4, k_eff=kwargs["k_eff"], + site="s", + ), reason + assert lane.route_counters()["route_declines"] == {f"s: {reason}": 1} + + +def test_the_decline_key_stays_bounded_while_the_numbers_track(armed_verify): + """One key per (site, reason); the tokens and blocks ride as min/max. + + A server sees every context length there is, so a key carrying the count + would grow without bound. + """ + + route = indexer()._sparse_decode_route + for blocks in (7, 33, 511, 33): + assert not route( + fixed_cache(bank=LONG_BANK), rows=4, k_eff=blocks, + site="select_eager_verify", + ) + counters = lane.route_counters() + assert counters["route_declines"] == {"select_eager_verify: partial_budget": 4} + assert counters["request_declines"] == 4 + extremes = counters["request_decline_extremes"] + assert extremes["blocks_min"] == 7 and extremes["blocks_max"] == 511 + assert extremes["tokens_min"] == extremes["tokens_max"] == LONG_BANK + assert counters["short_context_tokens"] == 2052 + + +def test_the_threshold_is_the_full_budget_the_abi_needs(): + assert lane.SHORT_CONTEXT_TOKENS == (lane.TOP_K + 1) * lane.COMPRESS_RATIO == 2052 + + +def test_a_request_shape_decline_prints_nothing(armed_verify, capsys): + """Once per QSA layer per request: a print here is a log flood.""" + + indexer()._sparse_decode_route( + fixed_cache(bank=DENSE_BANK), rows=4, k_eff=512, + site="select_eager_verify", + ) + captured = capsys.readouterr() + assert captured.out == "" + assert captured.err == "" + + +# --------------------------------------------------------------------------- +# The mirror against the native contract +# --------------------------------------------------------------------------- +NATIVE = (ROOT / "mtplx" / "native" / "__init__.py").read_text() + + +def native_decode_reasons(): + """Every string ``qsa_sparse_gqa_decode_unsupported_reason`` can return.""" + + tree = ast.parse(NATIVE) + fn = next( + node + for node in ast.walk(tree) + if isinstance(node, ast.FunctionDef) + and node.name == "qsa_sparse_gqa_decode_unsupported_reason" + ) + out = [] + for node in ast.walk(fn): + if not isinstance(node, ast.Return) or node.value is None: + continue + value = node.value + if isinstance(value, ast.Constant) and isinstance(value.value, str): + out.append(value.value) + return out + + +def test_every_request_shape_reason_the_kernel_can_return_is_mirrored(): + """The 1 K regression was a PARTIAL mirror; this is what closes it.""" + + native = native_decode_reasons() + for reason in lane.REQUEST_SHAPE_REASONS: + assert reason in native, reason + + +def test_the_context_branches_of_the_native_contract_are_all_claimed(): + """Any native reason mentioning the request's own size must be mirrored. + + A new context-dependent branch upstream fails this test instead of + reaching a served request as a 500. + """ + + request_words = ("total_tokens", "context", "query rows") + unclaimed = [ + reason + for reason in native_decode_reasons() + if any(word in reason for word in request_words) + and reason not in lane.REQUEST_SHAPE_REASONS + and "must be a host integer" not in reason + and "must be an exact host integer" not in reason + and "cannot be bool" not in reason + ] + assert unclaimed == [], unclaimed + + +def test_the_kernel_wrapper_names_a_mirror_bug_rather_than_blaming_the_request(): + source = inspect.getsource(lane.attention) + assert "REQUEST_SHAPE_REASONS" in source + assert "did not mirror the" in source + + +def test_context_decline_is_none_only_for_a_servable_shape(): + assert lane.context_decline( + total_tokens=LONG_BANK, rows=4, k_eff=512, capacity=LONG_BANK + ) is None + assert lane.context_decline( + total_tokens=DENSE_BANK, rows=4, k_eff=512, capacity=DENSE_BANK + ) == "short_context" + assert lane.context_decline( + total_tokens=LONG_BANK, rows=4, k_eff=511, capacity=LONG_BANK + ) == "partial_budget" + assert lane.context_decline( + total_tokens=LONG_BANK + 8, rows=4, k_eff=512, capacity=LONG_BANK + ) == "context_exceeds_capacity" + assert lane.context_decline( + total_tokens=lane.MAX_CONTEXT + 4, rows=4, k_eff=512, + capacity=lane.MAX_CONTEXT + 4, + ) == "context_above_limit" + + +def test_nothing_raises_when_the_flag_is_off(): + """The whole predicate is one host-side compare when nobody armed it.""" + + assert not indexer()._sparse_decode_route( + fixed_cache(decode_rows=0), + rows=4, + k_eff=17, + + site="select_eager_verify", + ) + assert lane.route_counters()["route_hits"] == 0 + + +# --------------------------------------------------------------------------- +# The per-layer proof inside the traced verify body +# --------------------------------------------------------------------------- +def required(cache, rows, *, indexer_present=True): + stub = SimpleNamespace(indexer=object() if indexer_present else None) + return Attention._sparse_decode_required(stub, cache, rows) + + +def test_the_armed_verify_width_on_a_fixed_cache_is_required(armed_verify): + assert required(fixed_cache(), 4) + + +def test_a_width_the_flag_does_not_arm_is_not_required(armed_verify): + assert not required(fixed_cache(), 1) + + +def test_a_growable_cache_is_never_required(armed_verify): + assert not required(SimpleNamespace(fixed_capacity=False), 4) + + +def test_a_layer_without_an_indexer_is_never_required(armed_verify): + assert not required(fixed_cache(), 4, indexer_present=False) + + +def guard(sel_mask, *, rows=4, before=None): + stub = SimpleNamespace() + return Attention._require_sparse_decode_lane( + stub, + sel_mask, + rows=rows, + before=before or {"route_hits": 0, "request_declines": 0}, + ) + + +def test_the_guard_accepts_the_sparse_lane(armed_verify): + guard(("sparse_blocks", object())) + + +def test_the_guard_accepts_a_forward_that_routed_for_short_context(armed_verify): + """The HumanEval case: the indexer declined, and that is correct.""" + + before = lane.route_snapshot() + lane.note_request_decline( + "select_eager_verify", "short_context", total_tokens=2048, blocks=512 + ) + guard(("gather_rows", object(), object()), before=before) + + +def test_the_guard_refuses_a_full_budget_forward_on_another_lane(armed_verify): + with pytest.raises(RuntimeError, match="gather_rows"): + guard(("gather_rows", object(), object())) + + +@pytest.mark.parametrize( + "sel_mask, name", + [ + (("flash", object(), 0), "flash"), + (("flash_prefill", object(), object()), "flash_prefill"), + (None, "no_selection"), + (object(), "dense_mask"), + ], +) +def test_the_guard_names_the_lane_attention_actually_got(armed_verify, sel_mask, name): + with pytest.raises(RuntimeError, match=name): + guard(sel_mask) + + +def test_the_guard_sits_before_every_other_attention_branch(): + source = inspect.getsource(Attention.__call__) + call = source.index("self._require_sparse_decode_lane(") + for lane_name in ("flash", "flash_prefill", "sparse_blocks", "gather_rows"): + assert call < source.index(f'sel_mask[0] == "{lane_name}"') + + +def test_the_guard_samples_before_the_indexer_runs(): + """The routing happens inside ``self.indexer(...)``, so the baseline must + be taken above it or the delta would always be zero.""" + + source = inspect.getsource(Attention.__call__) + assert source.index("sparse_before = _sparse_route_snapshot()") < source.index( + "self.indexer(" + ) + + +def test_the_guard_skips_vision_requests(): + source = inspect.getsource(Attention.__call__) + assert "if vrope is None and sparse_required:" in source + + +# --------------------------------------------------------------------------- +# The graph-level proof +# --------------------------------------------------------------------------- +ZERO = {"route_hits": 0, "request_declines": 0} + + +def test_assert_traced_raises_when_the_armed_lane_missed_the_graph(armed_verify): + with pytest.raises(lane.SparseDecodeContractError, match="not in the traced"): + lane.assert_traced(4, before=ZERO, where="compiled verify") + + +def test_assert_traced_passes_once_a_route_hit_landed(armed_verify): + lane.note_route_hit("select_eager_verify") + lane.assert_traced(4, before=ZERO, where="compiled verify") + + +def test_assert_traced_accepts_a_forward_that_routed_for_short_context( + armed_verify, +): + """The stock lane IS correct below 2,052 tokens; the assertion says so.""" + + lane.note_request_decline( + "select_eager_verify", "short_context", total_tokens=2048, blocks=512 + ) + lane.assert_traced(4, before=ZERO, where="compiled verify") + + +def test_assert_traced_still_raises_on_an_unbound_full_budget_forward( + armed_verify, +): + """A short forward earlier in the run must not vouch for this one.""" + + lane.note_request_decline( + "select_eager_verify", "short_context", total_tokens=2048, blocks=512 + ) + before = lane.route_snapshot() + with pytest.raises( + lane.SparseDecodeContractError, match="shape is one the lane can serve" + ): + lane.assert_traced(4, before=before, where="compiled verify") + + +def test_assert_traced_measures_the_delta_not_the_total(armed_verify): + """A hit from a PREVIOUS trace must not vouch for this one.""" + + lane.note_route_hit("select_eager_verify") + with pytest.raises(lane.SparseDecodeContractError): + lane.assert_traced( + 4, before={"route_hits": 1, "request_declines": 0}, where="compiled verify" + ) + + +def test_assert_traced_is_silent_off_the_armed_width(armed_verify): + lane.assert_traced(1, before=ZERO, where="compiled verify") + lane.assert_traced(16, before=ZERO, where="compiled verify") + + +def test_assert_traced_is_silent_when_nothing_is_armed(): + lane.assert_traced(4, before=ZERO, where="compiled verify") + + +def test_the_snapshot_carries_both_ways_a_forward_can_prove_engagement(): + assert lane.route_snapshot() == {"route_hits": 0, "request_declines": 0} + lane.note_route_hit("x") + lane.note_request_decline("x", "short_context", total_tokens=2048, blocks=9) + assert lane.route_snapshot() == {"route_hits": 1, "request_declines": 1} + + +def test_the_compiled_verify_body_samples_and_asserts_across_the_forward(): + source = (ROOT / "mtplx" / "graphbank.py").read_text() + body = source[source.index("def verify_step(input_ids, *args):") :] + body = body[: body.index(" return verify_step")] + sample = body.index("sparse_route_before = ") + forward = body.index("result = live._runtime_forward(") + check = body.index("_qsa_sparse_lane.assert_traced(") + assert sample < forward < check + + +# --------------------------------------------------------------------------- +# Construction owns the "cache built without the lane" gate +# --------------------------------------------------------------------------- +def make_cache(monkeypatch, *, decode=False, armed=False): + from mtplx import graphbank + + monkeypatch.setattr( + graphbank, "qsa_sparse_decode_enabled", lambda: bool(armed) + ) + kv = SimpleNamespace(cache=[None, None, None], step=256, keys=None, values=None) + return graphbank.TensorOffsetQSACache( + kv, + None, + None, + compress_ratio=4, + rows_gather_kv_m4=None, + qsa_sparse_decode=decode, + ) + + +def test_an_armed_flag_on_a_cache_built_without_the_lane_raises(monkeypatch): + with pytest.raises(RuntimeError, match="constructed without the lane"): + make_cache(monkeypatch, decode=False, armed=True) + + +def test_an_unarmed_process_builds_the_cache_untouched(monkeypatch): + cache = make_cache(monkeypatch, decode=False, armed=False) + assert cache.qsa_sparse_decode_rows == 0 + + +def test_a_disabled_probe_now_fails_the_build_rather_than_the_arm(monkeypatch): + monkeypatch.setattr(lane, "install", lambda *a, **k: False) + monkeypatch.setattr(lane, "_DISABLED_REASON", "parity probe failed on 'x'") + with pytest.raises(RuntimeError, match="parity probe failed"): + make_cache(monkeypatch, decode=True, armed=True) + + +def test_a_passing_probe_wires_the_verify_rows(monkeypatch): + monkeypatch.setattr(lane, "install", lambda *a, **k: True) + cache = make_cache(monkeypatch, decode=True, armed=True) + assert cache.qsa_sparse_decode_rows == lane.VERIFY_ROWS + + +def test_the_shadow_twins_carry_the_lane_without_a_defaulting_getattr(): + """A twin that silently defaulted to False would revert to stock. + + Both twin sites build a ``TensorOffsetQSACache`` from an entry that always + sets this attribute, so reading it directly is what makes a site that + forgot it loud. + """ + + source = (ROOT / "mtplx" / "graphbank.py").read_text() + assert 'getattr(\n entry, "qsa_sparse_decode"' not in source + assert source.count("qsa_sparse_decode=entry.qsa_sparse_decode") == 2 + + +# --------------------------------------------------------------------------- +# The evidence a receipt and a log carry +# --------------------------------------------------------------------------- +def test_the_install_verdict_reaches_stderr_not_only_the_logger(capsys): + lane._emit("[mtplx] qsa_sparse_decode: probe line") + assert "[mtplx] qsa_sparse_decode: probe line" in capsys.readouterr().err + + +def test_install_emits_both_the_line_and_the_json(armed_verify): + source = inspect.getsource(lane.install) + assert source.count("_emit(engagement_line(enabled=True))") == 1 + assert source.count("_emit(engagement_line(enabled=False))") == 1 + assert source.count('"[mtplx] qsa_sparse_decode install: "') == 2 + + +def test_the_engagement_line_states_rows_tile_splits_and_the_probe(armed_verify): + lane._PROBE_REPORT["worst"] = { + "cell": "verify-4096", + "vs_fp32": {"max_abs_ulps": 0.75, "rel_l2": 1.2e-5, "top1": 1.0}, + "vs_shipped": {"rel_l2": 4.78e-3}, + } + lane._COUNTS["cache_installs"] = 12 + lane._COUNTS["probe_runs"] = 2 + line = lane.engagement_line(enabled=True) + assert line.startswith("[mtplx] qsa_sparse_decode armed:") + for token in ("rows=4", "tile=128:32", "splits=17", "caches=12", "probe_runs=2"): + assert token in line + assert "'verify-4096'" in line + + +def test_the_off_line_carries_the_reason(monkeypatch): + monkeypatch.setattr(lane, "_DISABLED_REASON", "parity probe failed on 'x'") + assert lane.engagement_line(enabled=False) == ( + "[mtplx] qsa_sparse_decode: off (parity probe failed on 'x')" + ) + + +def test_the_receipt_never_raises_while_the_probe_is_pending(armed_verify): + block = lane.receipt() + assert block["armed"] is True + assert block["pending"] is True + assert block["installed"] is False + assert block["disabled_reason"] is None + + +def test_the_receipt_carries_everything_the_owner_asked_for(armed_verify): + block = lane.receipt() + for key in ( + "armed", + "installed", + "disabled_reason", + "tile", + "splits", + "probe", + "route_hits", + "route_sites", + "route_declines", + "kernel_calls", + "cache_installs", + ): + assert key in block, key + assert block["tile"] == [128, 32] + assert block["splits"] == 17 + + +def test_route_hits_and_kernel_calls_are_separate_counters(armed_verify): + lane.note_route_hit("select_eager_verify") + lane._COUNTS["verify_kernel"] += 3 + block = lane.receipt() + assert block["route_hits"] == 1 + assert block["kernel_calls"] == {"verify_kernel": 3} + + +def test_reset_clears_the_route_state_too(): + lane.note_route_hit("a") + lane.note_route_decline("b") + lane.note_request_decline("a", "short_context", total_tokens=2048, blocks=9) + lane.reset_for_tests() + assert lane.route_counters() == { + "route_hits": 0, + "route_sites": {}, + "route_declines": {}, + "request_declines": 0, + "request_decline_extremes": {}, + "short_context_tokens": 2052, + } + + +# --------------------------------------------------------------------------- +# The driver's engagement check +# --------------------------------------------------------------------------- + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/tests/test_qsa_sparse_gqa_native.py b/tests/test_qsa_sparse_gqa_native.py new file mode 100644 index 000000000..4a6b6596c --- /dev/null +++ b/tests/test_qsa_sparse_gqa_native.py @@ -0,0 +1,324 @@ +"""CPU-only gates for the native direct-index sparse-GQA QSA kernel. + +Two things are covered, neither of which dispatches Metal: + +1. ``mtplx.native.qsa_sparse_gqa_unsupported_reason`` -- the shape/dtype/ + position contract. MLX arrays are constructed but never evaluated, and + the reason function only reads metadata, so no kernel is compiled or run. +2. the microbenchmark's scoring logic -- the + visible-set identity checker and the bf16-ULP tolerance measure. Those + are pure host arithmetic and are exercised on numpy inputs, including the + failure cases, because the whole parity story rests on them. + +Numeric Metal parity belongs to the operator-controlled guarded window; see +the harness docstring for the command. +""" + +from __future__ import annotations + +import importlib.util +import os +from pathlib import Path + +import numpy as np +import pytest + +ROOT = Path(__file__).resolve().parents[1] + +TOTAL = 16_384 +CAPACITY = 32_768 +ROWS = 64 +TOP_K = 512 +RATIO = 4 +SCALE = 256 ** -0.5 + + + + + + +# --------------------------------------------------------------------------- +# 1. binding validation +# --------------------------------------------------------------------------- +@pytest.fixture(scope="module") +def native(): + mx = pytest.importorskip("mlx.core") + native = pytest.importorskip("mtplx.native") + if not native.native_qsa_available(): + pytest.skip("the native QSA extension is not built on this host") + return mx, native + + +def _inputs(mx, *, rows=ROWS, ids_dtype=None, q_dtype=None): + ids_dtype = ids_dtype or mx.int32 + q_dtype = q_dtype or mx.bfloat16 + return { + "queries": mx.zeros((1, 24, rows, 256), q_dtype), + "keys": mx.zeros((1, 2, CAPACITY, 256), q_dtype), + "values": mx.zeros((1, 2, CAPACITY, 256), q_dtype), + "block_ids": mx.zeros((rows, TOP_K), ids_dtype), + } + + +def _reason(native_mod, arrays, **kwargs): + params = { + "pos_start": TOTAL - ROWS, + "total_tokens": TOTAL, + "scale": SCALE, + } + params.update(kwargs) + return native_mod.qsa_sparse_gqa_unsupported_reason( + arrays["queries"], + arrays["keys"], + arrays["values"], + arrays["block_ids"], + **params, + ) + + +def test_production_contract_is_accepted(native): + mx, native_mod = native + assert _reason(native_mod, _inputs(mx)) is None + + +def test_four_dimensional_block_ids_are_accepted(native): + mx, native_mod = native + arrays = _inputs(mx) + arrays["block_ids"] = arrays["block_ids"].reshape(1, 1, ROWS, TOP_K) + assert _reason(native_mod, arrays) is None + + +def test_uint32_block_ids_are_accepted(native): + """The metallib instantiates int32 too, so no astype is forced on the lane.""" + + mx, native_mod = native + assert _reason(native_mod, _inputs(mx, ids_dtype=mx.uint32)) is None + + +@pytest.mark.parametrize( + "tile,expected_ok", + [((128, 32), True), ((256, 32), True), ((64, 64), True), ((128, 64), True), + ((32, 32), False), ((128, 128), False), ((64, 32), False)], +) +def test_only_instantiated_tiles_are_accepted(native, tile, expected_ok): + mx, native_mod = native + reason = _reason( + native_mod, _inputs(mx), key_tile=tile[0], dimension_tile=tile[1] + ) + assert (reason is None) is expected_ok + + +def test_float32_queries_are_refused(native): + mx, native_mod = native + reason = _reason(native_mod, _inputs(mx, q_dtype=mx.float32)) + assert reason is not None and "float16 or bfloat16" in reason + + +def test_mixed_dtypes_are_refused(native): + mx, native_mod = native + arrays = _inputs(mx) + arrays["keys"] = mx.zeros((1, 2, CAPACITY, 256), mx.float16) + reason = _reason(native_mod, arrays) + assert reason is not None and "dtypes must match" in reason + + +def test_wrong_head_count_is_refused(native): + mx, native_mod = native + arrays = _inputs(mx) + arrays["queries"] = mx.zeros((1, 16, ROWS, 256), mx.bfloat16) + reason = _reason(native_mod, arrays) + assert reason is not None and "[1, 24, S, 256]" in reason + + +def test_block_id_width_must_be_the_budget(native): + mx, native_mod = native + arrays = _inputs(mx) + arrays["block_ids"] = mx.zeros((ROWS, 256), mx.int32) + reason = _reason(native_mod, arrays) + assert reason is not None and "[S, 512]" in reason + + +def test_block_id_rows_must_match_queries(native): + mx, native_mod = native + arrays = _inputs(mx) + arrays["block_ids"] = mx.zeros((ROWS + 1, TOP_K), mx.int32) + reason = _reason(native_mod, arrays) + assert reason is not None and "[S, 512]" in reason + + +def test_float_block_ids_are_refused(native): + mx, native_mod = native + arrays = _inputs(mx) + arrays["block_ids"] = mx.zeros((ROWS, TOP_K), mx.float32) + reason = _reason(native_mod, arrays) + assert reason is not None and "int32 or uint32" in reason + + +def test_sub_crossover_context_is_refused(native): + """2,048 tokens is exactly the boundary: 512 blocks is not yet sparse.""" + + mx, native_mod = native + reason = _reason( + native_mod, _inputs(mx), pos_start=2048 - ROWS, total_tokens=2048 + ) + assert reason is not None and "dense/sparse boundary" in reason + + +def test_context_beyond_the_backing_is_refused(native): + mx, native_mod = native + reason = _reason( + native_mod, + _inputs(mx), + pos_start=CAPACITY + 1, + total_tokens=CAPACITY + 1 + ROWS, + ) + assert reason is not None and "backing capacity" in reason + + +def test_non_causal_suffix_is_refused(native): + mx, native_mod = native + reason = _reason(native_mod, _inputs(mx), pos_start=TOTAL, total_tokens=TOTAL) + assert reason is not None and "causal suffix" in reason + + +def test_traced_scalars_are_refused(native): + mx, native_mod = native + reason = _reason(native_mod, _inputs(mx), pos_start=mx.array(0)) + assert reason is not None and "host integers" in reason + + +def test_bool_positions_are_refused(native): + mx, native_mod = native + reason = _reason(native_mod, _inputs(mx), pos_start=True) + assert reason is not None and "cannot be bool" in reason + + +def test_supported_predicate_matches_the_reason(native): + mx, native_mod = native + arrays = _inputs(mx) + assert native_mod.qsa_sparse_gqa_supported( + arrays["queries"], + arrays["keys"], + arrays["values"], + arrays["block_ids"], + pos_start=TOTAL - ROWS, + total_tokens=TOTAL, + scale=SCALE, + ) + assert not native_mod.qsa_sparse_gqa_supported( + arrays["queries"], + arrays["keys"], + arrays["values"], + arrays["block_ids"], + pos_start=TOTAL - ROWS, + total_tokens=TOTAL, + scale=SCALE, + key_tile=32, + dimension_tile=32, + ) + + +def test_a_refused_call_raises_rather_than_dispatching(native): + mx, native_mod = native + arrays = _inputs(mx) + with pytest.raises(ValueError, match="dense/sparse boundary"): + native_mod.qsa_sparse_gqa( + arrays["queries"], + arrays["keys"], + arrays["values"], + arrays["block_ids"], + pos_start=2048 - ROWS, + total_tokens=2048, + scale=SCALE, + ) + + +# --------------------------------------------------------------------------- +# 2. harness scoring logic (pure numpy) +# --------------------------------------------------------------------------- +def _selector_output(pos_start: int, rows: int, rng: np.random.Generator): + """Reproduce what ``_select_eager`` guarantees: ascending valid prefix.""" + + ids = np.zeros((rows, TOP_K), dtype=np.int64) + ok = np.zeros((rows, TOP_K), dtype=bool) + for r in range(rows): + complete = (pos_start + r + 1) // RATIO + n = min(TOP_K, complete) + if n: + ids[r, :n] = np.sort(rng.choice(complete, size=n, replace=False)) + ok[r, :n] = True + return ids, ok + + +def _cell(pos_start: int, rows: int, seed: int = 7): + rng = np.random.default_rng(seed) + ids, ok = _selector_output(pos_start, rows, rng) + return { + "rows": rows, + "pos_start": pos_start, + "block_ids": ids, + "block_valid": ok, + } + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +# --------------------------------------------------------------------------- +# 3. the seam with the `flash_prefill` selector branch +# --------------------------------------------------------------------------- +# Both lanes read one `_select_eager`. This kernel's correctness rests on the +# contract the `flash_prefill` branch builds out of `top_idx` (ascending ids, +# prefix validity, count `min(512, (pos+1)//4)`), and that branch never serves +# a single row, so pin both sides. +def test_the_flash_prefill_selector_gate_never_admits_a_single_row(): + qwen4_exp = pytest.importorskip("mtplx.models.qwen4_exp") + assert qwen4_exp._qsa_prefill_min_rows() >= 2 + + +def test_the_flash_prefill_branch_still_builds_the_prefix_contract(): + """The contract must survive an edit to the scoring branches around it. + + Reading the source is deliberate: the branch's value is a GPU expression, + but the three operations that establish the invariant (sort ascending, + gather validity along the SORTED ids, blank invalid ids) are structural + and a future edit that drops one would be silent wrong attention. + """ + + import inspect + + qwen4_exp = pytest.importorskip("mtplx.models.qwen4_exp") + source = inspect.getsource(qwen4_exp.QSAIndexer._select_eager) + branch = source[source.index("_qsa_large_prefill_enabled(S, total)") :] + branch = branch[: branch.index('return ("flash_prefill"')] + assert "mx.sort(top_idx" in branch # chronological + assert "mx.take_along_axis(" in branch # validity gathered by SORTED id + assert "block_ids.astype(mx.int64)" in branch + assert "mx.where(" in branch # invalid slots blanked, flag kept separately diff --git a/tests/test_qwen4_prefill_mask_fuse.py b/tests/test_qwen4_prefill_mask_fuse.py index 68448e764..3cf41d84a 100644 --- a/tests/test_qwen4_prefill_mask_fuse.py +++ b/tests/test_qwen4_prefill_mask_fuse.py @@ -513,7 +513,14 @@ def test_refusal_is_native_on_a_build_without_fused_kernels(monkeypatch): q, kv, kv, mask=_lane_mask(total - rows, rows, total), scale=1.0 ) message = err.getvalue() - assert "require a GPU (Metal) stream" in message + # The refusal is NATIVE: MLX declined force_fused (a CPU stream has no + # fused kernel), the lane caught it and logged one per-class line with + # MLX's own message appended. The exact native wording differs by MLX + # build (0.32.0 said one thing, 0.32.2's CPU stream says "require a GPU + # (Metal) stream"), so assert the version-independent lane line rather + # than any one build's native string. + assert "has no fused SDPA for shape class" in message + assert "MTPLX_QWEN4_PREFILL_MASK_FUSE" in message # Two classes learned -- one per mask kind at this one geometry -- and # nothing said about any other shape. assert set(qwen4_exp._PREFILL_MASK_FUSE_UNAVAILABLE) == { From 64e81643341a0c75c62a82b86e5d704a9f147e12 Mon Sep 17 00:00:00 2001 From: davidtai Date: Mon, 7 Sep 2026 08:55:57 -0500 Subject: [PATCH 03/30] fix(qwen4): arm MTPLX_QSA_SPARSE_DECODE at install time, not import Served via `mtplx serve` (cli-resolved Turbo, no lane flags), the QSA split-K decode lane did not engage: qsa_sparse_decode_enabled() read the environment at IMPORT and cached the default (False), but the fixed-M4 auto-arm stamps MTPLX_QSA_SPARSE_DECODE (native-gated) into the environment AFTER runtime_options is imported, so the cache froze the default before the stamp landed -- the lane was absent from /health with neither an "armed:" nor a "declined to stock" line. hc_m4 escaped only because its reader is read on a path where the module was imported after the stamp. Resolve the flag lazily on the FIRST read (which is the graphbank cache install, after the overrides are applied), then cache; the _QSA_SPARSE_DECODE module global stays (tests force it to a bool) and the native-gated default in the server auto-arm is unchanged. The env is frozen once serving starts, so a lazy first read is still a single cached bool on the hot path. Regression tests, the shape that would have caught this: - the reader picks up a stamp applied AFTER import (an import-frozen reader fails it), - the fixed-M4 auto-arm block stamps the lane when the native extension is built, or prints the declined-to-stock verdict and leaves it unstamped when it is not. --- mtplx/runtime_options.py | 29 ++++++-- tests/test_qsa_sparse_decode_wiring.py | 95 ++++++++++++++++++++++++++ 2 files changed, 118 insertions(+), 6 deletions(-) diff --git a/mtplx/runtime_options.py b/mtplx/runtime_options.py index 52b414619..dbc7539e4 100644 --- a/mtplx/runtime_options.py +++ b/mtplx/runtime_options.py @@ -268,27 +268,44 @@ def qwen4_hc_m4_enabled() -> bool: #: failure at install: this kernel is rounding-class, so a parity miss is a #: numerical verdict, and the lane disables itself for the process and #: reports the measured deltas. -def _qsa_sparse_decode_import_default() -> bool: +def _resolve_qsa_sparse_decode() -> bool: # New key wins for any non-empty value (including "0" for the per-key # opt-out); the old MTPLX_FABLE_QSA_SPARSE_DECODE name is honoured as an - # alias only when the new key is unset. Read once at import. + # alias only when the new key is unset. raw = os.environ.get("MTPLX_QSA_SPARSE_DECODE") if raw is None or not str(raw).strip(): return env_bool("MTPLX_FABLE_QSA_SPARSE_DECODE", default=False) return env_bool("MTPLX_QSA_SPARSE_DECODE", default=False) -_QSA_SPARSE_DECODE = _qsa_sparse_decode_import_default() +#: ``None`` = resolve from the environment on every read; a test may force a +#: bool (the fixtures set this directly to arm/disarm the lane). +#: +#: The flag is read at USE, never frozen at import. The server's fixed-M4 +#: auto-arm stamps ``MTPLX_QSA_SPARSE_DECODE`` into the environment AFTER this +#: module is imported (``openai.py:_server_runtime_env_overrides``, gated on +#: the built native extension). An import-time read (or a first-use read that +#: happened to fire before the stamp) froze the default (False) before the +#: stamp landed, so the served lane never engaged even with the extension +#: built -- the bug the battery caught 2026-09-07. Reading the environment on +#: each call means the graphbank cache install (after the overrides are +#: applied) always sees the resolved value; the read is a dict lookup and the +#: env is frozen once serving starts. +_QSA_SPARSE_DECODE = None def qsa_sparse_decode_enabled() -> bool: - """True when the QSA split-K decode flag armed this process at import. + """True when the QSA split-K decode flag is armed for this process. Armed by ``MTPLX_QSA_SPARSE_DECODE`` (or the old - ``MTPLX_FABLE_QSA_SPARSE_DECODE`` alias). + ``MTPLX_FABLE_QSA_SPARSE_DECODE`` alias). Read at use, not frozen at + import -- see the note on :data:`_QSA_SPARSE_DECODE`. A test may set that + global to a bool to force the answer. """ - return _QSA_SPARSE_DECODE + if _QSA_SPARSE_DECODE is not None: + return bool(_QSA_SPARSE_DECODE) + return _resolve_qsa_sparse_decode() def _parse_sparse_decode_tile(raw: str | None) -> tuple[int, int]: diff --git a/tests/test_qsa_sparse_decode_wiring.py b/tests/test_qsa_sparse_decode_wiring.py index 4fe724296..bddd86461 100644 --- a/tests/test_qsa_sparse_decode_wiring.py +++ b/tests/test_qsa_sparse_decode_wiring.py @@ -769,3 +769,98 @@ def test_reset_clears_the_route_state_too(): + + +# --------------------------------------------------------------------------- +# Regression: the served-launch arming ordering (battery, 2026-09-07) +# --------------------------------------------------------------------------- +def test_the_reader_resolves_lazily_not_at_import(monkeypatch): + """A stamp applied AFTER runtime_options is imported must still arm. + + The served bug: qsa_sparse_decode_enabled() froze the env at IMPORT + (default False), so openai.py's fixed-M4 auto-arm setdefault -- which runs + later and stamps MTPLX_QSA_SPARSE_DECODE=1 when the native ext is built -- + landed after the cache and the lane never engaged. The reader must read at + first USE (install time), which is after the overrides are applied. + """ + + from mtplx import runtime_options as ro + + monkeypatch.delenv("MTPLX_QSA_SPARSE_DECODE", raising=False) + monkeypatch.delenv("MTPLX_FABLE_QSA_SPARSE_DECODE", raising=False) + # As at a fresh import, with nothing stamped yet. + monkeypatch.setattr(ro, "_QSA_SPARSE_DECODE", None) + assert ro.qsa_sparse_decode_enabled() is False + # The auto-arm stamps the key AFTER import; a lazily-resolved reader that + # has not yet been forced to a value must pick it up. + monkeypatch.setattr(ro, "_QSA_SPARSE_DECODE", None) + monkeypatch.setenv("MTPLX_QSA_SPARSE_DECODE", "1") + assert ro.qsa_sparse_decode_enabled() is True + # The old FABLE name is honoured only when the new key is unset. + monkeypatch.setattr(ro, "_QSA_SPARSE_DECODE", None) + monkeypatch.delenv("MTPLX_QSA_SPARSE_DECODE", raising=False) + monkeypatch.setenv("MTPLX_FABLE_QSA_SPARSE_DECODE", "1") + assert ro.qsa_sparse_decode_enabled() is True + + +def test_the_fixed_m4_auto_arm_stamps_or_declines_the_sparse_decode_lane( + tmp_path, monkeypatch +): + """Apply the fixed-M4 auto-arm block; the sparse-decode install must be + ATTEMPTED -- stamped (native ext built) or declined-to-stock (not built) -- + and never silently absent, which is exactly what the served launch showed. + """ + + import io + import contextlib + import json as _json + from types import SimpleNamespace as _NS + + import mtplx.server.openai as openai + from mtplx import runtime_options as ro + from mtplx.native import native_qsa_available + from mtplx.profiles import ( + MODEL_RUNTIME_ENV_OVERRIDE_KEYS, + normalize_runtime_env_overrides, + ) + + (tmp_path / "config.json").write_text( + _json.dumps({"model_type": "qwen4_exp"}), encoding="utf-8" + ) + args = _NS( + generation_mode="mtp", + verify_strategy="capture_commit", + model=str(tmp_path), + ) + # Force the fixed-M4 predicate on rather than crafting a full fixed-verify + # config; the sparse-decode default gates only on that predicate + the + # built native extension. + monkeypatch.setattr(openai, "_served_model_is_qwen4_fixed_m4", lambda a: True) + for key in ("MTPLX_QSA_SPARSE_DECODE", "MTPLX_FABLE_QSA_SPARSE_DECODE"): + monkeypatch.delenv(key, raising=False) + + err = io.StringIO() + with contextlib.redirect_stderr(err): + overrides = openai._server_runtime_env_overrides(args, {}) + log = err.getvalue() + + # The key is registered, so a stamped value survives the boot-time + # validator (the check that only runs inside apply_profile_env). + assert "MTPLX_QSA_SPARSE_DECODE" in MODEL_RUNTIME_ENV_OVERRIDE_KEYS + assert normalize_runtime_env_overrides(overrides) == overrides + + if native_qsa_available(): + # Ext built: the lane is armed by default -> install attempted. + assert overrides.get("MTPLX_QSA_SPARSE_DECODE") == "1" + # And the reader, resolved at install time AFTER the stamp is applied + # to the environment, arms (the bug returned False here). + monkeypatch.setattr(ro, "_QSA_SPARSE_DECODE", None) + monkeypatch.setenv( + "MTPLX_QSA_SPARSE_DECODE", overrides["MTPLX_QSA_SPARSE_DECODE"] + ) + assert ro.qsa_sparse_decode_enabled() is True + else: + # No ext: decline-to-stock, key NOT stamped, a verdict line printed so + # a wheel without the extension still serves and says so. + assert "MTPLX_QSA_SPARSE_DECODE" not in overrides + assert "MTPLX_QSA_SPARSE_DECODE declined to stock" in log From 16997fd9269f276e46b3b88a6e62d42e8a804d0b Mon Sep 17 00:00:00 2001 From: davidtai Date: Mon, 7 Sep 2026 09:13:21 -0500 Subject: [PATCH 04/30] fix(qwen4): resolve every remainder-lane flag at use, not at import Second arming failure (battery, 2026-09-07): served as `mtplx serve` launches it, hc_m4 was OFF for the same reason the QSA decode lane was in commit 3 -- its reader froze the environment at import (default off) while the fixed-M4 auto-arm stamps the lane keys into the environment AFTER runtime_options is imported. The earlier claim that hc_m4 read on a post-stamp path did not hold for the served path. Resolve every remainder-lane flag at USE (the install / route path, which runs after the overrides are applied), never at import: - runtime_options: qwen4_hc_m4_enabled, qsa_sparse_decode_tile and qsa_sparse_decode_splits (qsa_sparse_decode_enabled was fixed in commit 3); each keeps its module global as a test override (None = read env). - models/qwen4_exp: _prefill_mask_fuse_enabled drops @lru_cache (its body already reads os.environ), so a stamp landing after import is seen. - qwen4_prefill_chunk.resolve_query_tile_rows already read at use. The native-gated default and the MTPLX_FABLE_* aliases are unchanged. Upstream's own MTPLX_QWEN4_OPDIET / MTPLX_QWEN4_VERIFY_GLUE readers are left as-is (not part of this remainder set). Test reproducing the served order (tests/test_qwen4_remainder_arming.py): import mtplx.runtime + mtplx.server.openai FIRST, assert all four readers off, run _server_runtime_env_overrides for the fixed-M4 pack, apply it to os.environ, then assert all four arm -- the decode lane armed when the native extension is built, else an explicit declined-to-stock verdict, never silent absence. The hc_m4 read-once test is rewritten to assert read-at-use, and the mask-fuse test drops its now-defunct cache_clear() calls. --- mtplx/models/qwen4_exp.py | 14 ++-- mtplx/runtime_options.py | 77 ++++++++++++------ tests/test_qwen4_hc_m4.py | 26 ++++-- tests/test_qwen4_prefill_mask_fuse.py | 4 - tests/test_qwen4_remainder_arming.py | 113 ++++++++++++++++++++++++++ 5 files changed, 191 insertions(+), 43 deletions(-) create mode 100644 tests/test_qwen4_remainder_arming.py diff --git a/mtplx/models/qwen4_exp.py b/mtplx/models/qwen4_exp.py index 14b2a95c7..fe5ac98c6 100644 --- a/mtplx/models/qwen4_exp.py +++ b/mtplx/models/qwen4_exp.py @@ -44,7 +44,6 @@ import struct import time from dataclasses import dataclass, field -from functools import lru_cache from pathlib import Path from typing import Any, Dict, List, Optional, Union @@ -1921,7 +1920,6 @@ def _qsa_prefill_gather_tile_rows() -> int: return 64 -@lru_cache(maxsize=1) def _prefill_mask_fuse_enabled() -> bool: """MTPLX_QWEN4_PREFILL_MASK_FUSE: ask MLX for the fused masked SDPA. @@ -1980,12 +1978,14 @@ def _prefill_mask_fuse_enabled() -> bool: ``engaged:`` line -- the absence of a refusal is not a receipt. """ - # Cached (maxsize=1) so the hot lane does not touch os.environ every - # call; ``cache_clear()`` is the escape hatch a test or an in-process A/B - # uses to re-read. The renamed key wins when set to any non-empty value + # Read at USE, never frozen at import: the server's fixed-M4 auto-arm + # stamps this key into the environment after the module is imported, so an + # import-time freeze would miss the stamp (the arming-ordering class the + # battery caught). The read is one dict lookup, and it is only reached on + # the prefill path. The renamed key wins when set to any non-empty value # (including "0" for the per-key opt-out); the old - # MTPLX_FABLE_PREFILL_MASK_FUSE name is honoured as an alias only when - # the new key is unset. + # MTPLX_FABLE_PREFILL_MASK_FUSE name is honoured as an alias only when the + # new key is unset. raw = os.environ.get("MTPLX_QWEN4_PREFILL_MASK_FUSE") if raw is None or not str(raw).strip(): raw = os.environ.get("MTPLX_FABLE_PREFILL_MASK_FUSE") diff --git a/mtplx/runtime_options.py b/mtplx/runtime_options.py index dbc7539e4..047b5dfe0 100644 --- a/mtplx/runtime_options.py +++ b/mtplx/runtime_options.py @@ -225,30 +225,40 @@ def reset_qwen4_verify_glue_cache(env: Mapping[str, str] | None = None) -> None: #: Verify-width fused hyper-connection read (mtplx/kernels/qwen4_m4_hyper_read). #: -#: Read ONCE at import so the hot verify path never touches ``os.environ`` and -#: two traces of the same compiled graph cannot disagree about which chain they -#: carry. ``MTPLX_QWEN4_HC_M4`` is the key; the old ``MTPLX_FABLE_HC_M4`` name -#: is honoured as an alias only when the new key is unset (the new key wins for -#: any non-empty value, including ``0`` for the per-key opt-out). The kernel -#: RAISES on a family-contract miss rather than falling back, so an -#: armed-but-inert lane is unreachable. -def _qwen4_hc_m4_import_default() -> bool: +#: Read at USE, never frozen at import. ``MTPLX_QWEN4_HC_M4`` is the key; the +#: old ``MTPLX_FABLE_HC_M4`` name is honoured as an alias only when the new key +#: is unset (the new key wins for any non-empty value, including ``0`` for the +#: per-key opt-out). The kernel RAISES on a family-contract miss rather than +#: falling back, so an armed-but-inert lane is unreachable. +#: +#: ``None`` = resolve from the environment on each read; a test may force a +#: bool. This must NOT freeze at import: the server's fixed-M4 auto-arm stamps +#: ``MTPLX_QWEN4_HC_M4`` into the environment AFTER this module is imported, so +#: an import-time read froze the default (False) and the served lane never +#: armed -- the second arming failure the battery caught 2026-09-07 (the first +#: was the sibling QSA sparse-decode reader). The install check runs after the +#: overrides are applied, so reading the environment there sees the stamp. +_QWEN4_HC_M4 = None + + +def _resolve_qwen4_hc_m4() -> bool: raw = os.environ.get("MTPLX_QWEN4_HC_M4") if raw is None or not str(raw).strip(): return env_bool("MTPLX_FABLE_HC_M4", default=False) return env_bool("MTPLX_QWEN4_HC_M4", default=False) -_QWEN4_HC_M4 = _qwen4_hc_m4_import_default() - - def qwen4_hc_m4_enabled() -> bool: - """True when the HC_M4 flag armed this process at import. + """True when the HC_M4 flag is armed for this process. Armed by ``MTPLX_QWEN4_HC_M4`` (or the old ``MTPLX_FABLE_HC_M4`` alias). + Read at use, not frozen at import; a test may set :data:`_QWEN4_HC_M4` to + a bool to force the answer. """ - return _QWEN4_HC_M4 + if _QWEN4_HC_M4 is not None: + return bool(_QWEN4_HC_M4) + return _resolve_qwen4_hc_m4() @@ -339,10 +349,10 @@ def _parse_sparse_decode_tile(raw: str | None) -> tuple[int, int]: QSA_SPARSE_DECODE_TILES = ((128, 32), (256, 32), (64, 64), (128, 64)) QSA_SPARSE_DECODE_MAX_SPLITS = 64 -_QSA_SPARSE_DECODE_TILE = _parse_sparse_decode_tile( - os.environ.get("MTPLX_QSA_SPARSE_DECODE_TILE") - or os.environ.get("MTPLX_FABLE_QSA_SPARSE_DECODE_TILE") -) +#: ``None`` = resolve from the environment on each read; a test may force a +#: (key_tile, dim_tile) tuple. Read at use, not frozen at import (same +#: server-arming ordering as the master flag above). +_QSA_SPARSE_DECODE_TILE = None #: MEASURED default (2026-09-02, guarded micro, M=4, 16K, 12 layers). The @@ -385,22 +395,39 @@ def _parse_sparse_decode_splits(raw: str | None) -> int: return value -_QSA_SPARSE_DECODE_SPLITS = _parse_sparse_decode_splits( - os.environ.get("MTPLX_QSA_SPARSE_DECODE_SPLITS") - or os.environ.get("MTPLX_FABLE_QSA_SPARSE_DECODE_SPLITS") -) +#: ``None`` = resolve from the environment on each read; a test may force an +#: int. Read at use, not frozen at import. +_QSA_SPARSE_DECODE_SPLITS = None def qsa_sparse_decode_tile() -> tuple[int, int]: - """The armed ``(key_tile, dimension_tile)`` for the decode kernel.""" + """The armed ``(key_tile, dimension_tile)`` for the decode kernel. - return _QSA_SPARSE_DECODE_TILE + Read at use, not frozen at import; a test may set + :data:`_QSA_SPARSE_DECODE_TILE` to force the answer. + """ + + if _QSA_SPARSE_DECODE_TILE is not None: + return _QSA_SPARSE_DECODE_TILE + return _parse_sparse_decode_tile( + os.environ.get("MTPLX_QSA_SPARSE_DECODE_TILE") + or os.environ.get("MTPLX_FABLE_QSA_SPARSE_DECODE_TILE") + ) def qsa_sparse_decode_splits() -> int: - """The armed KV-split target for the decode kernel.""" + """The armed KV-split target for the decode kernel. - return _QSA_SPARSE_DECODE_SPLITS + Read at use, not frozen at import; a test may set + :data:`_QSA_SPARSE_DECODE_SPLITS` to force the answer. + """ + + if _QSA_SPARSE_DECODE_SPLITS is not None: + return _QSA_SPARSE_DECODE_SPLITS + return _parse_sparse_decode_splits( + os.environ.get("MTPLX_QSA_SPARSE_DECODE_SPLITS") + or os.environ.get("MTPLX_FABLE_QSA_SPARSE_DECODE_SPLITS") + ) @dataclass(frozen=True) diff --git a/tests/test_qwen4_hc_m4.py b/tests/test_qwen4_hc_m4.py index 0607525c6..39e7d89e5 100644 --- a/tests/test_qwen4_hc_m4.py +++ b/tests/test_qwen4_hc_m4.py @@ -347,16 +347,28 @@ def test_armed_wrong_down_rows_raise(armed): mod._hc_m4_applies(mx.zeros((1, 4, HCD), dtype=mx.bfloat16)) -def test_env_flag_is_read_once(monkeypatch): - """A mid-run env change must not reach the hot path: two traces of the - same compiled verify graph would then disagree about which read they - contain.""" +def test_env_flag_is_read_at_use_not_frozen_at_import(monkeypatch): + """The reader must resolve the environment at USE, not freeze at import. + + The server's fixed-M4 auto-arm stamps MTPLX_QWEN4_HC_M4 AFTER + runtime_options is imported; an import-time freeze left the served lane + OFF (the battery's 2026-09-07 arming failure). A stamp applied after the + module exists must be seen. ``_QWEN4_HC_M4`` is left None (unforced) so + the reader consults the environment. + """ - before = runtime_options.qwen4_hc_m4_enabled() + monkeypatch.setattr(runtime_options, "_QWEN4_HC_M4", None) + monkeypatch.delenv("MTPLX_QWEN4_HC_M4", raising=False) + monkeypatch.delenv("MTPLX_FABLE_HC_M4", raising=False) + assert runtime_options.qwen4_hc_m4_enabled() is False monkeypatch.setenv("MTPLX_QWEN4_HC_M4", "1") - assert runtime_options.qwen4_hc_m4_enabled() is before + assert runtime_options.qwen4_hc_m4_enabled() is True + monkeypatch.delenv("MTPLX_QWEN4_HC_M4", raising=False) + assert runtime_options.qwen4_hc_m4_enabled() is False + # A test/A-B may still force a value by setting the module global. + monkeypatch.setattr(runtime_options, "_QWEN4_HC_M4", True) monkeypatch.delenv("MTPLX_QWEN4_HC_M4", raising=False) - assert runtime_options.qwen4_hc_m4_enabled() is before + assert runtime_options.qwen4_hc_m4_enabled() is True def test_env_flag_defaults_off(): diff --git a/tests/test_qwen4_prefill_mask_fuse.py b/tests/test_qwen4_prefill_mask_fuse.py index 3cf41d84a..f2498c9f4 100644 --- a/tests/test_qwen4_prefill_mask_fuse.py +++ b/tests/test_qwen4_prefill_mask_fuse.py @@ -71,7 +71,6 @@ def _clean_lane_state(monkeypatch): monkeypatch.delenv(MASK_FUSE_ENV, raising=False) monkeypatch.delenv(QUERY_TILE_ENV, raising=False) - qwen4_exp._prefill_mask_fuse_enabled.cache_clear() saved_counts = dict(qwen4_exp._QSA_PREFILL_COUNTS) saved_unavailable = dict(qwen4_exp._PREFILL_MASK_FUSE_UNAVAILABLE) saved_printed = qwen4_exp._MASK_FUSE_REFUSALS_PRINTED[0] @@ -83,7 +82,6 @@ def _clean_lane_state(monkeypatch): try: yield finally: - qwen4_exp._prefill_mask_fuse_enabled.cache_clear() qwen4_exp._QSA_PREFILL_COUNTS.clear() qwen4_exp._QSA_PREFILL_COUNTS.update(saved_counts) qwen4_exp._PREFILL_MASK_FUSE_UNAVAILABLE.clear() @@ -94,7 +92,6 @@ def _clean_lane_state(monkeypatch): def _arm(monkeypatch, value: str = "1") -> None: monkeypatch.setenv(MASK_FUSE_ENV, value) - qwen4_exp._prefill_mask_fuse_enabled.cache_clear() def _lane_mask(pos_start: int, rows: int, total: int) -> mx.array: @@ -717,7 +714,6 @@ def test_armed_flag_on_an_unavailable_build_still_returns_the_dense_answer( armed = qwen4_exp._qsa_dense_attention( q, kv, kv, mask=_lane_mask(pos_start, rows, total), scale=0.25 ) - qwen4_exp._prefill_mask_fuse_enabled.cache_clear() monkeypatch.delenv(MASK_FUSE_ENV, raising=False) stock = qwen4_exp._qsa_dense_attention( q, kv, kv, mask=_lane_mask(pos_start, rows, total), scale=0.25 diff --git a/tests/test_qwen4_remainder_arming.py b/tests/test_qwen4_remainder_arming.py new file mode 100644 index 000000000..ca62a972c --- /dev/null +++ b/tests/test_qwen4_remainder_arming.py @@ -0,0 +1,113 @@ +"""The PR 391 remainder lanes must arm when SERVED, not freeze at import. + +This reproduces the exact ``mtplx serve`` order that hid two arming failures +the battery caught (2026-09-07): the reader modules are imported BEFORE the +fixed-M4 auto-arm stamps the lane keys into the environment, so an +import-time-frozen reader returns its default (off) forever -- the served lane +is then absent from /health with no verdict line. Every remainder-lane reader +must resolve the environment at USE (the install/route path, which runs after +the overrides are applied), so a stamp that lands after import is still seen. + +Nothing here dispatches Metal; it drives the pure-Python auto-arm block and the +env readers. +""" + +from __future__ import annotations + +import contextlib +import io +import json +from types import SimpleNamespace + +# Import the server + runtime modules FIRST, exactly as `mtplx serve` does, so +# any import-time reader cache would be populated with the (unset) defaults +# before the auto-arm runs. If a reader froze here, the asserts below fail. +import mtplx.runtime # noqa: F401 (import-order fidelity) +import mtplx.server.openai as openai +from mtplx import runtime_options as ro +from mtplx.models import qwen4_exp +from mtplx.native import native_qsa_available +from mtplx.profiles import normalize_runtime_env_overrides +from mtplx.qwen4_prefill_chunk import resolve_query_tile_rows + +_LANE_ENV_KEYS = ( + "MTPLX_QWEN4_HC_M4", + "MTPLX_QWEN4_PREFILL_MASK_FUSE", + "MTPLX_QSA_PREFILL_QUERY_TILE", + "MTPLX_QSA_SPARSE_DECODE", + "MTPLX_QSA_SPARSE_DECODE_TILE", + "MTPLX_QSA_SPARSE_DECODE_SPLITS", + "MTPLX_FABLE_HC_M4", + "MTPLX_FABLE_PREFILL_MASK_FUSE", + "MTPLX_FABLE_PREFILL_QSA_QUERY_TILE", + "MTPLX_FABLE_QSA_SPARSE_DECODE", +) + + +def test_remainder_lanes_arm_when_served_not_frozen_at_import(tmp_path, monkeypatch): + # The state at a fresh import: reader globals unforced (None -> read env), + # every lane key unset. + for name in ( + "_QWEN4_HC_M4", + "_QSA_SPARSE_DECODE", + "_QSA_SPARSE_DECODE_TILE", + "_QSA_SPARSE_DECODE_SPLITS", + ): + monkeypatch.setattr(ro, name, None) + for key in _LANE_ENV_KEYS: + monkeypatch.delenv(key, raising=False) + + # Before the auto-arm: every reader off, exactly the served import-time + # state. (A reader frozen True/False at import would already fail here.) + assert ro.qwen4_hc_m4_enabled() is False + assert qwen4_exp._prefill_mask_fuse_enabled() is False + assert resolve_query_tile_rows() == 0 + assert ro.qsa_sparse_decode_enabled() is False + + # Run the fixed-M4 auto-arm and APPLY it to the environment, as the server + # does via apply_profile_env. Force the fixed-M4 predicate on rather than + # crafting a full fixed-verify config; the remainder defaults gate on that + # predicate (and, for the decode lane, the built native extension). + (tmp_path / "config.json").write_text( + json.dumps({"model_type": "qwen4_exp"}), encoding="utf-8" + ) + args = SimpleNamespace( + generation_mode="mtp", + verify_strategy="capture_commit", + model=str(tmp_path), + ) + monkeypatch.setattr(openai, "_served_model_is_qwen4_fixed_m4", lambda a: True) + err = io.StringIO() + with contextlib.redirect_stderr(err): + overrides = openai._server_runtime_env_overrides(args, {}) + log = err.getvalue() + + # Every stamped key survives the boot-time validator (the check that only + # runs inside apply_profile_env), then apply them as the server would. + assert normalize_runtime_env_overrides(overrides) == overrides + for key, value in overrides.items(): + monkeypatch.setenv(key, value) + + # The three unconditional remainder lanes are armed, and each reader -- + # read at USE, after the stamp was applied -- sees it. This is the class + # that failed when served: an import-frozen reader returns its default. + assert overrides.get("MTPLX_QWEN4_HC_M4") == "1" + assert ro.qwen4_hc_m4_enabled() is True + assert overrides.get("MTPLX_QWEN4_PREFILL_MASK_FUSE") == "1" + assert qwen4_exp._prefill_mask_fuse_enabled() is True + assert overrides.get("MTPLX_QSA_PREFILL_QUERY_TILE") == "2048" + assert resolve_query_tile_rows() == 2048 + + # The native-gated decode lane: armed (install attempted) when the + # extension is built, else an explicit declined-to-stock verdict so a + # wheel without it still serves -- never a silent absence. + if native_qsa_available(): + assert overrides.get("MTPLX_QSA_SPARSE_DECODE") == "1" + assert ro.qsa_sparse_decode_enabled() is True + # its companions resolve to the measured geometry at use, too + assert ro.qsa_sparse_decode_tile() == (128, 32) + assert ro.qsa_sparse_decode_splits() == 17 + else: + assert "MTPLX_QSA_SPARSE_DECODE" not in overrides + assert "MTPLX_QSA_SPARSE_DECODE declined to stock" in log + assert ro.qsa_sparse_decode_enabled() is False From 5dd05b7e6660dc6fa82cfcc06e1c5970402bc65d Mon Sep 17 00:00:00 2001 From: davidtai Date: Mon, 7 Sep 2026 11:43:36 -0500 Subject: [PATCH 05/30] fix(qwen4): resolve the four frozen-at-import verify-lane flags at use Arming audit (battery, 2026-09-07): `mtplx serve` imports generation / runtime / model modules before parse_args stamps the auto-arm env, so any flag whose reader resolves at module import freezes its default before the stamp lands and the auto-arm's setdefault is a silent no-op as launched. A column-0 scan of every module holding a stamped key's reader found exactly four such readers among the ~31 auto-armed keys; every other stamped key reads the environment at use or is consumed from config.json at model load. All four are decode-verify lanes, so the release control (71.17 tok/s at 16K) ran without them -- a plausible slice of the 71->81 gap. Resolve all four at use (read the environment each call; the module global stays a test/force override; the env is frozen once serving starts, so two traces of one graph still read the same value): - MTPLX_QWEN4_DRAFT_K20_PRESCATTER: qwen4_draft_k20_prescatter._ENABLED, and generation.py's cached _QWEN4_DRAFT_K20_PRESCATTER (removed; the one draft consult site calls the reader). - MTPLX_QWEN4_BLOCK_VERIFY: qwen4_block_verify._ENABLED, and generation.py's cached _QWEN4_BLOCK_VERIFY (removed; the accept-loop consult calls the reader). - MTPLX_QWEN4_OPDIET (+ _ITEMS): runtime_options. - MTPLX_QWEN4_VERIFY_GLUE (+ _ITEMS): runtime_options (reset hook kept, now forcing the globals). No default value changed; keys that already read at use are untouched. Upstream's STRICT_CLAIMS and BATCH_PAGED_OFFSETS are also import-frozen but are not auto-armed (operator sets them pre-launch), so they are left as-is. Tests: tests/test_qwen4_remainder_arming.py extended to assert all four arm in the served order (import first, stamp, read) and that the fixed-M4 auto-arm stamps OPDIET / BLOCK_VERIFY / VERIFY_GLUE and their readers then arm. The block-verify and draft-k20 source-inspection tests and the opdiet read-once test are rewritten to assert read-at-use. --- mtplx/generation.py | 44 ++++++------- mtplx/qwen4_block_verify.py | 23 ++++--- mtplx/qwen4_draft_k20_prescatter.py | 19 ++++-- mtplx/runtime_options.py | 78 ++++++++++++++++++------ tests/test_qwen4_block_verify.py | 17 ++++-- tests/test_qwen4_draft_k20_prescatter.py | 8 ++- tests/test_qwen4_opdiet.py | 13 ++-- tests/test_qwen4_remainder_arming.py | 78 ++++++++++++++++++++++++ 8 files changed, 217 insertions(+), 63 deletions(-) diff --git a/mtplx/generation.py b/mtplx/generation.py index 4281831b1..ac73a676d 100644 --- a/mtplx/generation.py +++ b/mtplx/generation.py @@ -308,11 +308,13 @@ def _env_falsey(name: str) -> bool: } -# MTPLX_QWEN4_DRAFT_K20_PRESCATTER -- read ONCE at import (in -# ``mtplx.qwen4_draft_k20_prescatter``), default OFF. When off this constant -# is False, no plan is claimed, `_draft_k20_prescatter_plan` stays None, and -# the one draft-read site below is behind `is not None`, so the retained stock -# lane runs the code it ran before this module existed. +# MTPLX_QWEN4_DRAFT_K20_PRESCATTER -- read AT USE via +# ``_qwen4_draft_k20_prescatter_enabled()`` (see mtplx.qwen4_draft_k20_prescatter), +# default OFF, NOT frozen at import. The server's fixed-M4 auto-arm stamps this +# key AFTER this module is imported, so a module-level constant read here froze +# the default and the served lane never engaged (arming audit 2026-09-07). When +# off the claim below (behind `is not None`) is never built, so the retained +# stock lane runs the code it ran before this module existed. # # When on (and the request is eligible -- the claim RAISES rather than falling # back) each draft step builds its K20 support from the FR-Spec head's 65,536 @@ -320,21 +322,19 @@ def _env_falsey(name: str) -> bool: # 65,536-lane `argpartition` and `logsumexp` instead of 248,320-lane ones, and # the same `(ids, probs)` support because the ranked id table is strictly # ascending. See that module's docstring for the exactness argument. -_QWEN4_DRAFT_K20_PRESCATTER = _qwen4_draft_k20_prescatter_enabled() - -# MTPLX_QWEN4_BLOCK_VERIFY -- read ONCE at import (in -# ``mtplx.qwen4_block_verify``), default OFF. When off this constant is False, -# no verifier is built, and the stock accept loop evaluates exactly the -# expressions it evaluated before -- same acceptance probability, same -# residual, same uniforms, same order. When on, the loop runs block -# verification (Sun et al. 2024, arXiv:2403.10444) instead of the per-token -# Leviathan-Chen law: it clips the RUNNING reach product at 1 rather than -# clipping each factor, water-fills the resulting budget across the depth d+1 -# draft support, and corrects from the SCALED residual (c*p - q)+. Both laws -# are exact samplers of the same target distribution; BV accepts deeper more -# often (+1.85% tokens/window measured offline on 381 real windows) and draws -# exactly the same number of uniforms. See ``mtplx/qwen4_block_verify.py``. -_QWEN4_BLOCK_VERIFY = _qwen4_block_verify_enabled() +# +# MTPLX_QWEN4_BLOCK_VERIFY -- likewise read AT USE via +# ``_qwen4_block_verify_enabled()`` (see mtplx.qwen4_block_verify), default OFF, +# NOT frozen at import (same served-arming reason). The env is frozen once +# serving starts, so the accept loop reads the same value at every step. When +# on, the loop runs block verification (Sun et al. 2024, arXiv:2403.10444) +# instead of the per-token Leviathan-Chen law: it clips the RUNNING reach +# product at 1 rather than clipping each factor, water-fills the resulting +# budget across the depth d+1 draft support, and corrects from the SCALED +# residual (c*p - q)+. Both laws are exact samplers of the same target +# distribution; BV accepts deeper more often (+1.85% tokens/window measured +# offline on 381 real windows) and draws exactly the same number of uniforms. +# See ``mtplx/qwen4_block_verify.py``. def _family_capture_commit_enabled() -> bool: """qwen4_exp layer-owned capture-commit (``MTPLX_FAMILY_CAPTURE_COMMIT``). @@ -9953,7 +9953,7 @@ def emit_new_tokens() -> None: # both are passed as absent. _draft_k20_prescatter_plan = None _draft_k20_prescatter_receipt: dict[str, object] = {"installed": False} - if _QWEN4_DRAFT_K20_PRESCATTER: + if _qwen4_draft_k20_prescatter_enabled(): _draft_k20_prescatter_plan = _qwen4_draft_k20_prescatter_claim( rt, greedy_chain_enabled=_greedy_chain_eligible, @@ -12194,7 +12194,7 @@ def emit_new_tokens() -> None: _host_accept_drafts = draft_tokens _bv = None if ( - _QWEN4_BLOCK_VERIFY + _qwen4_block_verify_enabled() and _host_accept_drafts and sampler.temperature > 0 and target_prefix_tokens is None diff --git a/mtplx/qwen4_block_verify.py b/mtplx/qwen4_block_verify.py index a4f617bd4..576f8c569 100644 --- a/mtplx/qwen4_block_verify.py +++ b/mtplx/qwen4_block_verify.py @@ -125,17 +125,26 @@ def _env_truthy(name: str) -> bool: return os.environ.get(name, "").strip().lower() in {"1", "true", "yes", "on"} -#: Read exactly once, at import. Every call site in ``generation.py`` is -#: behind the module-level constant this feeds, so when the flag is unset the -#: accept loop evaluates the same expressions, in the same order, drawing the -#: same uniforms, as it did before this module existed. -_ENABLED = _env_truthy(_ENV_VAR) +#: ``None`` = resolve from the environment on each read; a test may force a +#: bool via :func:`_configure_for_test`. Read at USE, never at import: the +#: server's fixed-M4 auto-arm stamps MTPLX_QWEN4_BLOCK_VERIFY into the +#: environment AFTER this module is imported, so an import-time read froze the +#: default (off) and the served accept loop never used it -- the arming audit, +#: 2026-09-07. The env is frozen once serving starts, so a per-call read is the +#: same value at every accept step, drawing the same uniforms in the same +#: order as before. +_ENABLED = None def is_enabled() -> bool: - """True when ``MTPLX_QWEN4_BLOCK_VERIFY`` was set at import.""" + """True when ``MTPLX_QWEN4_BLOCK_VERIFY`` is set for this process. - return _ENABLED + Read at use, not frozen at import (a test may force :data:`_ENABLED`). + """ + + if _ENABLED is not None: + return bool(_ENABLED) + return _env_truthy(_ENV_VAR) def _configure_for_test(enabled: bool) -> None: diff --git a/mtplx/qwen4_draft_k20_prescatter.py b/mtplx/qwen4_draft_k20_prescatter.py index 654da6f28..0eb0c7211 100644 --- a/mtplx/qwen4_draft_k20_prescatter.py +++ b/mtplx/qwen4_draft_k20_prescatter.py @@ -196,14 +196,25 @@ def _env_truthy(name: str) -> bool: return os.environ.get(name, "").strip().lower() in {"1", "true", "yes", "on"} -#: Read exactly once, at import. -_ENABLED = _env_truthy(_ENV_VAR) +#: ``None`` = resolve from the environment on each read; a test may force a +#: bool via :func:`_configure_for_test`. Read at USE, never at import: the +#: server's fixed-M4 auto-arm stamps MTPLX_QWEN4_DRAFT_K20_PRESCATTER into the +#: environment AFTER this module is imported (via mtplx.server.openai's +#: generation import), so an import-time read froze the default (off) and the +#: served lane never engaged -- the arming audit, 2026-09-07. The env is frozen +#: once serving starts, so a per-call read returns the same value every time. +_ENABLED = None def is_enabled() -> bool: - """True when ``MTPLX_QWEN4_DRAFT_K20_PRESCATTER`` was set at import.""" + """True when ``MTPLX_QWEN4_DRAFT_K20_PRESCATTER`` is set for this process. - return _ENABLED + Read at use, not frozen at import (a test may force :data:`_ENABLED`). + """ + + if _ENABLED is not None: + return bool(_ENABLED) + return _env_truthy(_ENV_VAR) def _configure_for_test(enabled: bool) -> None: diff --git a/mtplx/runtime_options.py b/mtplx/runtime_options.py index 047b5dfe0..eb442d177 100644 --- a/mtplx/runtime_options.py +++ b/mtplx/runtime_options.py @@ -58,7 +58,14 @@ def env_bool( #: with it on the rewritten sites are value-identical by construction (see #: tests/test_qwen4_opdiet.py, which proves each rewrite against its original #: on random inputs). -_QWEN4_OPDIET = env_bool("MTPLX_QWEN4_OPDIET", default=False) +#: ``None`` = resolve from the environment on each read; a test may force a +#: bool. Read at USE, not frozen at import: the server's fixed-M4 auto-arm +#: stamps MTPLX_QWEN4_OPDIET into the environment AFTER this module is imported +#: (via the generation import), so an import-time read froze the default and +#: the served compiled verifier ran without the op diet (arming audit +#: 2026-09-07). The env is frozen once serving starts, so two traces of the +#: same graph still read the same value. +_QWEN4_OPDIET = None #: The independently selectable rewrites behind the master switch. #: @@ -105,25 +112,37 @@ def parse_opdiet_items( return frozenset(tokens) -_QWEN4_OPDIET_SELECTED = parse_opdiet_items( - os.environ.get("MTPLX_QWEN4_OPDIET_ITEMS") -) +#: ``None`` = resolve from the environment on each read; a test may force a +#: frozenset. Read at use, not frozen at import (same server-arming reason). +_QWEN4_OPDIET_SELECTED = None def qwen4_opdiet_enabled(item: str | None = None) -> bool: """True when the op diet is armed, and this item is selected. ``item=None`` answers only the master switch. Every gated call site names - its item so ``MTPLX_QWEN4_OPDIET_ITEMS`` can isolate one rewrite. + its item so ``MTPLX_QWEN4_OPDIET_ITEMS`` can isolate one rewrite. Read at + use, not frozen at import; a test may set :data:`_QWEN4_OPDIET` / + :data:`_QWEN4_OPDIET_SELECTED` to force the answer. """ - if not _QWEN4_OPDIET: + master = ( + _QWEN4_OPDIET + if _QWEN4_OPDIET is not None + else env_bool("MTPLX_QWEN4_OPDIET", default=False) + ) + if not master: return False if item is None: return True if item not in QWEN4_OPDIET_ITEMS: raise ValueError(f"unknown op-diet item {item!r}") - return item in _QWEN4_OPDIET_SELECTED + selected = ( + _QWEN4_OPDIET_SELECTED + if _QWEN4_OPDIET_SELECTED is not None + else parse_opdiet_items(os.environ.get("MTPLX_QWEN4_OPDIET_ITEMS")) + ) + return item in selected #: W70 -- fused glue inside the compiled fixed-M4 verify body. @@ -156,7 +175,13 @@ def qwen4_opdiet_enabled(item: str | None = None) -> bool: #: ``kernels/qwen4_m4_hyper_read`` already measured at 13.2 tok/s against 67.8. QWEN4_VERIFY_GLUE_ITEMS = ("qsa_rope", "qsa_rope_idx") -_QWEN4_VERIFY_GLUE = env_bool("MTPLX_QWEN4_VERIFY_GLUE", default=False) +#: ``None`` = resolve from the environment on each read; a test may force a +#: bool (directly or via :func:`reset_qwen4_verify_glue_cache`). Read at USE, +#: not frozen at import: the server's fixed-M4 auto-arm stamps +#: MTPLX_QWEN4_VERIFY_GLUE into the environment AFTER this module is imported, +#: so an import-time read froze the default and the served verify body ran +#: without the fused glue (arming audit 2026-09-07). +_QWEN4_VERIFY_GLUE = None def parse_verify_glue_items( @@ -188,29 +213,46 @@ def parse_verify_glue_items( return frozenset(tokens) -_QWEN4_VERIFY_GLUE_SELECTED = parse_verify_glue_items( - os.environ.get("MTPLX_QWEN4_VERIFY_GLUE_ITEMS") -) +#: ``None`` = resolve from the environment on each read; a test may force a +#: frozenset. Read at use, not frozen at import (same server-arming reason). +_QWEN4_VERIFY_GLUE_SELECTED = None def qwen4_verify_glue_enabled(item: str | None = None) -> bool: - """True when the verify-glue flag is armed, and this item is selected.""" + """True when the verify-glue flag is armed, and this item is selected. - if not _QWEN4_VERIFY_GLUE: + Read at use, not frozen at import; a test may set :data:`_QWEN4_VERIFY_GLUE` + / :data:`_QWEN4_VERIFY_GLUE_SELECTED` (directly or via + :func:`reset_qwen4_verify_glue_cache`) to force the answer. + """ + + master = ( + _QWEN4_VERIFY_GLUE + if _QWEN4_VERIFY_GLUE is not None + else env_bool("MTPLX_QWEN4_VERIFY_GLUE", default=False) + ) + if not master: return False if item is None: return True if item not in QWEN4_VERIFY_GLUE_ITEMS: raise ValueError(f"unknown verify-glue item {item!r}") - return item in _QWEN4_VERIFY_GLUE_SELECTED + selected = ( + _QWEN4_VERIFY_GLUE_SELECTED + if _QWEN4_VERIFY_GLUE_SELECTED is not None + else parse_verify_glue_items( + os.environ.get("MTPLX_QWEN4_VERIFY_GLUE_ITEMS") + ) + ) + return item in selected def reset_qwen4_verify_glue_cache(env: Mapping[str, str] | None = None) -> None: - """Re-read the verify-glue gates from the environment. Tests only. + """Force the verify-glue gates from a given environment. Tests only. - The hot path reads these once at import on purpose; this exists so a test - can arm one item without a subprocess, and it is never called by the - runtime. + The runtime reads these at use and never calls this; it exists so a test + can arm one item without a subprocess by FORCING the module globals (which + then win over the environment until reset again). """ global _QWEN4_VERIFY_GLUE, _QWEN4_VERIFY_GLUE_SELECTED diff --git a/tests/test_qwen4_block_verify.py b/tests/test_qwen4_block_verify.py index 080ea35b3..85cdbab9e 100644 --- a/tests/test_qwen4_block_verify.py +++ b/tests/test_qwen4_block_verify.py @@ -287,14 +287,19 @@ def _accept_loop_source(source: str) -> str: return source[start:end] -def test_the_gate_is_read_once_at_import_and_only_in_one_place(): +def test_the_gate_is_read_at_use_and_only_in_one_place(): source = _generation_source() - assert "_QWEN4_BLOCK_VERIFY = _qwen4_block_verify_enabled()" in source - # generation.py never reads the variable itself. + # generation.py consults the reader at USE, not a module constant frozen + # at import: the served fixed-M4 auto-arm stamps the key after import, so a + # frozen constant would leave the lane off as launched (arming audit). + assert "_qwen4_block_verify_enabled()" in source + assert "_QWEN4_BLOCK_VERIFY = " not in source + # generation.py never reads the env var itself. assert 'os.environ.get("MTPLX_QWEN4_BLOCK_VERIFY"' not in source assert '_env_truthy("MTPLX_QWEN4_BLOCK_VERIFY")' not in source module = (REPO_ROOT / "mtplx" / "qwen4_block_verify.py").read_text() assert module.count('_ENV_VAR = "MTPLX_QWEN4_BLOCK_VERIFY"') == 1 + # the env is read in exactly one place (is_enabled), at use. assert module.count("_env_truthy(_ENV_VAR)") == 1 assert "os.environ" not in inspect.getsource(bv_mod.BlockVerifier) assert "os.environ" not in inspect.getsource(bv_mod.build_verifier) @@ -356,9 +361,9 @@ def test_the_verifier_is_built_once_before_the_loop_and_is_gated(): ) assert built < loop assert body.count("_qwen4_build_block_verifier(") == 1 - gate = body.rindex("_QWEN4_BLOCK_VERIFY", 0, built) - # The construction sits under the module-level gate, and under the - # preconditions that make a block law meaningful at all. + gate = body.rindex("_qwen4_block_verify_enabled()", 0, built) + # The construction sits under the at-use gate, and under the preconditions + # that make a block law meaningful at all. guard = body[gate:built] assert "sampler.temperature > 0" in guard assert "target_prefix_tokens is None" in guard diff --git a/tests/test_qwen4_draft_k20_prescatter.py b/tests/test_qwen4_draft_k20_prescatter.py index fc970fee2..c5b569efc 100644 --- a/tests/test_qwen4_draft_k20_prescatter.py +++ b/tests/test_qwen4_draft_k20_prescatter.py @@ -1013,10 +1013,14 @@ def test_claim_declines_without_top_k(armed): # --------------------------------------------------------------------------- -def test_gate_is_off_by_default_in_generation(): +def test_gate_is_off_by_default_in_generation(monkeypatch): from mtplx import generation + from mtplx import qwen4_draft_k20_prescatter as prescatter - assert generation._QWEN4_DRAFT_K20_PRESCATTER is False + # Read at USE (not frozen at import): unforced global + unset env => off. + monkeypatch.setattr(prescatter, "_ENABLED", None) + monkeypatch.delenv("MTPLX_QWEN4_DRAFT_K20_PRESCATTER", raising=False) + assert generation._qwen4_draft_k20_prescatter_enabled() is False assert "draft_k20_prescatter" in generation.GenerationStats.__dataclass_fields__ diff --git a/tests/test_qwen4_opdiet.py b/tests/test_qwen4_opdiet.py index 123c105e6..214fb49e0 100644 --- a/tests/test_qwen4_opdiet.py +++ b/tests/test_qwen4_opdiet.py @@ -79,12 +79,17 @@ def _kernel_primitives(*outputs: mx.array) -> list[str]: # -------------------------------------------------------------------------- -def test_opdiet_defaults_off_and_is_read_once_at_import(monkeypatch): +def test_opdiet_defaults_off_and_is_read_at_use(monkeypatch): + # Read at USE, not frozen at import: the served fixed-M4 auto-arm stamps + # MTPLX_QWEN4_OPDIET AFTER this module is imported, so a frozen import-time + # read left the compiled verifier without the op diet (arming audit). The + # env is frozen once serving starts, so two traces of one graph still agree. + monkeypatch.setattr(runtime_options, "_QWEN4_OPDIET", None) + monkeypatch.setattr(runtime_options, "_QWEN4_OPDIET_SELECTED", None) + monkeypatch.delenv("MTPLX_QWEN4_OPDIET", raising=False) assert runtime_options.qwen4_opdiet_enabled() is False - # A late env change must NOT flip the hot path: the value was frozen at - # import so two traces of one graph can never disagree. monkeypatch.setenv("MTPLX_QWEN4_OPDIET", "1") - assert runtime_options.qwen4_opdiet_enabled() is False + assert runtime_options.qwen4_opdiet_enabled() is True def test_every_gated_module_reads_the_same_flag(): diff --git a/tests/test_qwen4_remainder_arming.py b/tests/test_qwen4_remainder_arming.py index ca62a972c..df21b226a 100644 --- a/tests/test_qwen4_remainder_arming.py +++ b/tests/test_qwen4_remainder_arming.py @@ -111,3 +111,81 @@ def test_remainder_lanes_arm_when_served_not_frozen_at_import(tmp_path, monkeypa assert "MTPLX_QSA_SPARSE_DECODE" not in overrides assert "MTPLX_QSA_SPARSE_DECODE declined to stock" in log assert ro.qsa_sparse_decode_enabled() is False + + +def test_upstream_verify_lanes_read_at_use_not_frozen_at_import(monkeypatch): + """The four upstream fixed-M4 verify lanes the arming audit found dead when + served must resolve the environment at USE too. + + ``mtplx.generation`` was imported at module load above (before any stamp), + so an import-frozen reader would return its default forever. Each lane's + reader must see a stamp applied after that import -- the served order. + """ + + import mtplx.generation as generation + import mtplx.qwen4_block_verify as block_verify + import mtplx.qwen4_draft_k20_prescatter as prescatter + + cases = [ + # (set-unforced, reader, env key) + ( + lambda: monkeypatch.setattr(ro, "_QWEN4_OPDIET", None), + ro.qwen4_opdiet_enabled, + "MTPLX_QWEN4_OPDIET", + ), + ( + lambda: monkeypatch.setattr(ro, "_QWEN4_VERIFY_GLUE", None), + ro.qwen4_verify_glue_enabled, + "MTPLX_QWEN4_VERIFY_GLUE", + ), + ( + lambda: monkeypatch.setattr(prescatter, "_ENABLED", None), + generation._qwen4_draft_k20_prescatter_enabled, + "MTPLX_QWEN4_DRAFT_K20_PRESCATTER", + ), + ( + lambda: monkeypatch.setattr(block_verify, "_ENABLED", None), + generation._qwen4_block_verify_enabled, + "MTPLX_QWEN4_BLOCK_VERIFY", + ), + ] + for unforce, reader, env_key in cases: + unforce() + monkeypatch.delenv(env_key, raising=False) + assert reader() is False, env_key + monkeypatch.setenv(env_key, "1") + assert reader() is True, env_key # the stamp lands AFTER import -> seen + monkeypatch.delenv(env_key, raising=False) + + +def test_the_fixed_m4_auto_arm_stamps_the_upstream_verify_lanes(tmp_path, monkeypatch): + """The fixed-M4 auto-arm stamps OPDIET, BLOCK_VERIFY and VERIFY_GLUE, and + the readers (read at use, after the stamp is applied) arm. + + (DRAFT_K20_PRESCATTER is stamped only on a q8/g64 lm_head pack, gated by the + FR-Spec draft; its read-at-use is covered above.) + """ + + for name in ("_QWEN4_OPDIET", "_QWEN4_OPDIET_SELECTED", "_QWEN4_VERIFY_GLUE", + "_QWEN4_VERIFY_GLUE_SELECTED"): + monkeypatch.setattr(ro, name, None) + for key in ("MTPLX_QWEN4_OPDIET", "MTPLX_QWEN4_VERIFY_GLUE", + "MTPLX_QWEN4_VERIFY_GLUE_ITEMS"): + monkeypatch.delenv(key, raising=False) + + (tmp_path / "config.json").write_text( + json.dumps({"model_type": "qwen4_exp"}), encoding="utf-8" + ) + args = SimpleNamespace( + generation_mode="mtp", verify_strategy="capture_commit", model=str(tmp_path) + ) + monkeypatch.setattr(openai, "_served_model_is_qwen4_fixed_m4", lambda a: True) + overrides = openai._server_runtime_env_overrides(args, {}) + assert normalize_runtime_env_overrides(overrides) == overrides + assert overrides.get("MTPLX_QWEN4_OPDIET") == "1" + assert overrides.get("MTPLX_QWEN4_BLOCK_VERIFY") == "1" + assert overrides.get("MTPLX_QWEN4_VERIFY_GLUE") == "1" + for key, value in overrides.items(): + monkeypatch.setenv(key, value) + assert ro.qwen4_opdiet_enabled() is True + assert ro.qwen4_verify_glue_enabled() is True From 3729c095c3efd41763e7708e13349d680ebeb447 Mon Sep 17 00:00:00 2001 From: davidtai Date: Mon, 7 Sep 2026 11:56:30 -0500 Subject: [PATCH 06/30] feat(qwen4): surface the three no-observable verify lanes in /health The arming audit's lesson: gate on the install verdict, not the env. Three decode-verify lanes had no per-window observable in /health qwen4_install_reports -- draft_k20_prescatter, block_verify, opdiet -- so a served window could not confirm they engaged. Add read-only reports (no behaviour change, no defaults touched): - draft_k20_prescatter: {armed (read at use), engaged (first-use latch set when claim_draft_route installs the route), receipt (the last install receipt)}. - block_verify: {armed, engaged (latched when a block verifier is built for the accept loop)} plus a one-shot "[mtplx] MTPLX_QWEN4_BLOCK_VERIFY armed:" log. - opdiet: {armed, items (configured selection), applied (first-use latch of the items that actually ran at a gated site)}. Each appears only when ARMED (read at use, gate-able without a request), so an unarmed lane stays absent (== off) like the other lanes; the engaged/applied latch rides inside the armed report. CPU test: tests/test_qwen4_remainder_arming.py asserts the three reports are absent when off / =0 and present with armed True under a served-order stamp. --- mtplx/qwen4_block_verify.py | 33 ++++++++++++++++++ mtplx/qwen4_draft_k20_prescatter.py | 47 ++++++++++++++++++++++++- mtplx/runtime_options.py | 39 ++++++++++++++++++++- mtplx/server/openai.py | 30 ++++++++++++++++ tests/test_qwen4_remainder_arming.py | 52 ++++++++++++++++++++++++++++ 5 files changed, 199 insertions(+), 2 deletions(-) diff --git a/mtplx/qwen4_block_verify.py b/mtplx/qwen4_block_verify.py index 576f8c569..d7f1a75dc 100644 --- a/mtplx/qwen4_block_verify.py +++ b/mtplx/qwen4_block_verify.py @@ -154,6 +154,29 @@ def _configure_for_test(enabled: bool) -> None: _ENABLED = bool(enabled) +#: First-use engagement latch: the accept loop has no other per-window +#: observable, so /health surfaces this so the battery can confirm the block +#: law actually ran instead of trusting the env. Read-only reporting. +_ENGAGED = [False] + + +def engagement_report() -> dict: + """Install/first-use verdict for ``/health qwen4_install_reports.block_verify``. + + ``armed`` is read at use (reflects the served auto-arm stamp, gate-able + without a request); ``engaged`` latches True the first window a block + verifier is actually built for the accept loop. + """ + + return {"armed": is_enabled(), "engaged": bool(_ENGAGED[0])} + + +def reset_engagement_for_test() -> None: + """Clear the first-use latch (tests only).""" + + _ENGAGED[0] = False + + # --------------------------------------------------------------------------- # Row preparation -- mirrors offline_block_verification.prepared_row exactly. # --------------------------------------------------------------------------- @@ -500,6 +523,16 @@ def build_verifier( vocab_size = 1 + int( max(int(ids.max()) for ids, _ in (*draft_rows, *target_rows)) ) + if not _ENGAGED[0]: + _ENGAGED[0] = True + import sys as _sys + + print( + "[mtplx] MTPLX_QWEN4_BLOCK_VERIFY armed: block verification engaged " + f"in the accept loop (depth={depth})", + file=_sys.stderr, + flush=True, + ) return BlockVerifier( draft_tokens=draft_tokens, draft_rows=draft_rows, diff --git a/mtplx/qwen4_draft_k20_prescatter.py b/mtplx/qwen4_draft_k20_prescatter.py index 0eb0c7211..0be48d431 100644 --- a/mtplx/qwen4_draft_k20_prescatter.py +++ b/mtplx/qwen4_draft_k20_prescatter.py @@ -224,6 +224,45 @@ def _configure_for_test(enabled: bool) -> None: _ENABLED = bool(enabled) +#: First-use engagement latch + the last install receipt, surfaced at /health +#: so the battery can gate on the install verdict instead of the env. The +#: receipt (``{installed, rows, ...}``) is per-request in generation.py; this +#: latches the last one a claim installed. Read-only reporting. +_ENGAGED = [False] +_LAST_RECEIPT: dict[str, object] = {} + + +def _note_engaged(receipt: dict[str, object] | None = None) -> None: + """Latch first-use for /health (called by ``claim_draft_route`` on install).""" + + _ENGAGED[0] = True + if receipt: + _LAST_RECEIPT.clear() + _LAST_RECEIPT.update(receipt) + + +def engagement_report() -> dict: + """Install/first-use verdict for ``/health qwen4_install_reports.draft_k20_prescatter``. + + ``armed`` is read at use (reflects the served auto-arm stamp, gate-able + without a request); ``engaged`` latches True the first time a claim installs + the pre-scatter route, and the last install receipt rides along. + """ + + return { + "armed": is_enabled(), + "engaged": bool(_ENGAGED[0]), + "receipt": dict(_LAST_RECEIPT), + } + + +def reset_engagement_for_test() -> None: + """Clear the first-use latch and receipt (tests only).""" + + _ENGAGED[0] = False + _LAST_RECEIPT.clear() + + class DraftK20PrescatterIneligible(RuntimeError): """The armed flag cannot work in THIS PROCESS at all. @@ -416,7 +455,7 @@ def claim_draft_route( if not _ENABLED: return None try: - return _claim_draft_route( + plan = _claim_draft_route( rt, draft_sampler=draft_sampler, draft_core=draft_core, @@ -446,6 +485,12 @@ def claim_draft_route( receipt.clear() receipt.update(stamped) return None + if plan is not None: + try: + _note_engaged(plan.to_dict()) + except Exception: + _note_engaged() + return plan def _claim_draft_route( diff --git a/mtplx/runtime_options.py b/mtplx/runtime_options.py index eb442d177..5323ae6d2 100644 --- a/mtplx/runtime_options.py +++ b/mtplx/runtime_options.py @@ -116,6 +116,12 @@ def parse_opdiet_items( #: frozenset. Read at use, not frozen at import (same server-arming reason). _QWEN4_OPDIET_SELECTED = None +#: First-use latch: the op-diet items actually applied at a gated site in the +#: compiled fixed-M4 verify graph. The graph has no other per-window observable, +#: so /health surfaces this so the battery can confirm which rewrites ran +#: instead of trusting the env. Read-only reporting. +_QWEN4_OPDIET_APPLIED: set[str] = set() + def qwen4_opdiet_enabled(item: str | None = None) -> bool: """True when the op diet is armed, and this item is selected. @@ -142,7 +148,38 @@ def qwen4_opdiet_enabled(item: str | None = None) -> bool: if _QWEN4_OPDIET_SELECTED is not None else parse_opdiet_items(os.environ.get("MTPLX_QWEN4_OPDIET_ITEMS")) ) - return item in selected + applied = item in selected + if applied: + _QWEN4_OPDIET_APPLIED.add(item) + return applied + + +def qwen4_opdiet_report() -> dict: + """Install/first-use verdict for ``/health qwen4_install_reports.opdiet``. + + ``armed`` is read at use (reflects the served auto-arm stamp, gate-able + without a request); ``items`` is the configured selection; ``applied`` is + the first-use latch of items that actually ran at a gated site. + """ + + if not qwen4_opdiet_enabled(): + return {"armed": False, "items": [], "applied": sorted(_QWEN4_OPDIET_APPLIED)} + selected = ( + _QWEN4_OPDIET_SELECTED + if _QWEN4_OPDIET_SELECTED is not None + else parse_opdiet_items(os.environ.get("MTPLX_QWEN4_OPDIET_ITEMS")) + ) + return { + "armed": True, + "items": sorted(selected), + "applied": sorted(_QWEN4_OPDIET_APPLIED), + } + + +def reset_qwen4_opdiet_applied_for_test() -> None: + """Clear the applied-items latch (tests only).""" + + _QWEN4_OPDIET_APPLIED.clear() #: W70 -- fused glue inside the compiled fixed-M4 verify body. diff --git a/mtplx/server/openai.py b/mtplx/server/openai.py index 93820aa7e..855d08780 100644 --- a/mtplx/server/openai.py +++ b/mtplx/server/openai.py @@ -18698,6 +18698,36 @@ def _qwen4_install_reports(state: Any) -> dict[str, Any]: glue = getattr(runtime, "_mtplx_qwen4_verify_glue", None) if isinstance(glue, dict): out["verify_glue"] = glue + # PR #391 remainder / arming audit: three decode-verify lanes with no + # per-window observable. Read-only {armed (read at use, gate-able without a + # request) + first-use engaged/applied latch}, so the battery gates on the + # install verdict instead of trusting the env. Present only when ARMED, so an + # unarmed lane stays absent (== off) like the others; the engaged/applied + # latch rides inside the armed report. + try: + from mtplx import qwen4_draft_k20_prescatter as _k20 + + report = _k20.engagement_report() + if report.get("armed"): + out["draft_k20_prescatter"] = report + except Exception: + pass + try: + from mtplx import qwen4_block_verify as _bv + + report = _bv.engagement_report() + if report.get("armed"): + out["block_verify"] = report + except Exception: + pass + try: + from mtplx.runtime_options import qwen4_opdiet_report + + report = qwen4_opdiet_report() + if report.get("armed"): + out["opdiet"] = report + except Exception: + pass try: model = getattr(runtime, "model", None) text = getattr(model, "language_model", model) diff --git a/tests/test_qwen4_remainder_arming.py b/tests/test_qwen4_remainder_arming.py index df21b226a..25c086fde 100644 --- a/tests/test_qwen4_remainder_arming.py +++ b/tests/test_qwen4_remainder_arming.py @@ -189,3 +189,55 @@ def test_the_fixed_m4_auto_arm_stamps_the_upstream_verify_lanes(tmp_path, monkey monkeypatch.setenv(key, value) assert ro.qwen4_opdiet_enabled() is True assert ro.qwen4_verify_glue_enabled() is True + + +def test_health_surfaces_the_three_no_observable_verify_lanes(monkeypatch): + """The three verify lanes with no per-window observable appear in + /health qwen4_install_reports with an at-use ``armed`` the battery can gate + on (True under a served stamp, False when the key is =0). + """ + + import mtplx.qwen4_block_verify as block_verify + import mtplx.qwen4_draft_k20_prescatter as prescatter + + state = SimpleNamespace(runtime=SimpleNamespace(model=None)) + # Unforced globals + cleared first-use latches (the served import-time + # state); the modules were imported at the top of this file, before any stamp. + monkeypatch.setattr(ro, "_QWEN4_OPDIET", None) + monkeypatch.setattr(ro, "_QWEN4_OPDIET_SELECTED", None) + monkeypatch.setattr(block_verify, "_ENABLED", None) + monkeypatch.setattr(prescatter, "_ENABLED", None) + ro.reset_qwen4_opdiet_applied_for_test() + block_verify.reset_engagement_for_test() + prescatter.reset_engagement_for_test() + keys = ( + "MTPLX_QWEN4_OPDIET", + "MTPLX_QWEN4_BLOCK_VERIFY", + "MTPLX_QWEN4_DRAFT_K20_PRESCATTER", + ) + for k in keys: + monkeypatch.delenv(k, raising=False) + + # Off -> absent (== off), so an unarmed lane never claims a verdict. + rep = openai._qwen4_install_reports(state) + assert "opdiet" not in rep + assert "block_verify" not in rep + assert "draft_k20_prescatter" not in rep + + # Served stamp applied AFTER import -> present with armed True at /health + # (the fix: gate-able from the install verdict, not the env). + for k in keys: + monkeypatch.setenv(k, "1") + rep = openai._qwen4_install_reports(state) + assert rep["opdiet"]["armed"] is True + assert rep["block_verify"] == {"armed": True, "engaged": False} + assert rep["draft_k20_prescatter"]["armed"] is True + assert rep["draft_k20_prescatter"]["engaged"] is False + + # Per-key opt-out (=0) -> absent again. + for k in keys: + monkeypatch.setenv(k, "0") + rep = openai._qwen4_install_reports(state) + assert "opdiet" not in rep + assert "block_verify" not in rep + assert "draft_k20_prescatter" not in rep From 27d5ff6b02d6c954d5b7bd7930592db4fc7a6408 Mon Sep 17 00:00:00 2001 From: davidtai Date: Mon, 7 Sep 2026 07:12:49 -0500 Subject: [PATCH 07/30] perf(qwen4): two exact decode lanes for Qwen3.8 Flash-Next, on the PR 391 remainder port Rebase of PR #475 (cached async PLE + pooled-key rowsel) onto the PR 391 remainder port head (perf/qwen38-391-remainder-main = upstream main 2.11.2 + the four remainder lanes: HC_M4, prefill mask fuse, QSA query tile, and the QSA split-K sparse-GQA decode extension. The two lanes here add to the same fixed-M4 lane_defaults / _QWEN4_PORT_KEYS block, and their native loader (ple_cpu_rows) and wheel-signer entry union with the QSA sparse-decode extension's (mtplx_native_qsa), all cleanly additive). PR 391 is closed and its Fable delivery stack (mtplx/full_stack_env.py, the turbo-full-stack profile, mtplx/native's QSA sparse-GQA loader) was never merged, but the maintainer independently re-landed the Flash-Next stack under the MTPLX_QWEN4_*/MTPLX_QSA_* namespace, so both lanes' base contracts are present upstream and both apply: the fixed-M4 compiled-verify auxiliary plane (mtplx/qwen4_fixed_verify.py) and the QSA indexer pooled-key kernel (mtplx/kernels/qsa_indexer_prepare._pool_keys_kernel). The work here is re-siting the arming and native loading off the absent full_stack_env onto upstream's own machinery. Lanes (armed by default for a served fixed-M4 Flash-Next pack, per-lane opt-out): - ple_cached_aux: a native CPU-stream provider stages the fixed 64 M4 n-gram rows and the auxiliary embedding plane is produced with mx.async_eval outside the compiled verifier. The stock owner-side row cache is preserved; declines to stock with a printed reason when the ple_cpu_rows extension is not built. - qsa_pooled_rowsel: the twelve QSA indexers' pooled-key preparation binds the pool kernel metadata once per indexer and shares one inv_freq object. Rebase changes vs the closed-PR commit: - mtplx/qwen4_aux_lanes.py rewritten off full_stack_env: primary keys are MTPLX_QWEN4_PLE_CACHED_AUX / MTPLX_QSA_POOLED_ROWSEL, the PR 391 MTPLX_FABLE_* names kept as aliases (primary wins when both set). - mtplx/server/openai.py: the two keys join the fixed-M4 lane_defaults and _QWEN4_PORT_KEYS (so the existing pop-loop kill-switch honours KEY=0), with an alias pre-step mirroring an operator's MTPLX_FABLE_* export onto the primary. - mtplx/qsa_pooled_rowsel.py: op-diet contract re-pointed from the absent fable_opdiet_enabled to upstream's qwen4_opdiet_enabled (MTPLX_QWEN4_OPDIET). - mtplx/runtime.py: the two installs run after the fixed-M4 verify install, logging instead of the removed _print_install_receipt. - mtplx/profiles.py: the two keys added to MODEL_RUNTIME_ENV_OVERRIDE_KEYS so normalize_runtime_env_overrides accepts the server-stamped values. - mtplx/native/__init__.py: a minimal PLE-only loader (load_ple_cpu_rows_extension / ple_cpu_rows_unavailable_reason); PR 391's qsa_sparse_gqa loader is not reproduced (upstream loads native QSA via kernels/qsa_prefill_direct.py). - scripts/bundle_native_runtime_wheel.py: accepts and Developer-ID signs the new mtplx_native_ple_cpu_rows extension alongside mtplx_qsa_kernels, so a notarized release wheel carries a signed ple_cpu_rows Mach-O. - scripts/fable/setup_over100_venv.sh: builds only ple_cpu_rows. 772f5bea's opt-in interleaved n-gram row cache (MTPLX_NGRAM_ROW_FILE, default off) touches the same _SidecarGather rows but at the disk-layout layer; it is orthogonal to this runtime-scheduling lane and does not subsume it. CPU tests (venv mlx 0.32.2, no GPU): 84 passed across the six lane test files plus the wheel-bundler test. --- docs/perf/pr391-aux-lanes.md | 263 +++++++ .../pr391-charts/pr391-aux-lanes-decode.svg | 70 ++ mtplx/native/__init__.py | 124 ++++ mtplx/ple_cached_aux.py | 560 +++++++++++++++ mtplx/ple_cached_row_handoff.py | 369 ++++++++++ mtplx/profiles.py | 6 + mtplx/qsa_pooled_rowsel.py | 412 +++++++++++ mtplx/qwen4_aux_lanes.py | 139 ++++ mtplx/runtime.py | 92 +++ mtplx/server/openai.py | 43 ++ native_extensions/ple_cpu_rows/.gitignore | 5 + native_extensions/ple_cpu_rows/CMakeLists.txt | 139 ++++ native_extensions/ple_cpu_rows/MANIFEST.in | 6 + native_extensions/ple_cpu_rows/bindings.cpp | 326 +++++++++ .../ple_cpu_rows/cached_sidecar_primitive.cpp | 153 ++++ .../ple_cpu_rows/cached_sidecar_primitive.h | 36 + .../ple_cpu_rows/cached_sidecar_producer.cpp | 491 +++++++++++++ .../ple_cpu_rows/cached_sidecar_producer.h | 182 +++++ .../ple_cpu_rows/host_provider.cpp | 398 +++++++++++ .../ple_cpu_rows/host_provider.h | 127 ++++ .../mtplx_native_ple_cpu_rows/__init__.py | 23 + .../ple_cpu_rows/ple_cpu_rows.cpp | 180 +++++ native_extensions/ple_cpu_rows/ple_cpu_rows.h | 36 + native_extensions/ple_cpu_rows/pyproject.toml | 8 + .../ple_cpu_rows/request_state.h | 196 ++++++ native_extensions/ple_cpu_rows/setup.py | 17 + .../ple_cpu_rows/sidecar_primitive.cpp | 139 ++++ .../ple_cpu_rows/sidecar_producer.cpp | 160 +++++ .../ple_cpu_rows/sidecar_producer.h | 114 +++ scripts/bundle_native_runtime_wheel.py | 13 +- scripts/fable/setup_over100_venv.sh | 76 ++ tests/test_bundle_native_runtime_wheel.py | 58 ++ ...test_pr391_cached_sidecar_primitive_cpu.py | 158 +++++ .../test_pr391_cached_sidecar_producer_cpu.py | 655 ++++++++++++++++++ tests/test_pr391_fixed_m4_pool_install_cpu.py | 468 +++++++++++++ tests/test_pr391_ple_cached_aux_cpu.py | 477 +++++++++++++ .../test_pr391_ple_cached_row_handoff_cpu.py | 324 +++++++++ tests/test_qwen4_aux_lanes.py | 210 ++++++ 38 files changed, 7249 insertions(+), 4 deletions(-) create mode 100644 docs/perf/pr391-aux-lanes.md create mode 100644 docs/perf/pr391-charts/pr391-aux-lanes-decode.svg create mode 100644 mtplx/ple_cached_aux.py create mode 100644 mtplx/ple_cached_row_handoff.py create mode 100644 mtplx/qsa_pooled_rowsel.py create mode 100644 mtplx/qwen4_aux_lanes.py create mode 100644 native_extensions/ple_cpu_rows/.gitignore create mode 100644 native_extensions/ple_cpu_rows/CMakeLists.txt create mode 100644 native_extensions/ple_cpu_rows/MANIFEST.in create mode 100644 native_extensions/ple_cpu_rows/bindings.cpp create mode 100644 native_extensions/ple_cpu_rows/cached_sidecar_primitive.cpp create mode 100644 native_extensions/ple_cpu_rows/cached_sidecar_primitive.h create mode 100644 native_extensions/ple_cpu_rows/cached_sidecar_producer.cpp create mode 100644 native_extensions/ple_cpu_rows/cached_sidecar_producer.h create mode 100644 native_extensions/ple_cpu_rows/host_provider.cpp create mode 100644 native_extensions/ple_cpu_rows/host_provider.h create mode 100644 native_extensions/ple_cpu_rows/mtplx_native_ple_cpu_rows/__init__.py create mode 100644 native_extensions/ple_cpu_rows/ple_cpu_rows.cpp create mode 100644 native_extensions/ple_cpu_rows/ple_cpu_rows.h create mode 100644 native_extensions/ple_cpu_rows/pyproject.toml create mode 100644 native_extensions/ple_cpu_rows/request_state.h create mode 100644 native_extensions/ple_cpu_rows/setup.py create mode 100644 native_extensions/ple_cpu_rows/sidecar_primitive.cpp create mode 100644 native_extensions/ple_cpu_rows/sidecar_producer.cpp create mode 100644 native_extensions/ple_cpu_rows/sidecar_producer.h create mode 100755 scripts/fable/setup_over100_venv.sh create mode 100644 tests/test_pr391_cached_sidecar_primitive_cpu.py create mode 100644 tests/test_pr391_cached_sidecar_producer_cpu.py create mode 100644 tests/test_pr391_fixed_m4_pool_install_cpu.py create mode 100644 tests/test_pr391_ple_cached_aux_cpu.py create mode 100644 tests/test_pr391_ple_cached_row_handoff_cpu.py create mode 100644 tests/test_qwen4_aux_lanes.py diff --git a/docs/perf/pr391-aux-lanes.md b/docs/perf/pr391-aux-lanes.md new file mode 100644 index 000000000..ebb62ce07 --- /dev/null +++ b/docs/perf/pr391-aux-lanes.md @@ -0,0 +1,263 @@ +# Two stacked decode lanes for Qwen3.8 Flash-Next + +This report covers two lanes from an external optimization pass that stack on top +of the Qwen3.8 Flash-Next stack. Both lanes keep the output identical to the stock path. Both lanes change +only the decode timing. + +Terms: + +- PLE: per-layer embedding, fed from an n-gram sidecar table. +- QSA: Qwen Sparse Attention. +- M4: the fixed four-row speculative verify width. +- lane: the unit an operator turns off with one switch. +- tok/s: tokens per second. + +The cell for every number in this report is the canonical served cell: a 16,384 +token templated prompt, 1,024 output tokens, temperature 1, top-p 0.95, top-k 20, +reasoning effort `xhigh`, native MTP depth 3, seeds 20260829 / 20260830 / +20260831, a pre-warmed n-gram table, cross-request prefix restore off (each seed +prefills cold), a 40 degree Celsius thermal hold, and fans at maximum. + +--- + +## 1. What each lane changes + +### 1.1 Cached async PLE (lane `ple_cached_aux`) + +The stock fixed-M4 route builds the auxiliary PLE embedding plane inside the +compiled verifier. This lane moves that work out of the compiled graph. A native +provider hashes the fixed 64 M4 n-gram row IDs and reads the cold rows once. The +provider binds the existing stock row cache; it does not add a second cache. The +lane then produces the auxiliary embedding plane with `mx.async_eval`, so the PLE +rows are produced while the compiled verifier replays. The compiled graph +arithmetic does not change, so the output is identical to the stock path. + +The lane keeps the stock owner-side cache and the stock warm handoff. It holds at +most two pending native tickets. It fails the model load if the native +installation fails; it does not run as stock after a failure. + +- Env key: `MTPLX_FABLE_PLE_CACHED_AUX`. Default on for a served Flash-Next pack. +- Off switch: `MTPLX_FABLE_PLE_CACHED_AUX=0`, or + `--disable-optimization ple_cached_aux`, or `MTPLX_FABLE_DISABLE=ple_cached_aux`. +- Requirement: the native extension `native_extensions/ple_cpu_rows`. Build it + with `scripts/fable/setup_over100_venv.sh`. When the extension is not built, + the lane declines with a printed reason and the server serves the stock path. +- Files: `mtplx/ple_cached_aux.py`, `mtplx/ple_cached_row_handoff.py`, + `native_extensions/ple_cpu_rows/`, the loader in `mtplx/native/__init__.py`. +- Install verdict: `[fable] ple_cached_aux ...` on the server log, and a + `/health` engagement report under `engagement_reports.ple_cached_aux`. + +### 1.2 Fixed-M4 pooled-key rowsel (lane `qsa_pooled_rowsel`) + +The twelve real QSA indexers prepare pooled keys once per decode call. This lane +replaces that preparation with a construction-bound rowsel method. The method +binds the existing pool-keys Metal kernel metadata once per indexer. It shares +one `inv_freq` object across the twelve indexers. It removes the per-call pooled +key setup from the decode path. It uses stock MLX; it needs no native extension. + +The lane is exact by construction. The install checks the 48-layer QSA layout, the +per-indexer geometry, the RMS-norm epsilon, the RoPE scale, the shared `inv_freq` +object identity, and the rope and bank op-diet items. The install reports bank +mode `rowsel` with no weight copies. A contract failure fails the model load. + +The method derives the number of new pooled blocks from the write width, not from +an assumed single block. So the method stays correct if a width-parameterized +verify writes 5 or 6 rows in one step +(`MTPLX_QWEN4_FIXED_VERIFY_ROWS`, which lands separately): four rows fill one +pooled block, and five to eight rows fill two. + +- Env key: `MTPLX_FABLE_QSA_POOLED_ROWSEL`. Default on for a served Flash-Next + pack. +- Off switch: `MTPLX_FABLE_QSA_POOLED_ROWSEL=0`, or + `--disable-optimization qsa_pooled_rowsel`, or + `MTPLX_FABLE_DISABLE=qsa_pooled_rowsel`. +- Files: `mtplx/qsa_pooled_rowsel.py`. +- Install verdict: `[fable] qsa_pooled_rowsel ...` on the server log, and a + `/health` engagement report under `engagement_reports.qsa_pooled_rowsel`. + +### 1.3 Arming + +The server arms both lanes for a served Flash-Next pack, the same way it arms the +retained stack: a `setdefault` behind the served-config check. An operator export +beats the default. The two lanes stay out of the retained 44-key stack and out of +the full-stack self-check, so the PR-391 battery counts and the committed flag +files do not move. The two lanes share the retained stack's off switches only: +`--disable-optimization`, `MTPLX_FABLE_DISABLE`, and `all`. + +`GET /health` reports the resolved state of both lanes under +`aux_lane_defaults`, beside the retained stack's `fable_defaults` block. + +--- + +## 2. The 16,384 / 1,024 ABBA retest + +The retest ran six guarded windows in the order control, cached, pooled, pooled, +cached, control (c1, a1, b1, b2, a2, c2). The bracketing controls cancel the +linear drift across the run. Every window ran the three seeds cold. The engine is +MTPLX 2.10.2 on MLX 0.32.2, served from the branch, with the 100 GiB memory cap. + +![Decode tok/s per window, three seeds, ABBA retest](pr391-charts/pr391-aux-lanes-decode.svg) + +### 2.1 Decode tok/s per window, per seed + +| Window | Arm | 20260829 | 20260830 | 20260831 | Mean | +| --- | --- | ---: | ---: | ---: | ---: | +| c1 | control | 80.357 | 75.851 | 82.757 | 79.655 | +| a1 | cached | 82.261 | 77.439 | 84.219 | 81.306 | +| b1 | pooled | 81.325 | 76.683 | 83.920 | 80.642 | +| b2 | pooled | 82.310 | 77.720 | 84.569 | 81.533 | +| a2 | cached | 82.825 | 77.574 | 84.680 | 81.693 | +| c2 | control | 81.691 | 76.973 | 83.953 | 80.872 | + +Prefill stays flat across the arms (control 1,316.0, cached 1,317.0, pooled +1,316.3 tok/s mean). Peak memory stays flat (92.62 GB mean). Both arms produce +output that is byte-identical to control: the reasoning digest and the text +digest match control on all three seeds, for both cached windows and both pooled +windows. + +### 2.2 Paired deltas + +The window-mean delta compares the mean of an arm's two windows against the mean +of the two bracketing controls. + +| Arm | Decode mean tok/s | vs control 80.264 | 17,408 (history) | +| --- | ---: | ---: | ---: | +| control (c1, c2) | 80.264 | n/a | n/a | +| cached async PLE (a1, a2) | 81.500 | +1.236 (+1.54%) | +1.20% | +| fixed-M4 pooled (b1, b2) | 81.088 | +0.824 (+1.03%) | +0.77% | + +The per-seed delta pairs the two control windows against the two candidate +windows at the same seed. + +| Seed | control | cached | Δ cached | pooled | Δ pooled | control window spread | +| --- | ---: | ---: | ---: | ---: | ---: | ---: | +| 20260829 | 81.024 | 82.543 | +1.87% | 81.818 | +0.98% | 1.65% | +| 20260830 | 76.412 | 77.506 | +1.43% | 77.201 | +1.03% | 1.47% | +| 20260831 | 83.355 | 84.449 | +1.31% | 84.245 | +1.07% | 1.43% | +| mean | n/a | n/a | +1.54% | n/a | +1.03% | 1.52% | + +### 2.3 Noise band and verdicts + +The seed-to-seed swing is large, because it follows the completion length; the +same-seed pairing cancels it. The load-bearing noise is the window-to-window +spread: the two control windows differ by about 1.52%. A claim near 1% therefore +needs the separation test, not the raw magnitude. + +Separation test: does every candidate window beat every control window? + +- Cached async PLE: yes, with no overlap. At every seed the slower cached window + is above the faster control window, and the cached window-mean minimum (81.306) + is above the control window-mean maximum (80.872). +- Fixed-M4 pooled: no. The slower pooled window falls below the faster control + window at all three seeds, and at the window-mean level (80.642 below 80.872). + +Verdicts: + +- **Cached async PLE reproduces.** The gain is small (about +1.5% decode) and + near the window-to-window noise band, but it separates from noise: every cached + measurement beats its matched control, per seed and per window-mean. The output + is byte-identical and prefill and peak are flat. +- **Fixed-M4 pooled does not separate from noise at this shape.** The direction + is positive and stable (+0.98% to +1.07% per seed, echoing the +0.77% at 17,408), but + the effect overlaps the window-to-window drift with two windows per arm. The + output is byte-identical. + +--- + +## 3. Original numbers at 17,408 tokens (history) + +The lanes were first measured at a different shape: 17,408 templated prompt tokens +(1,024 Python input plus 16,384 prefill) and 1,024 output tokens. That shape is +not the canonical cell, which is why the section 2 retest exists. The original +numbers are kept here as history. + +| Lane | candidate | control | Delta | +| --- | ---: | ---: | ---: | +| Cached async PLE (best, A2) | 80.8518 tok/s | 79.8925 tok/s | +1.20% decode, -0.71% wall | +| Fixed-M4 pooled rowsel | 80.4197 tok/s | 79.8085 tok/s | +0.77% decode, -0.39% wall | + +--- + +## 4. Same-build served pair + +Two guarded windows ran on one build of this worktree, on the canonical cell. +The `stack-both` window serves the defaults, so the server arms both lanes; the +`control` window serves the same build with both lanes off +(`MTPLX_FABLE_PLE_CACHED_AUX=0`, `MTPLX_FABLE_QSA_POOLED_ROWSEL=0`). The server +log confirms engagement on the `stack-both` window: `[fable] ple_cached_aux` +installs `variant=async_aux`, and `[fable] qsa_pooled_rowsel` installs 12 rowsel +bindings with one shared `inv_freq` object. + +| Seed | Arm | prefill tok/s | decode tok/s | peak GB | wall s | TTFT s | gen (finish) | reasoning sha | text sha | +| --- | --- | ---: | ---: | ---: | ---: | ---: | --- | --- | --- | +| 20260829 | stack-both | 1,284.2 | 82.952 | 89.20 | 25.32 | 12.96 | 1024 (length) | cfc57ad86ebd | e3b0c44298fc | +| 20260829 | control | 1,287.3 | 81.232 | 89.20 | 25.55 | 12.93 | 1024 (length) | cfc57ad86ebd | e3b0c44298fc | +| 20260830 | stack-both | 1,350.4 | 77.530 | 93.50 | 20.19 | 12.30 | 610 (stop) | a69cd27623b6 | cf5e14d1a99c | +| 20260830 | control | 1,353.2 | 76.571 | 93.50 | 20.26 | 12.27 | 610 (stop) | a69cd27623b6 | cf5e14d1a99c | +| 20260831 | stack-both | 1,351.5 | 84.761 | 95.17 | 24.02 | 12.32 | 990 (stop) | 0b28bbfa9fad | 2baf608e1946 | +| 20260831 | control | 1,353.4 | 83.602 | 95.17 | 24.17 | 12.30 | 990 (stop) | 0b28bbfa9fad | 2baf608e1946 | +| mean | stack-both | 1,328.7 | 81.747 | 92.62 | 23.17 | 12.52 | n/a | n/a | n/a | +| mean | control | 1,331.3 | 80.468 | 92.62 | 23.33 | 12.50 | n/a | n/a | n/a | + +The two lanes together add **+1.279 decode tok/s, +1.59%** over the same build +with both lanes off (per seed +2.12% / +1.25% / +1.39%). Prefill and peak memory +are flat, and TTFT is flat. The output is byte-identical: the reasoning digest +and the text digest match between the two arms on all three seeds. This +same-build pair is the number this pull request cites. + + +--- + +## 5. Changes that did not work + +Other candidates were measured on the same family. Every row below is measured and +rejected. The code of these rows is not in this pull request. The numbers are +from the lane inventory built on 2026-09-06. + +| Lane | Measured effect at 17,408 (or noted shape) | Reason it is not here | +| --- | --- | --- | +| Native sidecar sync raw | 73.03 tok/s (about -8.5%) | The uncached predecessor of the cached lane; the synchronous native read sits on the critical path. | +| Native sidecar async raw | 78.34 tok/s (about -1.9%) | The uncached predecessor; the cached lane supersedes it. | +| GDN conv/norm fused rows | -0.53 tok/s, +0.12 s wall | The component win did not survive the full model. | +| Empty finalization | -0.33% tok/s, +0.15% wall | Near flat, and it needs a custom profiler MLX build. | +| Native compiled-graph replay slot plan | -0.07% tok/s, +0.38% wall | An MLX-core change that measured flat to slightly slower. | +| Command-buffer timing extension | -0.18% tok/s | Measurement instrument, not a speedup; it costs to enable. | +| GPU-stream PLE transport | -1.88% tok/s, +0.86% wall | Replacing the CPU queue transport with a GPU-stream factory did not help. | +| Deferred CPU PLE serving | -1.93% tok/s, +1.02% wall | The schedule trades old overlap for new overlap and shows no gain. | +| Queued sampled-D3 selector | -4.49% tok/s, +2.52% wall | The queued GPU-to-CPU-to-GPU boundary adds draft cost. | +| MTP depth 4 | 55.77 tok/s (about -30%) | Closed by the maintainer; the deeper draft costs more than it accepts. | +| Draft temperature 0.85 | about 77.76 tok/s against 79.79 | Closed by the maintainer; the benchmark contract is temperature 1. | + +--- + +## 6. How to run and how to disable + +Build the venv and both native extensions: + +```bash +scripts/fable/setup_over100_venv.sh +``` + +Serve the pack. The server arms both lanes by default: + +```bash +mtplx serve \ + --model ~/.mtplx/models/Youssofal--Qwen3.8-Flash-Next-MTPLX-Optimized-Speed \ + --model-id mtplx-flash-next-optimized-speed +``` + +Read `GET /health` and check the `aux_lane_defaults` block. Confirm that the +`[fable] ple_cached_aux` and `[fable] qsa_pooled_rowsel` verdicts name both lanes. + +Turn one lane off: + +```bash +mtplx serve --disable-optimization ple_cached_aux --model ... +MTPLX_FABLE_QSA_POOLED_ROWSEL=0 mtplx serve --model ... +``` + +Turn both stacked lanes off: + +```bash +mtplx serve --disable-optimization ple_cached_aux,qsa_pooled_rowsel --model ... +``` diff --git a/docs/perf/pr391-charts/pr391-aux-lanes-decode.svg b/docs/perf/pr391-charts/pr391-aux-lanes-decode.svg new file mode 100644 index 000000000..1c40418df --- /dev/null +++ b/docs/perf/pr391-charts/pr391-aux-lanes-decode.svg @@ -0,0 +1,70 @@ + + +Qwen3.8 Flash-Next stacked lanes: decode tok/s per window (16,384 / 1,024, ABBA retest) + +74 + +76 + +78 + +80 + +82 + +84 + +86 + + +decode tok/s + + + + + + +control +c1 +cached +a1 +pooled +b1 +pooled +b2 +cached +a2 +control +c2 + + + + + + + + + + + + + + + + + + + + + + + +seed 20260829 + + +seed 20260830 + + +seed 20260831 +Run order c1, a1, b1, b2, a2, c2. Column tint marks the arm: grey control, blue cached PLE (a), green pooled rowsel (b). Each cached and pooled window is above both control windows at every seed. + \ No newline at end of file diff --git a/mtplx/native/__init__.py b/mtplx/native/__init__.py index 697419902..bf791d591 100644 --- a/mtplx/native/__init__.py +++ b/mtplx/native/__init__.py @@ -714,3 +714,127 @@ def qsa_sparse_gqa_decode( if stream is None: return extension.qsa_sparse_gqa_decode(*args) return extension.qsa_sparse_gqa_decode(*args, stream=stream) + + + + +# --------------------------------------------------------------------------- +# CPU-stream PLE row staging extension (mtplx_native_ple_cpu_rows) +# --------------------------------------------------------------------------- + +#: The API the cached lane requires from the extension. +_PLE_CPU_ROWS_REQUIRED = ( + "install_cached_sidecar_provider", + "compute_cached_row_ids", + "make_cached_sidecar_rows", + "drain_cached_completions", +) + + +def _ple_cpu_rows_extension_path() -> Path: + return ( + Path(__file__).resolve().parents[2] + / "native_extensions" + / "ple_cpu_rows" + ) + + +@lru_cache(maxsize=1) +def _ple_cpu_rows_abi_mismatch() -> str | None: + """nanobind-vs-mlx.core ABI reason for the PLE extension, or ``None``. + + A matching build and import that still cannot cast an ``mx::array``. Named + here so the reason is printed rather than surfacing as a bare ``TypeError`` + at first call. + """ + + ours = sorted(_ple_cpu_rows_extension_path().glob( + "mtplx_native_ple_cpu_rows/_ext*.so" + )) + if not ours: + return None + core = sorted(Path(mx.__file__).parent.glob("core*.so")) if mx.__file__ else [] + if not core: + return None + ext_version = _nanobind_internals_version(ours[0]) + mlx_version = _nanobind_internals_version(core[0]) + if ext_version is None or mlx_version is None or ext_version == mlx_version: + return None + return ( + f"the built PLE extension uses nanobind internals v{ext_version} but " + f"mlx.core uses v{mlx_version}, so it cannot resolve mlx::core::array " + "and every call raises TypeError; rebuild it with " + "-DMTPLX_NANOBIND_DIR set to a nanobind whose src/nb_abi.h says " + f"NB_INTERNALS_VERSION {mlx_version}" + ) + + +@lru_cache(maxsize=1) +def _load_ple_cpu_rows_extension() -> Any: + """Import the built PLE extension, or return the import/ABI error.""" + + native_path = str(_ple_cpu_rows_extension_path()) + if native_path not in sys.path: + sys.path.insert(0, native_path) + try: + import mtplx_native_ple_cpu_rows # noqa: PLC0415 + except Exception as exc: # pragma: no cover - depends on build state + return exc + mismatch = _ple_cpu_rows_abi_mismatch() + if mismatch is not None: # pragma: no cover - depends on build state + return RuntimeError(mismatch) + missing = [ + name + for name in _PLE_CPU_ROWS_REQUIRED + if not callable(getattr(mtplx_native_ple_cpu_rows, name, None)) + ] + if missing: # pragma: no cover - depends on build state + return RuntimeError( + "the built PLE extension lacks required callables: " + + ", ".join(missing) + ) + return mtplx_native_ple_cpu_rows + + +def native_ple_cpu_rows_available() -> bool: + """True when the built PLE extension imports and exposes its cached API.""" + + return not isinstance(_load_ple_cpu_rows_extension(), Exception) + + +def ple_cpu_rows_unavailable_reason() -> str | None: + """The reason the PLE extension cannot be used, or ``None`` when it can. + + A short, printable string for the cached PLE lane's graceful decline when + the extension has not been built. + """ + + loaded = _load_ple_cpu_rows_extension() + if not isinstance(loaded, Exception): + return None + if isinstance(loaded, ModuleNotFoundError): + return ( + "native_extensions/ple_cpu_rows is not built " + "(run scripts/fable/setup_over100_venv.sh)" + ) + return f"{type(loaded).__name__}: {loaded}" + + +def load_ple_cpu_rows_extension() -> Any: + """Return the imported PLE extension module, raising the stored error. + + The cached PLE lane calls :func:`ple_cpu_rows_unavailable_reason` first and + declines when it is not ``None``; this is the accessor for the armed path. + """ + + loaded = _load_ple_cpu_rows_extension() + if isinstance(loaded, Exception): + raise loaded + return loaded + + +__all__ += [ + "native_ple_cpu_rows_available", + "ple_cpu_rows_unavailable_reason", + "load_ple_cpu_rows_extension", +] diff --git a/mtplx/ple_cached_aux.py b/mtplx/ple_cached_aux.py new file mode 100644 index 000000000..1ec6f224c --- /dev/null +++ b/mtplx/ple_cached_aux.py @@ -0,0 +1,560 @@ +"""Construction-bound cached native PLE auxiliary for fixed M4. + +This lane reuses the stock fixed-M4 builder and the owner-thread stock sidecar +cache. A native provider computes the fixed 64 row IDs and reads cold rows; +Python owns only the compact hit/miss handoff and publication. The auxiliary +embedding plane is produced outside the compiled verifier via ``mx.async_eval`` +so PLE-row production overlaps compiled replay. It does not change the compiled +graph arithmetic, so the output is identical to the stock path. + +The module is inert on import and never imports MLX until an explicit installer +is called. The native provider it needs lives in the ``mtplx_native_ple_cpu_rows`` +extension (``native_extensions/ple_cpu_rows``), loaded through +:mod:`mtplx.native`. +""" + +from __future__ import annotations + +from dataclasses import dataclass +import os +from typing import Any, Callable + +import numpy as np + +from . import ple_cached_row_handoff + + +PENDING_LIMIT = 2 +NATIVE_CACHED_PROVIDER_API = frozenset( + { + "install_cached_sidecar_provider", + "compute_cached_row_ids", + "make_cached_sidecar_rows", + "drain_cached_completions", + } +) +_EMPTY_PACKED_MISSES = np.empty( + (0, ple_cached_row_handoff.PACKED_ROW_BYTES), dtype=np.uint8 +) +_EMPTY_PACKED_MISSES.flags.writeable = False + + +# --- the exact production fixed-M4 sidecar contract ------------------------- +# +# These constants and :func:`_validate_contract` freeze the one sidecar +# geometry the native provider is built for, and are computed once at +# installation. Nothing here is read again in the per-token call. +_EXPECTED_CONTEXT_LEN = 2 +_EXPECTED_NGRAM_SIZE = 3 +_EXPECTED_HEADS_PER_NGRAM = 8 +_EXPECTED_OUTPUT_DIM = 2560 +_EXPECTED_BITS = 4 +_EXPECTED_GROUP_SIZE = 32 +_EXPECTED_HEAD_COUNT = 16 +_EXPECTED_WEIGHT_SHAPE = (20,) +_EXPECTED_METADATA_SHAPE = (5,) +_EXPECTED_WEIGHT_ROW_BYTES = 80 +_EXPECTED_METADATA_ROW_BYTES = 10 +_PLANE_NAMES = ("weight", "scales", "biases") + + +def _shape(value: Any) -> tuple[int, ...]: + """Return a tuple shape without importing MLX or touching model state.""" + + return tuple(int(dimension) for dimension in value.shape) + + +def _validate_contract(inner: Any, embedding: Any, sidecar: Any) -> dict[str, Any]: + """Validate and freeze the exact production sidecar contract once.""" + + observed = ( + int(embedding.context_len), + int(embedding.ngram_size), + int(embedding.heads_per_ngram), + int(inner.args.ple_embed_dim), + int(sidecar.bits), + int(sidecar.group_size), + ) + expected = ( + _EXPECTED_CONTEXT_LEN, + _EXPECTED_NGRAM_SIZE, + _EXPECTED_HEADS_PER_NGRAM, + _EXPECTED_OUTPUT_DIM, + _EXPECTED_BITS, + _EXPECTED_GROUP_SIZE, + ) + if observed != expected: + raise ValueError( + "cached native PLE sidecar geometry mismatch: " + f"observed={observed} expected={expected}" + ) + + # _np_consts is the model's already-derived hash plan. Compute it once at + # installation and pass these exact values to the native provider; never + # reconstruct metadata in the per-token call. + mult, sizes, offsets = embedding._np_consts() + constants = (mult, sizes, offsets) + if tuple(_shape(value) for value in constants) != ( + (3,), + (_EXPECTED_HEAD_COUNT,), + (_EXPECTED_HEAD_COUNT,), + ): + raise ValueError( + "cached native PLE sidecar hash constants mismatch: " + f"shapes={tuple(_shape(value) for value in constants)}" + ) + + maps: dict[str, tuple[int, int]] = {} + row_count: int | None = None + expected_planes = { + "weight": (_EXPECTED_WEIGHT_SHAPE, "U32", _EXPECTED_WEIGHT_ROW_BYTES), + "scales": (_EXPECTED_METADATA_SHAPE, "BF16", _EXPECTED_METADATA_ROW_BYTES), + "biases": (_EXPECTED_METADATA_SHAPE, "BF16", _EXPECTED_METADATA_ROW_BYTES), + } + for name in _PLANE_NAMES: + try: + matrix, dtype_name = sidecar._maps[name] + except (AttributeError, KeyError, TypeError) as exc: + raise ValueError( + f"cached native PLE sidecar is missing {name} map" + ) from exc + shape = _shape(matrix) + expected_shape, expected_dtype, row_bytes = expected_planes[name] + if len(shape) != 2 or shape[1:] != expected_shape: + raise ValueError( + f"cached native PLE {name} shape mismatch: " + f"observed={shape} expected=(*,{expected_shape})" + ) + if str(dtype_name) != expected_dtype: + raise ValueError( + f"cached native PLE {name} dtype mismatch: " + f"observed={dtype_name!r} expected={expected_dtype!r}" + ) + current_rows = int(shape[0]) + if row_count is None: + row_count = current_rows + elif current_rows != row_count: + raise ValueError( + "cached native PLE sidecar plane row counts differ: " + f"{row_count} versus {current_rows} ({name})" + ) + if current_rows <= 0: + raise ValueError("cached native PLE sidecar row count must be positive") + try: + offset = int(matrix.offset) + nbytes = int(matrix.nbytes) + except (AttributeError, TypeError, ValueError) as exc: + raise ValueError( + f"cached native PLE {name} map lacks offset/nbytes" + ) from exc + if offset < 0 or nbytes != current_rows * row_bytes: + raise ValueError( + f"cached native PLE {name} byte span mismatch: " + f"offset={offset} nbytes={nbytes} " + f"expected_nbytes={current_rows * row_bytes}" + ) + maps[name] = (offset, nbytes) + + try: + source_fd = int(sidecar._fd) + except (AttributeError, TypeError, ValueError) as exc: + raise ValueError( + "cached native PLE sidecar has no readable descriptor" + ) from exc + + # Convert the one-time numpy metadata into the std::array-compatible tuples + # expected by nanobind. The values are copied now, not read again from the + # model during generation. + native_constants = tuple( + tuple(int(value) for value in np.asarray(array).reshape(-1)) + for array in constants + ) + return { + "embedding": embedding, + "sidecar": sidecar, + "row_count": int(row_count), + "maps": maps, + "source_fd": source_fd, + "multipliers": native_constants[0], + "sizes": native_constants[1], + "offsets": native_constants[2], + "eos": int(embedding.eos_id), + "output_dim": _EXPECTED_OUTPUT_DIM, + } + + +@dataclass(frozen=True, slots=True) +class CachedAuxInstallation: + """Installation-owned resources shared by every built aux wrapper.""" + + provider: Any + auxiliary_stream: Any + _state: Any + + @property + def pending_count(self) -> int: + return self._state.pending_count + + +class _CachedAuxInstallationState: + """Bounded provider state and terminal failure boundary.""" + + __slots__ = ( + "provider", + "handoff", + "compute_cached_row_ids", + "make_cached_sidecar_rows", + "drain_cached_completions", + "pending", + "failure", + ) + + def __init__( + self, + *, + provider: Any, + handoff: ple_cached_row_handoff.CachedRowHandoff, + compute_cached_row_ids: Callable[..., Any], + make_cached_sidecar_rows: Callable[..., Any], + drain_cached_completions: Callable[..., Any], + ) -> None: + self.provider = provider + self.handoff = handoff + self.compute_cached_row_ids = compute_cached_row_ids + self.make_cached_sidecar_rows = make_cached_sidecar_rows + self.drain_cached_completions = drain_cached_completions + self.pending: dict[Any, ple_cached_row_handoff.PreparedRows] = {} + self.failure: BaseException | None = None + + @property + def pending_count(self) -> int: + return len(self.pending) + + def ensure_healthy(self) -> None: + if self.failure is not None: + raise RuntimeError( + "cached native PLE installation failed; model reload required" + ) from self.failure + + def fail(self, error: BaseException) -> None: + if self.failure is None: + self.failure = error + # Native jobs are no longer safe to associate with Python cache state; + # release the bounded owner references and stop permanently. + self.pending.clear() + + def register( + self, ticket: Any, prepared: ple_cached_row_handoff.PreparedRows + ) -> None: + self.ensure_healthy() + try: + if ticket is None: + raise RuntimeError("miss-bearing cached PLE request returned no ticket") + if ticket in self.pending: + raise RuntimeError("cached native PLE ticket was registered twice") + if len(self.pending) >= PENDING_LIMIT: + raise RuntimeError( + f"cached native PLE pending limit {PENDING_LIMIT} exceeded" + ) + # Register before the caller submits the returned planes to MLX. + self.pending[ticket] = prepared + except BaseException as error: + self.fail(error) + raise + + def publish_all_hit(self, prepared: ple_cached_row_handoff.PreparedRows) -> None: + self.ensure_healthy() + try: + self.handoff.publish( + self.handoff.trusted_completion(prepared, _EMPTY_PACKED_MISSES) + ) + except BaseException as error: + self.fail(error) + raise + + def drain(self) -> None: + """Drain and publish all known completions before any new route work.""" + + self.ensure_healthy() + try: + completions = self.drain_cached_completions(self.provider) + pairs = [] + seen = set() + for ticket, packed in completions: + if ticket is None or ticket in seen or ticket not in self.pending: + raise RuntimeError("cached native PLE completion ticket is unknown") + seen.add(ticket) + pairs.append((ticket, self.pending[ticket], packed)) + # Validate association for the entire drained batch before any + # cache publication, so an unknown later ticket cannot partially + # publish an earlier completion. + for ticket, prepared, packed in pairs: + self.pending.pop(ticket) + self.handoff.publish(self.handoff.trusted_completion(prepared, packed)) + except BaseException as error: + self.fail(error) + raise + + +def _install_cached_provider(native_module: Any, contract: dict[str, Any]) -> Any: + """Install the cached provider with the original immutable layout/hash.""" + + duplicate_fd = os.dup(contract["source_fd"]) + try: + return native_module.install_cached_sidecar_provider( + duplicate_fd, + contract["row_count"], + contract["maps"]["weight"][0], + contract["maps"]["weight"][1], + contract["maps"]["scales"][0], + contract["maps"]["scales"][1], + contract["maps"]["biases"][0], + contract["maps"]["biases"][1], + contract["multipliers"], + contract["sizes"], + contract["offsets"], + contract["eos"], + io_workers=8, + ) + finally: + os.close(duplicate_fd) + + +class _CachedFixedM4SidecarAux: + """Stock warm ownership plus installation-shared cached row handoff.""" + + __slots__ = ( + "_auxiliary_stream", + "_compute_cached_row_ids", + "_dequantize", + "_make_cached_sidecar_rows", + "_previous_tokens", + "_state", + "_stock", + "_stream", + "_submit_embedding", + "_submit_planes", + "_output_dim", + ) + + def __init__( + self, + stock_aux: Any, + *, + state: _CachedAuxInstallationState, + auxiliary_stream: Any, + compute_cached_row_ids: Callable[..., Any], + make_cached_sidecar_rows: Callable[..., Any], + previous_tokens: Callable[..., tuple[int, int]], + submit_planes: Callable[..., Any], + submit_embedding: Callable[..., Any], + mx_module: Any, + output_dim: int, + ) -> None: + self._stock = stock_aux + self._state = state + self._auxiliary_stream = auxiliary_stream + self._compute_cached_row_ids = compute_cached_row_ids + self._make_cached_sidecar_rows = make_cached_sidecar_rows + self._previous_tokens = previous_tokens + self._submit_planes = submit_planes + self._submit_embedding = submit_embedding + self._dequantize = mx_module.dequantize + self._stream = mx_module.stream + self._output_dim = int(output_dim) + + def __call__( + self, + _input_ids: Any, + host_input_ids: Any, + completion_tokens: Any, + committed_count: int, + ) -> Any: + self._state.ensure_healthy() + self._state.drain() + + # Upstream's fixed-M4 aux (mtplx/qwen4_fixed_verify._FixedM4SidecarAux) + # is a pure producer: warm ownership lives on the _SidecarGather, which + # warms itself (qwen4_exp.py _submit_warm), so the aux no longer carries + # a pending-warm / install-owned-rows step. This lane only replaces plane + # production; hit/miss routing may differ from the stock owner cache but + # the gathered row bytes are identical, so the plane stays exact. + stock = self._stock + try: + previous = self._previous_tokens( + stock._prompt_tail, + completion_tokens, + committed_count, + ) + current_ids = tuple(int(value) for value in host_input_ids) + row_ids = self._compute_cached_row_ids( + self._state.provider, + previous, + current_ids, + ) + prepared = self._state.handoff.prepare(row_ids) + ticket, planes = self._make_cached_sidecar_rows( + self._state.provider, + prepared.source, + prepared.hit_packed, + prepared.miss_ids, + ) + if ticket is None: + if prepared.miss_count: + raise RuntimeError( + "cached native PLE miss request returned no completion ticket" + ) + self._state.publish_all_hit(prepared) + else: + if not prepared.miss_count: + raise RuntimeError( + "cached native PLE all-hit request returned a ticket" + ) + self._state.register(ticket, prepared) + + self._submit_planes(*planes) + with self._stream(self._auxiliary_stream): + embedding = self._dequantize( + planes[0], + planes[1], + planes[2], + group_size=_EXPECTED_GROUP_SIZE, + bits=_EXPECTED_BITS, + ).reshape(1, 4, self._output_dim) + self._submit_embedding(embedding) + return embedding + except BaseException as error: + self._state.fail(error) + raise + + +def _install_cached_builder( + runtime: Any, + *, + native_module: Any, + mx_module: Any, + stock_module: Any | None, + submit_planes: Callable[..., Any], + submit_embedding: Callable[..., Any], +) -> CachedAuxInstallation: + if stock_module is None: + from . import qwen4_fixed_verify as stock_module + if not callable(submit_planes) or not callable(submit_embedding): + raise ValueError("cached native PLE requires bound plane/embedding submitters") + original_builder = getattr(runtime, "build_fixed_m4_compiled_verify_aux", None) + if not callable(original_builder): + raise ValueError("cached native PLE requires the stock compiled-aux builder") + required = ( + "install_cached_sidecar_provider", + "compute_cached_row_ids", + "make_cached_sidecar_rows", + "drain_cached_completions", + ) + if any(not callable(getattr(native_module, name, None)) for name in required): + raise ValueError("cached native PLE provider API is incomplete") + + inner = stock_module._inner(runtime) + layer_index = int(inner._ple_stage_idx) + ple = inner.layers[layer_index].ple + embedding = ple.ple_embedding + sidecar = embedding.ngram_embedding._sidecar + if sidecar is None: + raise ValueError("cached native PLE requires the sidecar table") + contract = _validate_contract(inner, embedding, sidecar) + provider = _install_cached_provider(native_module, contract) + handoff = ple_cached_row_handoff.bind_stock_cache(sidecar) + state = _CachedAuxInstallationState( + provider=provider, + handoff=handoff, + compute_cached_row_ids=native_module.compute_cached_row_ids, + make_cached_sidecar_rows=native_module.make_cached_sidecar_rows, + drain_cached_completions=native_module.drain_cached_completions, + ) + auxiliary_stream = mx_module.new_stream(mx_module.gpu) + previous_tokens = stock_module._fixed_m4_previous_tokens + stock_aux_type = stock_module._FixedM4SidecarAux + + def build(*args: Any, **kwargs: Any) -> _CachedFixedM4SidecarAux: + state.ensure_healthy() + state.drain() + try: + stock_aux = original_builder(*args, **kwargs) + if not isinstance(stock_aux, stock_aux_type): + raise ValueError( + "cached native PLE requires stock materialized M4 aux" + ) + return _CachedFixedM4SidecarAux( + stock_aux, + state=state, + auxiliary_stream=auxiliary_stream, + compute_cached_row_ids=state.compute_cached_row_ids, + make_cached_sidecar_rows=state.make_cached_sidecar_rows, + previous_tokens=previous_tokens, + submit_planes=submit_planes, + submit_embedding=submit_embedding, + mx_module=mx_module, + output_dim=contract["output_dim"], + ) + except BaseException as error: + state.fail(error) + raise + + runtime.build_fixed_m4_compiled_verify_aux = build + return CachedAuxInstallation(provider, auxiliary_stream, state) + + +def install_fixed_m4_cached_aux_builder( + runtime: Any, + *, + native_module: Any, + mx_module: Any | None = None, + stock_module: Any | None = None, +) -> CachedAuxInstallation: + """Install the asynchronous cached native-PLE builder explicitly.""" + + if mx_module is None: + import mlx.core as mx_module + return _install_cached_builder( + runtime, + native_module=native_module, + mx_module=mx_module, + stock_module=stock_module, + submit_planes=mx_module.async_eval, + submit_embedding=mx_module.async_eval, + ) + + +def install_fixed_m4_sync_cached_aux_builder( + runtime: Any, + *, + native_module: Any, + mx_module: Any | None = None, + stock_module: Any | None = None, +) -> CachedAuxInstallation: + """Install the synchronous cached native-plane diagnostic builder.""" + + if mx_module is None: + import mlx.core as mx_module + return _install_cached_builder( + runtime, + native_module=native_module, + mx_module=mx_module, + stock_module=stock_module, + submit_planes=mx_module.eval, + submit_embedding=mx_module.async_eval, + ) + + +install_qwen4_fixed_cached_aux = install_fixed_m4_cached_aux_builder +install_qwen4_fixed_sync_cached_aux = install_fixed_m4_sync_cached_aux_builder + + +__all__ = [ + "CachedAuxInstallation", + "PENDING_LIMIT", + "NATIVE_CACHED_PROVIDER_API", + "install_fixed_m4_cached_aux_builder", + "install_fixed_m4_sync_cached_aux_builder", + "install_qwen4_fixed_cached_aux", + "install_qwen4_fixed_sync_cached_aux", +] diff --git a/mtplx/ple_cached_row_handoff.py b/mtplx/ple_cached_row_handoff.py new file mode 100644 index 000000000..87ef2541a --- /dev/null +++ b/mtplx/ple_cached_row_handoff.py @@ -0,0 +1,369 @@ +"""Owner-thread handoff for the stock Qwen4 PLE row cache. + +The native producer owns hashing and cold-row reads, but it never reads or +writes the Python LRU. This module is the narrow owner-side seam: it turns a +fixed 64-row native request into compact hit/miss storage and publishes a +trusted completed miss batch back into the existing stock cache. + +There is intentionally no MLX import here. The payload is the exact packed +sidecar row used by the stock cache: 20 uint32 weight values followed by five +uint16 scale values and five uint16 bias values (100 bytes total). +""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any + +import numpy as np + + +MAX_ROW_SLOTS = 64 +PACKED_ROW_BYTES = 100 +_HIT_FLAG = 0x80 +_INDEX_MASK = 0x3F +_PLANE_SPECS = ( + ("weight", np.dtype(np.uint32), (20,), 80), + ("scales", np.dtype(np.uint16), (5,), 10), + ("biases", np.dtype(np.uint16), (5,), 10), +) + + +@dataclass(frozen=True, slots=True) +class PreparedRows: + """Immutable compact request handed to the native row producer. + + ``source`` has one byte per fixed slot. Bit 7 selects ``hit_packed``; + otherwise it selects ``miss_ids``. The low six bits are the compact index + in the selected array. This keeps the native-side storage fixed at 64 + slots without retaining duplicate Python ID arrays. + """ + + source: np.ndarray + hit_packed: np.ndarray + miss_ids: np.ndarray + # Host-only replay order. The native ABI consumes only the first three + # arrays; the owner uses this immutable sorted-unique order to reproduce + # stock's insert-then-touch LRU sequence at publication. + touch_order: np.ndarray + # Host-only compact-index tags for ``touch_order``. A hit tag points into + # ``hit_packed`` so publication can restore a hit evicted while native I/O + # was in flight; a miss tag points into ``miss_ids``. + touch_source: np.ndarray + + @property + def hit_count(self) -> int: + return int(self.hit_packed.shape[0]) + + @property + def miss_count(self) -> int: + return int(self.miss_ids.shape[0]) + + +@dataclass(frozen=True, slots=True) +class _CompletedMisses: + """A completion ticket minted by one :class:`CachedRowHandoff`. + + The public checked constructor copies and validates an untrusted batch. + The native path can use ``trusted_completion`` after its extension has + established the immutable packed-buffer contract; ``publish`` then has no + repeated ID/shape validation in the enabled path. + """ + + prepared: PreparedRows + packed: np.ndarray + _owner_token: object + + +def _pack_into_trusted(destination: np.ndarray, payload: Any) -> None: + """Pack a construction-validated stock payload without hot checks.""" + + for value, (_name, _dtype, _shape, row_bytes), start in zip( + payload, + _PLANE_SPECS, + (0, 80, 90), + ): + destination[start : start + row_bytes] = ( + np.asarray(value).view(np.uint8).reshape(-1) + ) + + +def _pack_into_checked(destination: np.ndarray, payload: Any) -> None: + """Checked external boundary for one stock payload.""" + + for value, (_name, dtype, shape, row_bytes), start in zip( + payload, + _PLANE_SPECS, + (0, 80, 90), + ): + array = np.asarray(value) + # These are construction-bound stock cache values. Let a malformed + # external fake fail at the boundary rather than adding checks to the + # installed per-cycle handoff path. + if array.dtype != dtype or tuple(array.shape) != shape: + raise ValueError( + "stock PLE row payload geometry mismatch: " + f"observed={array.dtype}/{tuple(array.shape)} " + f"expected={dtype}/{shape}" + ) + destination[start : start + row_bytes] = array.view(np.uint8).reshape(-1) + + +def pack_row_payload(payload: Any) -> np.ndarray: + """Checked external helper producing one immutable packed stock row.""" + + packed = np.empty((PACKED_ROW_BYTES,), dtype=np.uint8) + _pack_into_checked(packed, payload) + packed.flags.writeable = False + return packed + + +def _snapshot_row_ids_checked(row_ids: Any) -> np.ndarray: + """Copy and validate the fixed native row-id boundary once.""" + + values = np.asarray(row_ids) + if values.shape != (MAX_ROW_SLOTS,): + raise ValueError( + f"native PLE row handoff requires exactly {MAX_ROW_SLOTS} IDs; " + f"got shape {values.shape}" + ) + if not np.issubdtype(values.dtype, np.integer): + raise ValueError("native PLE row IDs must have an integer dtype") + if np.any(values < 0) or np.any(values > np.iinfo(np.uint32).max): + raise ValueError("native PLE row IDs must fit uint32") + return np.ascontiguousarray(values, dtype=np.uint32) + + +def _snapshot_row_ids_trusted(row_ids: Any) -> np.ndarray: + """Copy the fixed native ABI without revalidating it per cycle.""" + + return np.array(row_ids, dtype=np.uint32, order="C", copy=True) + + +def _readonly_contiguous(array: Any, *, dtype: Any) -> np.ndarray: + result = np.ascontiguousarray(array, dtype=dtype) + result.flags.writeable = False + return result + + +class CachedRowHandoff: + """Bind owner-side operations to one existing stock sidecar cache.""" + + __slots__ = ( + "_hot", + "_hot_cap_rows", + "_hot_row_bytes", + "_row_specs", + "_owner_token", + ) + + def __init__(self, sidecar: Any) -> None: + self._hot = sidecar._hot + self._hot_row_bytes = int(sidecar._hot_row_bytes) + self._hot_cap_rows = int(sidecar._hot_cap_rows) + if self._hot_row_bytes != PACKED_ROW_BYTES: + raise ValueError( + "stock PLE cache row-byte policy changed: " + f"observed={self._hot_row_bytes} expected={PACKED_ROW_BYTES}" + ) + if self._hot_cap_rows < 0: + raise ValueError("stock PLE cache row limit cannot be negative") + self._row_specs = self._capture_row_specs(sidecar) + self._owner_token = object() + + @staticmethod + def _capture_row_specs(sidecar: Any) -> tuple[tuple[int, Any, tuple[int, ...], int], ...]: + specs = [] + for name, expected_dtype, expected_shape, row_bytes in _PLANE_SPECS: + try: + matrix, _dtype_name = sidecar._maps[name] + except (AttributeError, KeyError, TypeError) as exc: + raise ValueError(f"stock PLE sidecar lacks {name} map") from exc + observed_shape = tuple(int(value) for value in matrix.shape[1:]) + observed_dtype = np.dtype(matrix.dtype) + if observed_dtype != expected_dtype or observed_shape != expected_shape: + raise ValueError( + f"stock PLE {name} geometry mismatch: " + f"observed={observed_dtype}/{observed_shape} " + f"expected={expected_dtype}/{expected_shape}" + ) + if int(np.prod(observed_shape)) * observed_dtype.itemsize != row_bytes: + raise ValueError(f"stock PLE {name} row-byte calculation mismatch") + start = sum(item[3] for item in specs) + specs.append((start, observed_dtype, observed_shape, row_bytes)) + return tuple(specs) + + @property + def row_bytes(self) -> int: + return self._hot_row_bytes + + @property + def limit_rows(self) -> int: + return self._hot_cap_rows + + def prepare(self, row_ids: Any) -> PreparedRows: + """Resolve a trusted fixed-64 native request against the stock LRU. + + The native installer guarantees shape, dtype, range, and row geometry + once. ``checked_prepare`` is the diagnostic/untrusted boundary; this + enabled method intentionally does not call it or repeat those checks. + """ + + slots = _snapshot_row_ids_trusted(row_ids) + unique, inverse = np.unique(slots, return_inverse=True) + unique_count = int(unique.size) + hit_flags = np.empty((unique_count,), dtype=np.bool_) + for index, row in enumerate(unique): + hit_flags[index] = int(row) in self._hot + + hit_unique_indices = np.flatnonzero(hit_flags) + miss_unique_indices = np.flatnonzero(~hit_flags) + hit_count = int(hit_unique_indices.size) + miss_count = int(miss_unique_indices.size) + hit_packed = np.empty((hit_count, PACKED_ROW_BYTES), dtype=np.uint8) + compact_index = np.empty((unique_count,), dtype=np.uint8) + + # np.unique returns sorted IDs, exactly as the stock _rows_matrices + # path does. Publication inserts misses first, then replays this + # complete order before eviction, matching stock. A hit is packed + # here but deliberately not touched until publication so an owner + # gather can interleave safely while native misses are outstanding. + for compact, unique_index in enumerate(hit_unique_indices): + row = int(unique[unique_index]) + _pack_into_trusted(hit_packed[compact], self._hot[row]) + compact_index[unique_index] = _HIT_FLAG | compact + for compact, unique_index in enumerate(miss_unique_indices): + compact_index[unique_index] = compact + + source = _readonly_contiguous(compact_index[inverse], dtype=np.uint8) + miss_ids = _readonly_contiguous(unique[miss_unique_indices], dtype=np.uint32) + touch_order = _readonly_contiguous(unique, dtype=np.uint32) + touch_source = _readonly_contiguous(compact_index, dtype=np.uint8) + hit_packed.flags.writeable = False + return PreparedRows( + source=source, + hit_packed=hit_packed, + miss_ids=miss_ids, + touch_order=touch_order, + touch_source=touch_source, + ) + + def checked_prepare(self, row_ids: Any) -> PreparedRows: + """Checked diagnostic/test boundary for a native row request.""" + + snapshot = _snapshot_row_ids_checked(row_ids) + return self.prepare(snapshot) + + def checked_completion( + self, + prepared: PreparedRows, + completed_miss_ids: Any, + packed: Any, + ) -> _CompletedMisses: + """Validate and freeze an untrusted native/test completion boundary.""" + + if not isinstance(prepared, PreparedRows): + raise TypeError("completed PLE rows require PreparedRows") + observed_ids = np.asarray(completed_miss_ids) + if observed_ids.shape != prepared.miss_ids.shape: + raise ValueError("completed PLE miss IDs shape does not match miss IDs") + if observed_ids.dtype != np.uint32 or not np.array_equal( + observed_ids, prepared.miss_ids + ): + raise ValueError("completed PLE miss IDs do not match prepared miss IDs") + observed_packed = np.asarray(packed) + expected_shape = (prepared.miss_count, PACKED_ROW_BYTES) + if observed_packed.shape != expected_shape: + raise ValueError( + "completed PLE packed rows shape does not match prepared misses: " + f"observed={observed_packed.shape} expected={expected_shape}" + ) + if observed_packed.dtype != np.uint8: + raise ValueError("completed PLE packed rows must be uint8") + frozen = np.array(observed_packed, dtype=np.uint8, order="C", copy=True) + frozen.flags.writeable = False + return _CompletedMisses(prepared, frozen, self._owner_token) + + def trusted_completion( + self, prepared: PreparedRows, packed: np.ndarray + ) -> _CompletedMisses: + """Mint a ticket from an already-validated native immutable batch. + + The native producer is expected to return a C-contiguous read-only + uint8 ``(miss_count, 100)`` array whose rows correspond to the + immutable ``prepared.miss_ids``. The checked constructor above is for + untrusted test/extension boundaries; this method deliberately performs + no repeated shape or ID work in the installed route. + """ + + return _CompletedMisses(prepared, packed, self._owner_token) + + def publish(self, completion: _CompletedMisses) -> None: + """Install a trusted completion with stock ownership/eviction rules.""" + + prepared = completion.prepared + packed = completion.packed + # The ticket is trusted here. Each row is copied into its own 100-byte + # owner buffer before typed views are retained by _hot; retaining one + # cache row therefore cannot retain the producer's entire 6.4 KiB batch. + for index, row in enumerate(prepared.miss_ids): + self._hot[int(row)] = self._payload_from_packed_row(packed[index]) + # Stock _rows_matrices inserts all misses, touches every sorted unique + # ID, and only then evicts. Replaying the host-only order preserves + # that exact mixed hit/miss LRU result. + for index, row in enumerate(prepared.touch_order): + row_id = int(row) + source = int(prepared.touch_source[index]) + if source & _HIT_FLAG and row_id not in self._hot: + hit_index = source & _INDEX_MASK + self._hot[row_id] = self._payload_from_packed_row( + prepared.hit_packed[hit_index] + ) + self._hot.move_to_end(row_id) + while len(self._hot) > self._hot_cap_rows: + self._hot.popitem(last=False) + + def checked_publish(self, completion: _CompletedMisses) -> None: + """Checked diagnostic/test boundary around the trusted publisher.""" + + if not isinstance(completion, _CompletedMisses): + raise TypeError("publish requires a completed PLE row ticket") + if completion._owner_token is not self._owner_token: + raise ValueError("completed PLE row ticket belongs to another handoff") + self.publish(completion) + + def _payload_from_packed_row(self, packed_row: Any) -> tuple[np.ndarray, ...]: + owner = np.array( + packed_row, + dtype=np.uint8, + order="C", + copy=True, + ).reshape(PACKED_ROW_BYTES) + owner.flags.writeable = False + payload = [] + for start, dtype, shape, row_bytes in self._row_specs: + view = np.frombuffer( + owner, + dtype=dtype, + count=row_bytes // dtype.itemsize, + offset=start, + ).reshape(shape) + view.flags.writeable = False + payload.append(view) + return tuple(payload) + + +def bind_stock_cache(sidecar: Any) -> CachedRowHandoff: + """Capture one existing sidecar cache without creating a second cache.""" + + return CachedRowHandoff(sidecar) + + +__all__ = [ + "CachedRowHandoff", + "MAX_ROW_SLOTS", + "PACKED_ROW_BYTES", + "PreparedRows", + "bind_stock_cache", + "pack_row_payload", +] diff --git a/mtplx/profiles.py b/mtplx/profiles.py index 6ef879fae..6cc012e2c 100644 --- a/mtplx/profiles.py +++ b/mtplx/profiles.py @@ -428,6 +428,12 @@ def announce_runtime_gated_env( "MTPLX_QWEN4_VERIFY_GLUE", "MTPLX_QWEN4_VERIFY_GLUE_ITEMS", "MTPLX_QWEN4_PLE_FIRST_GATHER_EARLY", + # PR #475 aux lanes: the server auto-arms these for a served fixed-M4 + # pack, so they pass through normalize_runtime_env_overrides and must + # be accepted here or the boot-time validator raises on the server's + # own overrides before a weight is read. + "MTPLX_QWEN4_PLE_CACHED_AUX", + "MTPLX_QSA_POOLED_ROWSEL", "MTPLX_SESSION_BANK_SHED_BOUNDARIES", "MTPLX_SESSION_BANK_PROTECTED_TERMINAL", # PR #391 remainder ports (davidtai), stamped by the Flash-Next lane diff --git a/mtplx/qsa_pooled_rowsel.py b/mtplx/qsa_pooled_rowsel.py new file mode 100644 index 000000000..882e0ab75 --- /dev/null +++ b/mtplx/qsa_pooled_rowsel.py @@ -0,0 +1,412 @@ +"""Construction-bound fixed-M4 pooled-key installation for Qwen4 (rowsel). + +The module is inert on import. :func:`install_fixed_m4_pool` is called only +after the model weights and fixed-M4 stack are installed and before graph warmup. +It captures the actual QSA layer weights and the shared RoPE array, then +replaces only the fixed-capacity pooled-key preparation method with a +construction-bound "rowsel" method that binds the existing +``qsa_indexer_pool_keys_metal`` kernel metadata once per indexer and shares one +``inv_freq`` object across the twelve real QSA indexers. The surrounding cache +state, fixed/non-fixed route, and bank ownership remain stock behavior, so the +output is identical to the stock path. + +Stock MLX only: the pool helper is ``mtplx.kernels.qsa_indexer_prepare``'s +``_pool_keys_kernel`` (a ``mx.fast.metal_kernel``); there is no native +extension and no profiler build. +""" + +from __future__ import annotations + +import importlib +import sys +from types import MethodType +from typing import Any, Callable + + +QSA_LAYER_POSITIONS = tuple(range(3, 48, 4)) +EXPECTED_LAYER_COUNT = 48 +EXPECTED_RATIO = 4 +EXPECTED_HEAD_DIM = 128 +EXPECTED_ROTARY_DIM = 64 +EXPECTED_EPS = 1e-6 +EXPECTED_SCALE = 1.0 +POOL_KERNEL_GRID = (32, 1, 1) +POOL_KERNEL_THREADGROUP = (32, 1, 1) +POOL_KERNEL_OUTPUT_SHAPES = ((1, 1, EXPECTED_HEAD_DIM),) + + +class FixedM4PoolInstallError(RuntimeError): + """The construction-time fixed-M4 pool contract was not met.""" + + +class PoolKernelBinding: + """Immutable metadata and callable for one real QSA indexer.""" + + __slots__ = ( + "kernel", + "norm_weight", + "inv_freq", + "head_dim", + "rotary_dim", + "ratio", + "eps", + "scale", + "dtype", + "template", + "grid", + "threadgroup", + "output_shapes", + "output_dtypes", + "mx_module", + ) + + def __init__( + self, + *, + kernel: Callable[..., Any], + norm_weight: Any, + inv_freq: Any, + head_dim: int, + rotary_dim: int, + ratio: int, + eps: float, + scale: float, + dtype: Any, + ) -> None: + self.kernel = kernel + self.norm_weight = norm_weight + self.inv_freq = inv_freq + self.head_dim = int(head_dim) + self.rotary_dim = int(rotary_dim) + self.ratio = int(ratio) + self.eps = float(eps) + self.scale = float(scale) + self.dtype = dtype + self.template = (("T", dtype),) + self.grid = POOL_KERNEL_GRID + self.threadgroup = POOL_KERNEL_THREADGROUP + self.output_shapes = POOL_KERNEL_OUTPUT_SHAPES + self.output_dtypes = (dtype,) + + def pool(self, raw_keys: Any, block_start: Any) -> Any: + """Run the already-bound one-block helper with dynamic input leaves.""" + + result = self.kernel( + inputs=[raw_keys, self.norm_weight, self.inv_freq, block_start], + template=self.template, + grid=self.grid, + threadgroup=self.threadgroup, + output_shapes=self.output_shapes, + output_dtypes=self.output_dtypes, + ) + return result[0] + + +def _text_model(runtime: Any) -> Any: + return getattr(runtime.model, "language_model", runtime.model) + + +def _inner(runtime: Any) -> Any: + text = _text_model(runtime) + inner = getattr(text, "model", None) + if inner is None: + raise FixedM4PoolInstallError("runtime has no Qwen4 text model") + return inner + + +def _resolve_mx(mx_module: Any | None) -> Any: + if mx_module is not None: + return mx_module + return importlib.import_module("mlx.core") + + +def _resolve_kernel_factory(kernel_factory: Callable[..., Any] | None) -> Callable[..., Any]: + if kernel_factory is not None: + return kernel_factory + helper = importlib.import_module("mtplx.kernels.qsa_indexer_prepare") + factory = getattr(helper, "_pool_keys_kernel", None) + if not callable(factory): + raise FixedM4PoolInstallError("QSA pool helper lacks _pool_keys_kernel") + return factory + + +def _loaded_graphbank() -> Any | None: + """Return the already-loaded graphbank module without importing it.""" + + return sys.modules.get("mtplx.graphbank") + + +def _graphbank_has_current_runtime_entries(graphbank: Any, runtime: Any) -> bool: + """Detect fixed graphs already cached for this runtime, cold-only.""" + + runtime_id = id(runtime) + for cache_name in ("_SHARED_VERIFY_STEPS", "_SHARED_OVERLAP_SPLITS"): + entries = getattr(graphbank, cache_name, None) + if not isinstance(entries, dict): + continue + for key, entry in tuple(entries.items()): + if not isinstance(key, tuple) or not key or key[0] != runtime_id: + continue + if not isinstance(entry, (tuple, list)) or not entry: + continue + runtime_ref = entry[-1] + owner = runtime_ref() if callable(runtime_ref) else None + if owner is runtime: + return True + # A stale entry can have the same first key after Python recycles + # an object id. It is not this runtime's trace and must not make a + # cold installation fail. Do not mutate the shared graph cache + # while merely auditing construction state. + return False + + +def _validate_normal_opdiet() -> dict[str, bool]: + """Validate the measured rope+rowsel bank mode once at installation.""" + + options = importlib.import_module("mtplx.runtime_options") + rope = bool(options.qwen4_opdiet_enabled("rope")) + bank = bool(options.qwen4_opdiet_enabled("bank")) + if not rope or not bank: + raise FixedM4PoolInstallError( + "fixed-M4 pool requires the qwen4 op-diet rope and bank items " + "(MTPLX_QWEN4_OPDIET)" + ) + return {"rope": rope, "bank": bank} + + +def _validate_indexers(runtime: Any, mx: Any) -> tuple[Any, ...]: + inner = _inner(runtime) + layers = tuple(getattr(inner, "layers", ())) + if len(layers) != EXPECTED_LAYER_COUNT: + raise FixedM4PoolInstallError( + f"fixed-M4 pool requires {EXPECTED_LAYER_COUNT} layers; got {len(layers)}" + ) + + indexers: list[Any] = [] + for position, layer in enumerate(layers): + expected_qsa = position in QSA_LAYER_POSITIONS + is_linear = bool(getattr(layer, "is_linear", False)) + if expected_qsa == is_linear: + role = "QSA" if expected_qsa else "linear" + raise FixedM4PoolInstallError( + f"layer {position} is not the expected {role} position" + ) + if not expected_qsa: + continue + attention = getattr(layer, "self_attn", None) + indexer = getattr(attention, "indexer", None) + if indexer is None: + raise FixedM4PoolInstallError(f"QSA layer {position} has no indexer") + indexers.append(indexer) + + if tuple(position for position in QSA_LAYER_POSITIONS) != tuple( + position + for position, layer in enumerate(layers) + if not bool(getattr(layer, "is_linear", False)) + ): + raise FixedM4PoolInstallError("QSA layer positions differ from the production 48-layer layout") + + if len(indexers) != len(QSA_LAYER_POSITIONS): + raise FixedM4PoolInstallError( + f"expected {len(QSA_LAYER_POSITIONS)} QSA indexers; got {len(indexers)}" + ) + return tuple(indexers) + + +def _validate_indexer_contract(indexer: Any, mx: Any, position: int) -> tuple[Any, Any, int, int, float, float, Any]: + ratio = int(getattr(indexer, "ratio", -1)) + head_dim = int(getattr(indexer, "head_dim", -1)) + eps = float(getattr(indexer, "rms_norm_eps", float("nan"))) + scale = float(getattr(indexer, "_rope_attention_scaling", float("nan"))) + inv_freq = getattr(indexer, "_inv_freq", None) + norm_module = getattr(indexer, "k_layernorm", None) + norm_weight = getattr(norm_module, "weight", None) + if ratio != EXPECTED_RATIO or head_dim != EXPECTED_HEAD_DIM: + raise FixedM4PoolInstallError( + f"QSA layer {position} geometry mismatch: ratio={ratio} head_dim={head_dim}" + ) + if eps != EXPECTED_EPS or scale != EXPECTED_SCALE: + raise FixedM4PoolInstallError( + f"QSA layer {position} RoPE metadata mismatch: eps={eps} scale={scale}" + ) + if inv_freq is None or norm_weight is None: + raise FixedM4PoolInstallError(f"QSA layer {position} lacks norm or inv_freq") + if int(getattr(inv_freq, "ndim", -1)) != 1: + raise FixedM4PoolInstallError(f"QSA layer {position} inv_freq must be rank 1") + if int(getattr(inv_freq, "shape", (0,))[0]) * 2 != EXPECTED_ROTARY_DIM: + raise FixedM4PoolInstallError(f"QSA layer {position} rotary dimension is not 64") + if int(getattr(norm_weight, "ndim", -1)) != 1 or tuple(norm_weight.shape) != (EXPECTED_HEAD_DIM,): + raise FixedM4PoolInstallError(f"QSA layer {position} norm weight must have shape (128,)") + if getattr(inv_freq, "dtype", None) != mx.float32: + raise FixedM4PoolInstallError(f"QSA layer {position} inv_freq must be float32") + if getattr(norm_weight, "dtype", None) != mx.bfloat16: + raise FixedM4PoolInstallError(f"QSA layer {position} norm weight must be BF16") + norm_eps = float(getattr(norm_module, "eps", float("nan"))) + if norm_eps != eps: + raise FixedM4PoolInstallError( + f"QSA layer {position} norm epsilon mismatch: module={norm_eps} metadata={eps}" + ) + return norm_weight, inv_freq, head_dim, EXPECTED_ROTARY_DIM, eps, scale, norm_weight.dtype + + +def _make_binding( + indexer: Any, + *, + position: int, + mx: Any, + kernel_factory: Callable[..., Any], +) -> PoolKernelBinding: + norm_weight, inv_freq, head_dim, rotary_dim, eps, scale, dtype = _validate_indexer_contract( + indexer, mx, position + ) + kernel = kernel_factory( + head_dim, + rotary_dim, + EXPECTED_RATIO, + eps, + scale, + dtype, + ) + if not callable(kernel): + raise FixedM4PoolInstallError("QSA pool helper did not return a callable") + return PoolKernelBinding( + kernel=kernel, + norm_weight=norm_weight, + inv_freq=inv_freq, + head_dim=head_dim, + rotary_dim=rotary_dim, + ratio=EXPECTED_RATIO, + eps=eps, + scale=scale, + dtype=dtype, + ) + + +def _fixed_m4_pool_method(self: Any, cache: Any, total: Any) -> Any: + """Stock fixed-bank method with only pool arithmetic replaced. + + The Python body is entered at graph trace time. In particular, + ``pooled_capacity`` is read from the compiled bank shape here; it is not a + per-token eligibility check and no capacity is copied into the binding. + + ``max_new`` is derived from the write width (``cache._last_write_rows``), + not assumed to be one block, so the method stays correct when a + width-parameterized verify writes 5 or 6 rows in one step + (MTPLX_QWEN4_FIXED_VERIFY_ROWS): 4 rows fill one pooled block, 5-8 rows fill + two. Keep this derivation if that lane lands. + """ + + binding = self._mtplx_fixed_m4_pool_binding + mx = binding.mx_module + step_rows = int(getattr(cache, "_last_write_rows", 1)) + nb_old = cache.offset // self.ratio + nb_total = total // self.ratio + max_new = max(1, (step_rows + self.ratio - 1) // self.ratio) + pooled = cache.pooled + pooled_capacity = int(pooled.shape[1]) + for rel in range(max_new): + block = nb_old + rel + safe_block = mx.minimum( + block, mx.array(pooled_capacity - 1, dtype=block.dtype) + ) + start = safe_block * self.ratio + fresh = mx.slice( + cache.raw_keys, + start, + axes=(1,), + slice_size=(1, self.ratio, self.head_dim), + ) + block_start = safe_block.reshape(1).astype(mx.int32) + candidate = binding.pool(fresh, block_start) + old_row = mx.slice( + pooled, + safe_block, + axes=(1,), + slice_size=(1, 1, pooled.shape[2]), + ) + merged = mx.where( + nb_total > block, + candidate.astype(pooled.dtype), + old_row, + ) + pooled = mx.slice_update(pooled, merged, safe_block, axes=(1,)) + cache.pooled = pooled + return pooled + + +def install_fixed_m4_pool( + runtime: Any, + *, + mx_module: Any | None = None, + kernel_factory: Callable[..., Any] | None = None, +) -> dict[str, Any]: + """Cold-install the pooled preparation on all twelve real QSA indexers.""" + + if getattr(runtime, "_mtplx_fixed_m4_pool_installed", False): + raise FixedM4PoolInstallError("fixed-M4 pool installation already completed") + mx = _resolve_mx(mx_module) + factory = _resolve_kernel_factory(kernel_factory) + graphbank = _loaded_graphbank() + if graphbank is not None and _graphbank_has_current_runtime_entries(graphbank, runtime): + raise FixedM4PoolInstallError( + "fixed-M4 pool installation must be cold; graphbank already has this runtime" + ) + opdiet = _validate_normal_opdiet() + indexers = _validate_indexers(runtime, mx) + + shared_inv = getattr(indexers[0], "_inv_freq", None) + if shared_inv is None or any(indexer._inv_freq is not shared_inv for indexer in indexers): + raise FixedM4PoolInstallError( + "QSA indexers do not share the exact inv_freq object from TextArgs" + ) + + bindings = tuple( + _make_binding( + indexer, + position=QSA_LAYER_POSITIONS[position], + mx=mx, + kernel_factory=factory, + ) + for position, indexer in enumerate(indexers) + ) + if any(binding.inv_freq is not shared_inv for binding in bindings): + raise FixedM4PoolInstallError("pool bindings lost the shared inv_freq object") + + for indexer, binding in zip(indexers, bindings): + binding.mx_module = mx + original = indexer._extend_pooled_fixed + indexer._mtplx_fixed_m4_pool_original = original + indexer._mtplx_fixed_m4_pool_binding = binding + indexer._extend_pooled_fixed = MethodType(_fixed_m4_pool_method, indexer) + indexer._mtplx_fixed_m4_pool_installed = True + + report = { + "installed": True, + "qsa_layer_positions": list(QSA_LAYER_POSITIONS), + "shared_inv_freq_identity": True, + "shared_inv_freq_object_count": 1, + "kernel_binding_count": len(bindings), + "bank_mode": "rowsel", + "opdiet": opdiet, + "weights_copied": False, + "graph_warmup_required_after_install": True, + } + runtime._mtplx_fixed_m4_pool_installed = True + runtime._mtplx_fixed_m4_pool_install_report = dict(report) + return report + + +install_qwen4_fixed_m4_pool = install_fixed_m4_pool + + +__all__ = [ + "EXPECTED_EPS", + "EXPECTED_HEAD_DIM", + "EXPECTED_RATIO", + "EXPECTED_ROTARY_DIM", + "FixedM4PoolInstallError", + "PoolKernelBinding", + "QSA_LAYER_POSITIONS", + "install_fixed_m4_pool", + "install_qwen4_fixed_m4_pool", +] diff --git a/mtplx/qwen4_aux_lanes.py b/mtplx/qwen4_aux_lanes.py new file mode 100644 index 000000000..46fa5e236 --- /dev/null +++ b/mtplx/qwen4_aux_lanes.py @@ -0,0 +1,139 @@ +"""Two exact decode lanes stacked on the Qwen3.8 Flash-Next fixed-M4 stack. + +These are the two aux lanes from PR #391 / PR #475, rebased onto upstream main: + +* ``ple_cached_aux`` -- the cached async PLE auxiliary (``mtplx/ple_cached_aux.py``), + which needs the ``mtplx_native_ple_cpu_rows`` extension; and +* ``qsa_pooled_rowsel`` -- the fixed-M4 pooled-key rowsel install + (``mtplx/qsa_pooled_rowsel.py``), pure stock MLX. + +Both are exact by construction (timing-only, byte-identical output). The server +auto-arms them for a served fixed-M4 Flash-Next pack the same way it arms the +other PR #391 ports: a ``setdefault`` behind the fixed-M4 predicate in +``mtplx/server/openai.py``, keyed off upstream's ``MTPLX_QWEN4_*`` / ``MTPLX_QSA_*`` +namespace and honoured by the ``_QWEN4_LANE_KEYS`` ``pop`` kill-switch loop, so +an explicit operator export (``KEY=0``) always wins. + +PR #391 spelled these lanes ``MTPLX_FABLE_PLE_CACHED_AUX`` / +``MTPLX_FABLE_QSA_POOLED_ROWSEL`` and armed them through ``mtplx.full_stack_env``. +Upstream main has no ``full_stack_env``; the primary keys are the ``MTPLX_QWEN4_*`` +/ ``MTPLX_QSA_*`` names below and the old ``MTPLX_FABLE_*`` names are kept as +aliases so an existing launch line or pack contract keeps working. + +The module is inert on import: it imports no MLX and loads no extension. +""" + +from __future__ import annotations + +import os +from typing import Mapping + + +#: Boolean vocabulary, matching upstream's lenient ``os.environ.get`` readers. +TRUE_TOKENS = frozenset({"1", "true", "yes", "on"}) + +#: lane -> primary env key (upstream's namespace). +LANE_KEYS: dict[str, str] = { + "ple_cached_aux": "MTPLX_QWEN4_PLE_CACHED_AUX", + "qsa_pooled_rowsel": "MTPLX_QSA_POOLED_ROWSEL", +} + +#: primary env key -> the PR #391 alias an operator or pack may still use. +LANE_ALIASES: dict[str, str] = { + "MTPLX_QWEN4_PLE_CACHED_AUX": "MTPLX_FABLE_PLE_CACHED_AUX", + "MTPLX_QSA_POOLED_ROWSEL": "MTPLX_FABLE_QSA_POOLED_ROWSEL", +} + +#: Lane names, in order. +LANES: tuple[str, ...] = tuple(LANE_KEYS) + + +def _truthy(value: object) -> bool: + return str(value or "").strip().lower() in TRUE_TOKENS + + +def lane_enabled(lane: str, environ: Mapping[str, str] | None = None) -> bool: + """Is ``lane`` armed in ``environ``? Unset is off (the stock path). + + The primary ``MTPLX_QWEN4_*`` / ``MTPLX_QSA_*`` key is authoritative when it + is explicitly set (including ``=0`` as the off switch); an unset primary + falls back to the PR #391 ``MTPLX_FABLE_*`` alias. + """ + + if lane not in LANE_KEYS: + raise KeyError(f"unknown auxiliary lane {lane!r}; expected one of {LANES}") + source = os.environ if environ is None else environ + primary = LANE_KEYS[lane] + raw = source.get(primary) + if raw is not None and str(raw).strip(): + return _truthy(raw) + alias = LANE_ALIASES.get(primary) + if alias is not None: + return _truthy(source.get(alias)) + return False + + +def ple_cached_aux_enabled(environ: Mapping[str, str] | None = None) -> bool: + """``MTPLX_QWEN4_PLE_CACHED_AUX`` read the way its install site reads it.""" + + return lane_enabled("ple_cached_aux", environ) + + +def qsa_pooled_rowsel_enabled(environ: Mapping[str, str] | None = None) -> bool: + """``MTPLX_QSA_POOLED_ROWSEL`` read the way its install site reads it.""" + + return lane_enabled("qsa_pooled_rowsel", environ) + + +#: lane -> the runtime attribute that carries its load-time install report. +_RUNTIME_REPORT_ATTR: dict[str, str] = { + "ple_cached_aux": "ple_cached_aux_report", + "qsa_pooled_rowsel": "qsa_pooled_rowsel_report", +} + +#: install-report keys surfaced in the read-only /health entry, when present. +_HEALTH_REPORT_KEYS: tuple[str, ...] = ("native_ext", "reason", "bank_mode") + + +def health_report( + lane: str, + runtime: Any, + environ: Mapping[str, str] | None = None, +) -> dict[str, Any] | None: + """Read-only /health entry for ``lane``, or ``None`` when it is not armed. + + Shape mirrors the upstream verify-lane reports (commit 6): ``armed`` is read + at use (gate-able without a request), so an unarmed lane returns ``None`` and + stays absent from ``qwen4_install_reports`` (== off). When armed, the + load-time install report on the runtime (set by runtime.py at install) rides + inside: ``status`` (installed/declined) plus ``native_ext`` (ple_cached_aux, + installed), ``reason`` (ple_cached_aux, declined) or ``bank_mode`` + (qsa_pooled_rowsel). Never touches the GPU; never changes behaviour. + """ + + if lane not in LANE_KEYS: + raise KeyError(f"unknown auxiliary lane {lane!r}; expected one of {LANES}") + if not lane_enabled(lane, environ): + return None + entry: dict[str, Any] = {"armed": True} + report = getattr(runtime, _RUNTIME_REPORT_ATTR[lane], None) + if isinstance(report, Mapping): + status = report.get("status") + if status is not None: + entry["status"] = status + for key in _HEALTH_REPORT_KEYS: + value = report.get(key) + if value is not None: + entry[key] = value + return entry + + +__all__ = [ + "LANES", + "LANE_KEYS", + "LANE_ALIASES", + "TRUE_TOKENS", + "lane_enabled", + "ple_cached_aux_enabled", + "qsa_pooled_rowsel_enabled", +] diff --git a/mtplx/runtime.py b/mtplx/runtime.py index a150e56ef..796c83ac9 100644 --- a/mtplx/runtime.py +++ b/mtplx/runtime.py @@ -1021,6 +1021,98 @@ def load( qwen4_verify_report = install_qwen4_fixed_verify_route(runtime) runtime.qwen4_fixed_verify_report = qwen4_verify_report logger.info("[qwen4-fixed-M4-verify] %s", qwen4_verify_report) + # ---- stacked auxiliary lane: cached async PLE (PR #475) --------- + # Wraps the fixed-M4 compiled-verify aux builder installed just + # above so the auxiliary PLE plane is produced outside the compiled + # verifier via mx.async_eval. Exact by construction (the stock + # owner-side row cache is preserved). Needs the native + # ple_cpu_rows extension; declines with a printed reason and serves + # stock when it is not built. Any other failure (a contract miss) + # escapes and fails the load, keeping the exactness contract + # unhealthy-on-failure. + from .qwen4_aux_lanes import ( + ple_cached_aux_enabled, + qsa_pooled_rowsel_enabled, + ) + + if ple_cached_aux_enabled(): + from .native import ( + load_ple_cpu_rows_extension, + ple_cpu_rows_unavailable_reason, + ) + + decline = ple_cpu_rows_unavailable_reason() + if decline is not None: + ple_cached_aux_report = { + "lane": "ple_cached_aux", + "status": "declined", + "reason": decline, + } + else: + from .ple_cached_aux import ( + PENDING_LIMIT, + install_fixed_m4_cached_aux_builder, + ) + + native_module = load_ple_cpu_rows_extension() + installation = install_fixed_m4_cached_aux_builder( + runtime, native_module=native_module + ) + runtime._ple_cached_aux_installation = installation + ple_cached_aux_report = { + "lane": "ple_cached_aux", + "status": "installed", + "variant": "async_aux", + "pending_limit": PENDING_LIMIT, + "native_ext": getattr(native_module, "__file__", None), + } + runtime.ple_cached_aux_report = ple_cached_aux_report + logger.info("[qwen4-ple-cached-aux] %s", ple_cached_aux_report) + # A stderr install line (print, not logger.info, which the serve + # log drops at the default level) so a benchmark window can tell + # an engaged lane from a decline-to-stock, matching the + # qsa_sparse_decode convention. + if ple_cached_aux_report["status"] == "declined": + print( + "[mtplx] ple_cached_aux declined to stock: " + f"{ple_cached_aux_report['reason']}", + file=sys.stderr, + flush=True, + ) + else: + print( + "[mtplx] ple_cached_aux armed: " + f"{ple_cached_aux_report['variant']} " + f"pending_limit={ple_cached_aux_report['pending_limit']} " + f"native_ext={ple_cached_aux_report['native_ext']}", + file=sys.stderr, + flush=True, + ) + # ---- stacked auxiliary lane: fixed-M4 pooled-key rowsel --------- + # Rebinds the twelve QSA indexers' pooled-key preparation to the + # construction-bound rowsel method, sharing one inv_freq object. + # Exact by construction; contract failures escape (unhealthy). + if qsa_pooled_rowsel_enabled(): + from .qsa_pooled_rowsel import install_fixed_m4_pool + + pool_report = install_fixed_m4_pool(runtime) + qsa_pooled_rowsel_report = { + "lane": "qsa_pooled_rowsel", + "status": "installed", + **pool_report, + } + runtime.qsa_pooled_rowsel_report = qsa_pooled_rowsel_report + logger.info( + "[qwen4-qsa-pooled-rowsel] %s", qsa_pooled_rowsel_report + ) + print( + "[mtplx] qsa_pooled_rowsel armed: " + f"bank_mode={pool_report['bank_mode']} " + f"indexers={pool_report['kernel_binding_count']} " + f"shared_inv_freq_objects={pool_report['shared_inv_freq_object_count']}", + file=sys.stderr, + flush=True, + ) from .qwen4_m4_stage3 import ( install_qwen4_m4_stage3, qwen4_m4_stage3_flags, diff --git a/mtplx/server/openai.py b/mtplx/server/openai.py index 855d08780..221b92759 100644 --- a/mtplx/server/openai.py +++ b/mtplx/server/openai.py @@ -925,6 +925,15 @@ def _server_runtime_env_overrides( "MTPLX_QWEN4_BLOCK_VERIFY", "MTPLX_QWEN4_PLE_PREFILL_LOOKAHEAD", "MTPLX_QWEN4_PLE_FIRST_GATHER_EARLY", + # PR #475 (davidtai), measured at the 16,384/1,024 cell on the + # same fixed-M4 geometry: the cached async PLE auxiliary plane + # (native CPU-stream rows, produced outside the compiled + # verifier via mx.async_eval; declines to stock and prints a + # reason when the native extension is not built) and the + # construction-bound fixed-M4 pooled-key rowsel install. Both + # exact by construction (byte-identical output). + "MTPLX_QWEN4_PLE_CACHED_AUX", + "MTPLX_QSA_POOLED_ROWSEL", "MTPLX_SESSION_BANK_SHED_BOUNDARIES", "MTPLX_SESSION_BANK_PROTECTED_TERMINAL", # PR #391 remainder ports (davidtai), same fixed-M4 geometry. @@ -944,6 +953,15 @@ def _server_runtime_env_overrides( # The tail also requires the fused gate+up owners, so the # MTPLX_FUSED_GATE_UP kill switch drops it with them. lane_defaults.append("MTPLX_QWEN4_M4_STAGE3") + # PR #475 aux lanes: mirror an operator's old MTPLX_FABLE_* export + # onto the primary MTPLX_QWEN4_*/MTPLX_QSA_* key when the primary is + # unset, so the alias arms or kills the lane before the default + # stamp. setdefault below then leaves the mirrored value in place. + for _primary, _alias in _QWEN4_AUX_LANE_ALIASES.items(): + if os.environ.get(_primary) is None: + _alias_val = os.environ.get(_alias) + if _alias_val is not None and _alias_val.strip(): + overrides[_primary] = _alias_val for key in lane_defaults: if os.environ.get(key) is None: overrides.setdefault(key, "1") @@ -1143,6 +1161,11 @@ def _served_model_type_is_qwen4_exp(args: argparse.Namespace) -> bool: "MTPLX_QWEN4_VERIFY_GLUE_ITEMS", "MTPLX_QWEN4_PLE_PREFILL_LOOKAHEAD", "MTPLX_QWEN4_PLE_FIRST_GATHER_EARLY", + # PR #475 aux lanes (davidtai): the cached async PLE auxiliary and the + # fixed-M4 pooled-key rowsel install. Both exact by construction; each is + # its own kill switch through the pop loop below. + "MTPLX_QWEN4_PLE_CACHED_AUX", + "MTPLX_QSA_POOLED_ROWSEL", "MTPLX_SESSION_BANK_SHED_BOUNDARIES", "MTPLX_SESSION_BANK_PROTECTED_TERMINAL", # PR #391 remainder ports (davidtai): each is its own kill switch through @@ -1152,6 +1175,13 @@ def _served_model_type_is_qwen4_exp(args: argparse.Namespace) -> bool: "MTPLX_QSA_SPARSE_DECODE", "MTPLX_NGRAM_PREWARM", ) +# PR #475's aux lanes were spelled MTPLX_FABLE_* under PR #391. Upstream has no +# full_stack_env, so the primary keys are the MTPLX_QWEN4_*/MTPLX_QSA_* names +# above; the old names are honoured as aliases when the primary is unset. +_QWEN4_AUX_LANE_ALIASES = { + "MTPLX_QWEN4_PLE_CACHED_AUX": "MTPLX_FABLE_PLE_CACHED_AUX", + "MTPLX_QSA_POOLED_ROWSEL": "MTPLX_FABLE_QSA_POOLED_ROWSEL", +} # Every key the fixed-M4 lane defaults may stamp; an explicit operator # export (any non-empty value) always beats a stamped value for these. _QWEN4_LANE_KEYS = _QWEN4_PORT_KEYS + ( @@ -18728,6 +18758,19 @@ def _qwen4_install_reports(state: Any) -> dict[str, Any]: out["opdiet"] = report except Exception: pass + # PR #475 aux lanes: their own per-window observable. Read-only + # {armed (read at use) + the load-time install report on the runtime}, so a + # served window can tell an engaged lane from a decline-to-stock. Present + # only when ARMED, like the lanes above. + try: + from mtplx import qwen4_aux_lanes as _aux + + for _lane in ("ple_cached_aux", "qsa_pooled_rowsel"): + entry = _aux.health_report(_lane, runtime) + if entry is not None: + out[_lane] = entry + except Exception: + pass try: model = getattr(runtime, "model", None) text = getattr(model, "language_model", model) diff --git a/native_extensions/ple_cpu_rows/.gitignore b/native_extensions/ple_cpu_rows/.gitignore new file mode 100644 index 000000000..ff8e35ef3 --- /dev/null +++ b/native_extensions/ple_cpu_rows/.gitignore @@ -0,0 +1,5 @@ +build/ +*.egg-info/ +mtplx_native_ple_cpu_rows/*.so +mtplx_native_ple_cpu_rows/*.dylib +mtplx_native_ple_cpu_rows/__pycache__/ diff --git a/native_extensions/ple_cpu_rows/CMakeLists.txt b/native_extensions/ple_cpu_rows/CMakeLists.txt new file mode 100644 index 000000000..27620b4c5 --- /dev/null +++ b/native_extensions/ple_cpu_rows/CMakeLists.txt @@ -0,0 +1,139 @@ +cmake_minimum_required(VERSION 3.27) + +project(mtplx_native_ple_cpu_rows_ext LANGUAGES CXX) + +set(CMAKE_CXX_STANDARD 17) +set(CMAKE_CXX_STANDARD_REQUIRED ON) +set(CMAKE_POSITION_INDEPENDENT_CODE ON) + +option(BUILD_SHARED_LIBS "Build extensions as a shared library" ON) + +# CPU-stream PLE row staging primitive for the cached async PLE lane +# (mtplx/ple_cached_aux.py). The module links the active MLX wheel and must be +# built against the exact serving nanobind/mlx ABI; scripts/fable/setup_over100_venv.sh +# builds it the same way it builds native_extensions/qsa_sparse_gqa. Every +# translation unit lives in this directory -- there is no staging step and no +# absolute compiled-in path. + +find_package( + Python 3.11 + COMPONENTS Interpreter Development.Module + REQUIRED) + +set(MTPLX_NANOBIND_DIR "" CACHE PATH + "nanobind package root to build against (default: interpreter's)") +if(MTPLX_NANOBIND_DIR) + set(nanobind_ROOT "${MTPLX_NANOBIND_DIR}/cmake") + set(MTPLX_NANOBIND_SRC "${MTPLX_NANOBIND_DIR}") +else() + execute_process( + COMMAND "${Python_EXECUTABLE}" -m nanobind --cmake_dir + OUTPUT_STRIP_TRAILING_WHITESPACE + OUTPUT_VARIABLE nanobind_ROOT + RESULT_VARIABLE MTPLX_NANOBIND_RESULT) + if(NOT MTPLX_NANOBIND_RESULT EQUAL 0) + message(FATAL_ERROR "the active interpreter cannot locate nanobind") + endif() + get_filename_component(MTPLX_NANOBIND_SRC "${nanobind_ROOT}" DIRECTORY) +endif() +find_package(nanobind CONFIG REQUIRED) + +execute_process( + COMMAND "${Python_EXECUTABLE}" -m mlx --cmake-dir + OUTPUT_STRIP_TRAILING_WHITESPACE + OUTPUT_VARIABLE MTPLX_MLX_PACKAGE_ROOT + RESULT_VARIABLE MTPLX_MLX_RESULT) +if(NOT MTPLX_MLX_RESULT EQUAL 0) + message(FATAL_ERROR "the active interpreter cannot locate MLX headers") +endif() +set(MLX_DIR "${MTPLX_MLX_PACKAGE_ROOT}/share/cmake/MLX") +find_package(MLX CONFIG REQUIRED PATHS "${MLX_DIR}" NO_DEFAULT_PATH) + +# nanobind's type registry is part of the MLX Python ABI. Refuse a build when +# the selected nanobind internals version differs from the loaded mlx.core +# binary; otherwise the extension can build but cannot cast mlx arrays. +set(MTPLX_NB_INTERNALS "") +foreach(_hdr "src/nb_abi.h" "src/nb_internals.h") + if(NOT MTPLX_NB_INTERNALS AND EXISTS "${MTPLX_NANOBIND_SRC}/${_hdr}") + file(STRINGS "${MTPLX_NANOBIND_SRC}/${_hdr}" _nb_lines + REGEX "define[ \\t]+NB_INTERNALS_VERSION[ \\t]+[0-9]+") + foreach(_nb_line IN LISTS _nb_lines) + if(NOT MTPLX_NB_INTERNALS) + string(REGEX MATCH "NB_INTERNALS_VERSION[ \\t]+([0-9]+)" _m + "${_nb_line}") + if(CMAKE_MATCH_1) + set(MTPLX_NB_INTERNALS "${CMAKE_MATCH_1}") + endif() + endif() + endforeach() + endif() +endforeach() +if(NOT MTPLX_NB_INTERNALS) + message(FATAL_ERROR + "could not read NB_INTERNALS_VERSION from selected nanobind source; " + "refusing to build an unchecked MLX extension") +endif() + +set(MTPLX_MLX_NB_INTERNALS "") +file(GLOB MTPLX_MLX_CORE + "${MTPLX_MLX_PACKAGE_ROOT}/core*.so" + "${MTPLX_MLX_PACKAGE_ROOT}/core*.dylib") +if(NOT MTPLX_MLX_CORE) + message(FATAL_ERROR + "the active MLX package root has no core*.so/core*.dylib; refusing to " + "skip the nanobind ABI check") +endif() +list(GET MTPLX_MLX_CORE 0 MTPLX_MLX_CORE_SO) +file(STRINGS "${MTPLX_MLX_CORE_SO}" _mlx_tags + REGEX "v[0-9]+[0-9a-zA-Z._-]*_[0-9a-zA-Z_]*libcpp[0-9a-zA-Z_]*") +foreach(_tag IN LISTS _mlx_tags) + if(NOT MTPLX_MLX_NB_INTERNALS) + string(REGEX MATCH "v([0-9]+)[0-9a-zA-Z._-]*_[0-9a-zA-Z_]*libcpp" + _m "${_tag}") + if(CMAKE_MATCH_1) + set(MTPLX_MLX_NB_INTERNALS "${CMAKE_MATCH_1}") + endif() + endif() +endforeach() +if(NOT MTPLX_MLX_NB_INTERNALS) + message(FATAL_ERROR + "could not read nanobind internals ABI from the active mlx.core; " + "refusing to build an unchecked extension") +endif() +if(MTPLX_NB_INTERNALS AND MTPLX_MLX_NB_INTERNALS AND + NOT MTPLX_NB_INTERNALS STREQUAL MTPLX_MLX_NB_INTERNALS) + message(FATAL_ERROR + "nanobind ABI mismatch: extension v${MTPLX_NB_INTERNALS}, " + "mlx.core v${MTPLX_MLX_NB_INTERNALS}") +endif() + +add_library(mtplx_native_ple_cpu_rows) +target_sources( + mtplx_native_ple_cpu_rows + PUBLIC + ${CMAKE_CURRENT_LIST_DIR}/cached_sidecar_primitive.cpp + ${CMAKE_CURRENT_LIST_DIR}/cached_sidecar_producer.cpp + ${CMAKE_CURRENT_LIST_DIR}/ple_cpu_rows.cpp + ${CMAKE_CURRENT_LIST_DIR}/sidecar_primitive.cpp + ${CMAKE_CURRENT_LIST_DIR}/sidecar_producer.cpp + ${CMAKE_CURRENT_LIST_DIR}/host_provider.cpp) +target_include_directories( + mtplx_native_ple_cpu_rows + PUBLIC + ${CMAKE_CURRENT_LIST_DIR}) +target_link_libraries(mtplx_native_ple_cpu_rows PUBLIC mlx) + +nanobind_add_module( + _ext + NB_STATIC + STABLE_ABI + LTO + NOMINSIZE + NB_DOMAIN + mlx + ${CMAKE_CURRENT_LIST_DIR}/bindings.cpp) +target_link_libraries(_ext PRIVATE mtplx_native_ple_cpu_rows) + +if(BUILD_SHARED_LIBS) + target_link_options(_ext PRIVATE -Wl,-rpath,@loader_path) +endif() diff --git a/native_extensions/ple_cpu_rows/MANIFEST.in b/native_extensions/ple_cpu_rows/MANIFEST.in new file mode 100644 index 000000000..769dc6157 --- /dev/null +++ b/native_extensions/ple_cpu_rows/MANIFEST.in @@ -0,0 +1,6 @@ +include CMakeLists.txt +include *.cpp +include *.h +include pyproject.toml +include setup.py +recursive-include mtplx_native_ple_cpu_rows *.py diff --git a/native_extensions/ple_cpu_rows/bindings.cpp b/native_extensions/ple_cpu_rows/bindings.cpp new file mode 100644 index 000000000..29b1778f6 --- /dev/null +++ b/native_extensions/ple_cpu_rows/bindings.cpp @@ -0,0 +1,326 @@ +// SPDX-License-Identifier: Apache-2.0 + +#include +#include +#include +#include +#include + +#include +#include +#include +#include +#include + +#include "cached_sidecar_primitive.h" +#include "ple_cpu_rows.h" + +namespace nb = nanobind; +using namespace nb::literals; + +namespace { + +using SourceArray = nb::ndarray>; +using HitArray = nb::ndarray>; +using MissArray = nb::ndarray>; +using RowIdArray = nb::ndarray>; +using PackedCompletionArray = nb::ndarray>; + +using OwnedRowIds = std::array; + +void delete_owned_row_ids(void* pointer) noexcept { + delete static_cast(pointer); +} + +void delete_owned_packed(void* pointer) noexcept { + delete static_cast*>(pointer); +} + +RowIdArray owned_row_ids(const mtplx_native::ple_cpu_rows::SidecarRowIds& rows) { + auto* owner = new OwnedRowIds(rows); + nb::capsule capsule(owner, delete_owned_row_ids); + return RowIdArray(owner->data(), {64}, capsule); +} + +PackedCompletionArray owned_packed_completion( + const mtplx_native::ple_cpu_rows::CachedCompletion& completion) { + auto* owner = new std::vector( + static_cast(completion.count) * 100); + for (std::size_t index = 0; index < completion.count; ++index) { + std::memcpy(owner->data() + index * 100, + completion.payloads[index].data(), + 100); + } + nb::capsule capsule(owner, delete_owned_packed); + return PackedCompletionArray( + owner->data(), + {static_cast(completion.count), 100}, + capsule); +} + +mtplx_native::ple_cpu_rows::CachedRowHandoff cached_handoff_from_arrays( + const SourceArray& source, + const HitArray& hits, + const MissArray& misses) { + namespace rows = mtplx_native::ple_cpu_rows; + if (hits.ndim() != 2 || hits.shape(1) != 100 || hits.shape(0) > 64) { + throw nb::value_error( + "cached sidecar hits must be contiguous uint8 with shape (H, 100), H<=64"); + } + if (misses.ndim() != 1 || misses.shape(0) > 64) { + throw nb::value_error( + "cached sidecar misses must be contiguous uint32 with shape (M,), M<=64"); + } + + rows::CachedRowHandoff handoff{}; + std::memcpy(handoff.source.data(), source.data(), handoff.source.size()); + handoff.hit_count = static_cast(hits.shape(0)); + handoff.miss_count = static_cast(misses.shape(0)); + if (hits.size() != 0) { + std::memcpy(handoff.hit_packed.data(), + hits.data(), + hits.size() * sizeof(std::uint8_t)); + } + if (misses.size() != 0) { + std::memcpy(handoff.miss_ids.data(), + misses.data(), + misses.size() * sizeof(std::uint32_t)); + } + return handoff; +} + +mtplx_native::ple_cpu_rows::PackedPayload payload_from_bytes( + const nb::object& object) { + if (!nb::isinstance(object)) { + throw nb::type_error("payload must be exactly 6400 bytes"); + } + // Nanobind's standard-string caster only accepts Python unicode under + // nanobind 2.15; it rejects bytes before the payload can reach the native + // primitive. Keep the checked bytes handle borrowed and copy its binary + // storage, including embedded NUL bytes, into the immutable request payload. + const nb::bytes raw = nb::borrow(object); + if (raw.size() != mtplx_native::ple_cpu_rows::kPayloadBytes) { + throw nb::value_error("payload must contain exactly 6400 bytes"); + } + mtplx_native::ple_cpu_rows::PackedPayload payload{}; + std::memcpy(payload.data(), raw.c_str(), raw.size()); + return payload; +} + +} // namespace + +NB_MODULE(_ext, m) { + m.doc() = + "CPU-stream MLX-owned PLE row staging primitive; dequantization stays " + "in ordinary MLX operations"; + + m.def( + "make_cpu_rows", + [](const nb::object& payload, int delay_ms, bool fail, bool cancel) { + return mtplx_native::ple_cpu_rows::make_cpu_rows( + payload_from_bytes(payload), delay_ms, fail, cancel); + }, + "payload"_a, + "delay_ms"_a = 0, + "fail"_a = false, + "cancel"_a = false, + "Create fresh [64,20] U32 and two [64,5] BF16 MLX-owned planes." + " Submit async_eval(planes) before constructing a GPU consumer."); + + // The provider object has no mutable Python-side state. Its constructor + // duplicates/validates the descriptor and binds the complete hash plan; + // each subsequent call supplies only one copied M4 token snapshot. + nb::class_(m, + "SidecarProducer") + .def_prop_ro("row_count", + &mtplx_native::ple_cpu_rows::SidecarProducer::row_count); + + nb::class_( + m, "CachedSidecarProducer") + .def_prop_ro( + "row_count", + &mtplx_native::ple_cpu_rows::CachedSidecarProducer::row_count) + .def_prop_ro( + "io_workers", + &mtplx_native::ple_cpu_rows::CachedSidecarProducer::io_workers); + + m.def( + "install_sidecar_provider", + [](int descriptor, + std::uint64_t row_count, + std::uint64_t weights_offset, + std::uint64_t weights_length, + std::uint64_t scales_offset, + std::uint64_t scales_length, + std::uint64_t biases_offset, + std::uint64_t biases_length, + const std::array& multipliers, + const std::array& sizes, + const std::array& offsets, + std::int64_t eos) { + mtplx_native::host_provider::SidecarLayout layout{}; + layout.row_count = row_count; + layout.weights = {weights_offset, weights_length, 80}; + layout.scales = {scales_offset, scales_length, 10}; + layout.biases = {biases_offset, biases_length, 10}; + return mtplx_native::ple_cpu_rows::SidecarProducer::install( + descriptor, layout, multipliers, sizes, offsets, eos); + }, + "descriptor"_a, + "row_count"_a, + "weights_offset"_a, + "weights_length"_a, + "scales_offset"_a, + "scales_length"_a, + "biases_offset"_a, + "biases_length"_a, + "multipliers"_a, + "sizes"_a, + "offsets"_a, + "eos"_a, + "Install the immutable real-sidecar plan and duplicated reader."); + + m.def( + "install_cached_sidecar_provider", + [](int descriptor, + std::uint64_t row_count, + std::uint64_t weights_offset, + std::uint64_t weights_length, + std::uint64_t scales_offset, + std::uint64_t scales_length, + std::uint64_t biases_offset, + std::uint64_t biases_length, + const std::array& multipliers, + const std::array& sizes, + const std::array& offsets, + std::int64_t eos, + std::size_t io_workers) { + mtplx_native::host_provider::SidecarLayout layout{}; + layout.row_count = row_count; + layout.weights = {weights_offset, weights_length, 80}; + layout.scales = {scales_offset, scales_length, 10}; + layout.biases = {biases_offset, biases_length, 10}; + return mtplx_native::ple_cpu_rows::CachedSidecarProducer::install( + descriptor, + layout, + multipliers, + sizes, + offsets, + eos, + io_workers); + }, + "descriptor"_a, + "row_count"_a, + "weights_offset"_a, + "weights_length"_a, + "scales_offset"_a, + "scales_length"_a, + "biases_offset"_a, + "biases_length"_a, + "multipliers"_a, + "sizes"_a, + "offsets"_a, + "eos"_a, + "io_workers"_a = + mtplx_native::ple_cpu_rows::kCachedDefaultIoWorkers, + "Install the immutable cache-aware sidecar plan and fixed I/O pool."); + + m.def( + "compute_cached_row_ids", + [](const std::shared_ptr< + mtplx_native::ple_cpu_rows::CachedSidecarProducer>& producer, + const std::array& previous, + const std::array& ids) { + if (producer == nullptr) { + throw nb::value_error("cached sidecar producer is null"); + } + mtplx_native::ple_cpu_rows::SidecarJobInput input{previous, ids}; + return owned_row_ids(producer->compute_row_ids(input)); + }, + "provider"_a, + "previous"_a, + "current"_a, + "Compute one immutable uint32[64] fixed-M4 row-ID snapshot."); + + m.def( + "make_cached_sidecar_rows", + [](const std::shared_ptr< + mtplx_native::ple_cpu_rows::CachedSidecarProducer>& producer, + const SourceArray& source, + const HitArray& hits, + const MissArray& misses) { + const auto handoff = cached_handoff_from_arrays(source, hits, misses); + const auto submission = + mtplx_native::ple_cpu_rows::make_cached_sidecar_rows( + producer, handoff); + nb::object ticket = submission.ticket.has_value() + ? nb::cast(*submission.ticket) + : nb::none(); + return nb::make_tuple( + ticket, + nb::make_tuple(std::get<0>(submission.arrays), + std::get<1>(submission.arrays), + std::get<2>(submission.arrays))); + }, + "provider"_a, + "source"_a, + "hits"_a, + "misses"_a, + "Copy the fixed handoff and enqueue MLX-owned U32/BF16 planes."); + + m.def( + "drain_cached_completions", + [](const std::shared_ptr< + mtplx_native::ple_cpu_rows::CachedSidecarProducer>& producer) { + nb::list output; + for (const auto& completion : + mtplx_native::ple_cpu_rows::drain_cached_completions(producer)) { + output.append(nb::make_tuple( + completion.ticket, owned_packed_completion(completion))); + } + return output; + }, + "provider"_a, + "Drain owner-thread miss completions as immutable uint8[M,100] arrays."); + + m.def( + "make_sidecar_rows", + [](const std::shared_ptr& + producer, + const std::array& previous, + const std::array& ids) { + std::shared_ptr + installed = producer; + return mtplx_native::ple_cpu_rows::make_sidecar_rows( + installed, previous, ids); + }, + "provider"_a, + "previous"_a, + "ids"_a, + "Create fresh sidecar [64,20] U32 and two [64,5] BF16 planes."); +} diff --git a/native_extensions/ple_cpu_rows/cached_sidecar_primitive.cpp b/native_extensions/ple_cpu_rows/cached_sidecar_primitive.cpp new file mode 100644 index 000000000..c9ce5d853 --- /dev/null +++ b/native_extensions/ple_cpu_rows/cached_sidecar_primitive.cpp @@ -0,0 +1,153 @@ +// SPDX-License-Identifier: Apache-2.0 + +#include "cached_sidecar_primitive.h" + +#include +#include +#include +#include +#include + +#include "mlx/allocator.h" +#include "mlx/backend/cpu/encoder.h" +#include "mlx/primitives.h" + +namespace mtplx_native::ple_cpu_rows { + +namespace mx = mlx::core; + +// Defined by ple_cpu_rows.cpp. The synthetic and cached provider primitives +// deliberately share this one process-lifetime CPU stream and permit pool. +mx::Stream installed_cpu_stream(); +std::shared_ptr installed_permits(); + +namespace { + +using mx::array; + +void copy_packed_to_planes(const host::PackedRows& packed, + array& weight_output, + array& scales_output, + array& bias_output) { + auto* weight_dst = weight_output.data(); + auto* scales_dst = scales_output.data(); + auto* bias_dst = bias_output.data(); + for (std::size_t row = 0; row < kRows; ++row) { + const auto* source = packed.data() + row * 100; + std::memcpy(weight_dst + row * kWeightValuesPerRow, source, 80); + std::memcpy(scales_dst + row * kMetadataValuesPerRow, source + 80, 10); + std::memcpy(bias_dst + row * kMetadataValuesPerRow, source + 90, 10); + } +} + +class PleCachedSidecarRowsPrimitive final : public mx::Primitive { + public: + PleCachedSidecarRowsPrimitive(mx::Stream cpu_stream, + std::shared_ptr job) + : mx::Primitive(cpu_stream), job_(std::move(job)) {} + + ~PleCachedSidecarRowsPrimitive() override = default; + + void eval_cpu(const std::vector&, + std::vector& outputs) override { + auto& weight_output = outputs[0]; + auto& scales_output = outputs[1]; + auto& bias_output = outputs[2]; + + // These are MLX allocator buffers. The queued lambda captures array + // descriptors, retaining all three buffers through row reads and copies. + weight_output.set_data(mx::allocator::malloc(kWeightBytes)); + scales_output.set_data(mx::allocator::malloc(kMetadataBytes)); + bias_output.set_data(mx::allocator::malloc(kMetadataBytes)); + + auto weight = weight_output; + auto scales = scales_output; + auto bias = bias_output; + auto job = std::move(job_); + auto task = [job = std::move(job), + weight = std::move(weight), + scales = std::move(scales), + bias = std::move(bias)]() mutable { + try { + // A failed read must not expose partially populated planes. The + // scheduler propagates the exception; zeroing is not success. + std::memset(weight.data(), 0, kWeightBytes); + std::memset(scales.data(), 0, kMetadataBytes); + std::memset(bias.data(), 0, kMetadataBytes); + const auto packed = job->run(); + copy_packed_to_planes(packed, weight, scales, bias); + } catch (...) { + // The ordinary MLX CPU scheduler owns propagation. The job's state + // makes permit release idempotent on both success and failure. + job->release_permit(); + throw; + } + // This is the outer publication boundary: all three MLX copies have + // completed before the job's permit can be returned. + job->release_permit(); + }; + + auto& encoder = mx::cpu::get_command_encoder(stream()); + encoder.set_output_array(weight_output); + encoder.set_output_array(scales_output); + encoder.set_output_array(bias_output); + encoder.dispatch(std::move(task)); + } + + void eval_gpu(const std::vector&, + std::vector&) override { + throw std::runtime_error( + "[mtplx_native_ple_cpu_rows] cached sidecar primitive is CPU-stream-only"); + } + + const char* name() const override { return "MtplxCachedSidecarRows"; } + + std::vector output_shapes( + const std::vector&) override { + return {mx::Shape{64, 20}, mx::Shape{64, 5}, mx::Shape{64, 5}}; + } + + bool is_equivalent(const mx::Primitive&) const override { + // Each handoff owns an independent ticket, snapshot, and permit. + return false; + } + + private: + std::shared_ptr job_; +}; + +} // namespace + +CachedRowsSubmission make_cached_sidecar_rows( + const std::shared_ptr& producer, + const CachedRowHandoff& handoff) { + if (producer == nullptr) { + throw std::invalid_argument( + "[mtplx_native_ple_cpu_rows] cached sidecar producer is null"); + } + auto job = producer->make_job(handoff, installed_permits()); + std::optional ticket; + if (handoff.miss_count != 0) ticket = job->ticket(); + + auto primitive = std::make_shared( + installed_cpu_stream(), std::move(job)); + auto outputs = mx::array::make_arrays( + {mx::Shape{64, 20}, mx::Shape{64, 5}, mx::Shape{64, 5}}, + {mx::uint32, mx::bfloat16, mx::bfloat16}, + primitive, + std::vector{}); + return CachedRowsSubmission{ + std::move(ticket), + CachedRowsArrays{outputs.at(0), outputs.at(1), outputs.at(2)}}; +} + +std::vector drain_cached_completions( + const std::shared_ptr& producer) { + if (producer == nullptr) { + throw std::invalid_argument( + "[mtplx_native_ple_cpu_rows] cached sidecar producer is null"); + } + return producer->drain_completed(); +} + +} // namespace mtplx_native::ple_cpu_rows diff --git a/native_extensions/ple_cpu_rows/cached_sidecar_primitive.h b/native_extensions/ple_cpu_rows/cached_sidecar_primitive.h new file mode 100644 index 000000000..f9b6b841e --- /dev/null +++ b/native_extensions/ple_cpu_rows/cached_sidecar_primitive.h @@ -0,0 +1,36 @@ +// SPDX-License-Identifier: Apache-2.0 + +#pragma once + +#include +#include +#include +#include +#include + +#include "cached_sidecar_producer.h" +#include "mlx/array.h" +#include "mlx/stream.h" + +namespace mtplx_native::ple_cpu_rows { + +namespace mx = mlx::core; + +using CachedRowsArrays = std::tuple; + +// The ticket is present only when misses were admitted. The arrays are +// ordinary MLX-owned U32/BF16 planes; the primitive copies the native packed +// rows into them on the installed CPU stream before the caller dequantizes. +struct CachedRowsSubmission final { + std::optional ticket; + CachedRowsArrays arrays; +}; + +CachedRowsSubmission make_cached_sidecar_rows( + const std::shared_ptr& producer, + const CachedRowHandoff& handoff); + +std::vector drain_cached_completions( + const std::shared_ptr& producer); + +} // namespace mtplx_native::ple_cpu_rows diff --git a/native_extensions/ple_cpu_rows/cached_sidecar_producer.cpp b/native_extensions/ple_cpu_rows/cached_sidecar_producer.cpp new file mode 100644 index 000000000..031d1ccaf --- /dev/null +++ b/native_extensions/ple_cpu_rows/cached_sidecar_producer.cpp @@ -0,0 +1,491 @@ +// SPDX-License-Identifier: Apache-2.0 + +#include "cached_sidecar_producer.h" + +#include +#include +#include +#include +#include + +namespace mtplx_native::ple_cpu_rows { + +namespace { + +class RawReaderAdapter final : public CachedRowReader { + public: + explicit RawReaderAdapter(std::shared_ptr base) + : base_(std::move(base)) {} + + host::PackedRow read_one(std::uint32_t row_id) const override { + return base_->reader().read_one(row_id); + } + + private: + const std::shared_ptr base_; +}; + +} // namespace + +struct CachedSidecarJob::State final { + explicit State(std::shared_ptr permits) + : permits(std::move(permits)) {} + + ~State() { + // A producer may be torn down with undrained completion entries. The + // entry owns this state, so its final destruction is the last safe place + // to return a permit that the owner never drained. + std::lock_guard lock(mutex); + release_locked(); + } + + void start() { + std::lock_guard lock(mutex); + if (abandoned || started) { + throw std::runtime_error("cached sidecar job was already started or abandoned"); + } + started = true; + } + + bool mark_finished(bool has_completion) noexcept { + std::lock_guard lock(mutex); + run_finished = true; + completion_pending = has_completion && !abandoned; + if (abandoned) { + release_locked(); + } else { + release_if_ready_locked(); + } + return !abandoned; + } + + void mark_failed() noexcept { + std::lock_guard lock(mutex); + run_finished = true; + failed = true; + completion_pending = false; + release_locked(); + } + + void mark_output_published() noexcept { + std::lock_guard lock(mutex); + output_published = true; + release_if_ready_locked(); + } + + void mark_completion_drained() noexcept { + std::lock_guard lock(mutex); + completion_drained = true; + completion_pending = false; + release_if_ready_locked(); + } + + void abandon() noexcept { + std::lock_guard lock(mutex); + abandoned = true; + // A pre-dispatch job has no worker-owned work. Once run() has started, + // keep its permit until mark_finished() observes that every submitted + // read has been joined; otherwise repeated abandon() calls could admit + // more than two active I/O jobs. + if (!started || run_finished) { + release_locked(); + } + } + + void on_job_destroyed() noexcept { + std::lock_guard lock(mutex); + if (!started) { + abandoned = true; + release_locked(); + return; + } + if (!run_finished) { + abandoned = true; + return; + } + if (abandoned || failed) { + return; + } + if (!output_published) { + // run() completed, but the caller never published its output. Any + // queued completion is stale and will be purged by the owner drain. + abandoned = true; + completion_pending = false; + release_locked(); + return; + } + // A published miss completion must remain owned by the producer after + // the job handle dies. The owner drain is the only release boundary. + release_if_ready_locked(); + } + + bool is_abandoned() const noexcept { + std::lock_guard lock(mutex); + return abandoned; + } + + private: + void release_if_ready_locked() noexcept { + if (!run_finished || !output_published || failed) { + return; + } + if (completion_pending && !completion_drained) { + return; + } + release_locked(); + } + + void release_locked() noexcept { + if (permit_held) { + permit_held = false; + permits->release(); + } + } + + const std::shared_ptr permits; + mutable std::mutex mutex; + bool started = false; + bool run_finished = false; + bool output_published = false; + bool completion_pending = false; + bool completion_drained = false; + bool failed = false; + bool abandoned = false; + bool permit_held = true; +}; + +struct CachedSidecarProducer::CompletionEntry final { + CachedCompletion completion{}; + std::shared_ptr state; +}; + +class CachedSidecarProducer::IoPool final { + public: + explicit IoPool(std::size_t worker_count) { + workers_.reserve(worker_count); + try { + for (std::size_t index = 0; index < worker_count; ++index) { + workers_.emplace_back([this] { worker_loop(); }); + } + } catch (...) { + { + std::lock_guard lock(mutex_); + stopping_ = true; + } + condition_.notify_all(); + for (auto& worker : workers_) { + if (worker.joinable()) worker.join(); + } + throw; + } + } + + ~IoPool() { + { + std::lock_guard lock(mutex_); + stopping_ = true; + } + condition_.notify_all(); + for (auto& worker : workers_) { + if (worker.joinable()) worker.join(); + } + } + + IoPool(const IoPool&) = delete; + IoPool& operator=(const IoPool&) = delete; + + template + auto submit(Function&& function) + -> std::future> { + using Result = std::invoke_result_t; + auto task = std::make_shared>( + std::forward(function)); + auto future = task->get_future(); + { + std::lock_guard lock(mutex_); + if (stopping_) { + throw std::runtime_error("cached sidecar I/O pool is stopping"); + } + tasks_.emplace_back([task] { (*task)(); }); + } + condition_.notify_one(); + return future; + } + + private: + void worker_loop() noexcept { + for (;;) { + std::function task; + { + std::unique_lock lock(mutex_); + condition_.wait(lock, [this] { return stopping_ || !tasks_.empty(); }); + if (tasks_.empty()) { + if (stopping_) return; + continue; + } + task = std::move(tasks_.front()); + tasks_.pop_front(); + } + task(); + } + } + + std::mutex mutex_; + std::condition_variable condition_; + std::deque> tasks_; + std::vector workers_; + bool stopping_ = false; +}; + +CachedSidecarJob::CachedSidecarJob( + std::shared_ptr producer, + CachedRowHandoff handoff, + std::shared_ptr permits, + std::uint64_t ticket) + : producer_(std::move(producer)), + handoff_(std::move(handoff)), + ticket_(ticket), + state_(std::make_shared(std::move(permits))) {} + +CachedSidecarJob::~CachedSidecarJob() { state_->on_job_destroyed(); } + +void CachedSidecarJob::release_permit() noexcept { + state_->mark_output_published(); +} + +void CachedSidecarJob::abandon() noexcept { state_->abandon(); } + +host::PackedRows CachedSidecarJob::run() { + state_->start(); + + std::vector> futures; + futures.reserve(handoff_.miss_count); + std::exception_ptr first_error; + try { + for (std::size_t index = 0; index < handoff_.miss_count; ++index) { + const std::uint32_t row_id = handoff_.miss_ids[index]; + const auto reader = producer_->reader_; + futures.emplace_back(producer_->io_pool_->submit( + [reader, row_id] { return reader->read_one(row_id); })); + } + } catch (...) { + first_error = std::current_exception(); + } + + if (first_error) { + // packaged_task futures do not join on destruction. Drain every task + // already submitted before releasing this job's permit. + for (auto& future : futures) { + try { + (void)future.get(); + } catch (...) { + // Preserve the submission failure as the ordinary CPU error. + } + } + state_->mark_failed(); + std::rethrow_exception(first_error); + } + + std::array miss_payloads{}; + for (std::size_t index = 0; index < futures.size(); ++index) { + try { + miss_payloads[index] = futures[index].get(); + } catch (...) { + if (!first_error) first_error = std::current_exception(); + } + } + if (first_error) { + state_->mark_failed(); + std::rethrow_exception(first_error); + } + + host::PackedRows output{}; + for (std::size_t row = 0; row < host::kRowsPerWindow; ++row) { + const std::uint8_t source = handoff_.source[row]; + const std::size_t index = source & kCachedSourceMask; + const host::PackedRow& payload = + (source & kCachedHitBit) != 0 ? handoff_.hit_packed[index] + : miss_payloads[index]; + std::copy( + payload.begin(), payload.end(), + output.begin() + row * host::kPackedBytesPerRow); + } + + const bool enqueue_completion = + state_->mark_finished(handoff_.miss_count != 0); + if (handoff_.miss_count != 0 && enqueue_completion) { + CachedCompletion completion{}; + completion.ticket = ticket_; + completion.count = handoff_.miss_count; + for (std::size_t index = 0; index < handoff_.miss_count; ++index) { + completion.miss_ids[index] = handoff_.miss_ids[index]; + completion.payloads[index] = miss_payloads[index]; + } + try { + producer_->enqueue_completion(state_, std::move(completion)); + } catch (...) { + state_->mark_failed(); + throw; + } + } + return output; +} + +std::shared_ptr CachedSidecarProducer::install( + int descriptor, + host::SidecarLayout layout, + const std::array& multipliers, + const std::array& sizes, + const std::array& offsets, + std::int64_t eos, + std::size_t io_workers) { + auto base = SidecarProducer::install( + descriptor, layout, multipliers, sizes, offsets, eos); + auto reader = std::make_shared(base); + return std::shared_ptr(new CachedSidecarProducer( + std::move(base), std::move(reader), io_workers)); +} + +std::shared_ptr +CachedSidecarProducer::install_for_test( + int descriptor, + host::SidecarLayout layout, + const std::array& multipliers, + const std::array& sizes, + const std::array& offsets, + std::int64_t eos, + std::shared_ptr reader, + std::size_t io_workers) { + if (reader == nullptr) { + throw std::invalid_argument("cached sidecar test reader is null"); + } + auto base = SidecarProducer::install( + descriptor, layout, multipliers, sizes, offsets, eos); + return std::shared_ptr(new CachedSidecarProducer( + std::move(base), std::move(reader), io_workers)); +} + +CachedSidecarProducer::CachedSidecarProducer( + std::shared_ptr base, + std::shared_ptr reader, + std::size_t io_workers) + : base_(std::move(base)), + reader_(std::move(reader)), + io_pool_(nullptr), + io_workers_(io_workers), + completions_(std::make_unique>()) { + if (base_ == nullptr || reader_ == nullptr) { + throw std::invalid_argument("cached sidecar producer dependencies are null"); + } + if (io_workers_ == 0 || io_workers_ > kCachedMaxIoWorkers) { + throw std::invalid_argument("cached sidecar I/O worker count is out of bounds"); + } + io_pool_ = std::make_unique(io_workers_); +} + +CachedSidecarProducer::~CachedSidecarProducer() = default; + +std::uint64_t CachedSidecarProducer::row_count() const noexcept { + return base_->row_count(); +} + +SidecarRowIds CachedSidecarProducer::compute_row_ids( + const SidecarJobInput& input) const { + const host::NgramRowsResult result = + base_->plan().compute(input.previous, input.ids); + SidecarRowIds rows{}; + for (std::size_t index = 0; index < rows.size(); ++index) { + rows[index] = static_cast(result.rows[index]); + } + return rows; +} + +void CachedSidecarProducer::validate_handoff( + const CachedRowHandoff& handoff, + std::uint64_t row_count) { + if (handoff.hit_count > host::kRowsPerWindow || + handoff.miss_count > host::kRowsPerWindow) { + throw std::invalid_argument("cached sidecar handoff count exceeds M4 bound"); + } + for (std::size_t index = 0; index < handoff.source.size(); ++index) { + const std::uint8_t source = handoff.source[index]; + if ((source & 0x40U) != 0) { + throw std::invalid_argument( + "cached sidecar source uses reserved bit 6"); + } + const std::size_t compact_index = source & kCachedSourceMask; + const bool hit = (source & kCachedHitBit) != 0; + const std::size_t limit = hit ? handoff.hit_count : handoff.miss_count; + if (compact_index >= limit) { + throw std::invalid_argument("cached sidecar source index is out of bounds"); + } + } + for (std::size_t index = 0; index < handoff.miss_count; ++index) { + if (static_cast(handoff.miss_ids[index]) >= row_count) { + throw std::out_of_range("cached sidecar miss row exceeds row count"); + } + for (std::size_t prior = 0; prior < index; ++prior) { + if (handoff.miss_ids[prior] == handoff.miss_ids[index]) { + throw std::invalid_argument("cached sidecar miss IDs must be unique"); + } + } + } +} + +std::shared_ptr CachedSidecarProducer::make_job( + const CachedRowHandoff& handoff, + const std::shared_ptr& permits) { + if (permits == nullptr) { + throw std::invalid_argument("cached sidecar permit pool is null"); + } + validate_handoff(handoff, row_count()); + if (!permits->try_acquire()) { + throw std::runtime_error("[mtplx_native_ple_cpu_rows] bounded two-request queue is full"); + } + PermitLease lease(permits); + const std::uint64_t ticket = next_ticket_.fetch_add(1, std::memory_order_relaxed); + auto owned = std::unique_ptr(new CachedSidecarJob( + shared_from_this(), handoff, permits, ticket)); + lease.disarm(); + return std::shared_ptr(std::move(owned)); +} + +void CachedSidecarProducer::enqueue_completion( + const std::shared_ptr& state, + CachedCompletion completion) { + std::lock_guard lock(completion_mutex_); + // An explicit abandon may race the final output copy. In that case the + // permit is already released (after the reads have joined), and no stale + // completion should become visible to the owner. + if (state->is_abandoned()) return; + completions_->erase( + std::remove_if( + completions_->begin(), completions_->end(), + [](const CompletionEntry& entry) { + return entry.state == nullptr || entry.state->is_abandoned(); + }), + completions_->end()); + if (completions_->size() >= kMaxOutstanding) { + throw std::runtime_error("cached sidecar completion bound is full"); + } + if (state->is_abandoned()) return; + completions_->push_back(CompletionEntry{std::move(completion), state}); +} + +std::vector CachedSidecarProducer::drain_completed() { + std::deque pending; + { + std::lock_guard lock(completion_mutex_); + pending.swap(*completions_); + } + std::vector output; + output.reserve(pending.size()); + for (auto& entry : pending) { + if (entry.state == nullptr || entry.state->is_abandoned()) continue; + entry.state->mark_completion_drained(); + output.push_back(std::move(entry.completion)); + } + return output; +} + +} // namespace mtplx_native::ple_cpu_rows diff --git a/native_extensions/ple_cpu_rows/cached_sidecar_producer.h b/native_extensions/ple_cpu_rows/cached_sidecar_producer.h new file mode 100644 index 000000000..301a0211a --- /dev/null +++ b/native_extensions/ple_cpu_rows/cached_sidecar_producer.h @@ -0,0 +1,182 @@ +// SPDX-License-Identifier: Apache-2.0 +// +// MLX-free cache-aware CPU sidecar handoff. The stock Python owner remains +// the only persistent row cache. This component owns only bounded immutable +// handoff/completion bytes and a construction-time fixed I/O pool. + +#pragma once + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include "host_provider.h" +#include "request_state.h" +#include "sidecar_producer.h" + +namespace mtplx_native::ple_cpu_rows { + +namespace host = mtplx_native::host_provider; + +constexpr std::uint8_t kCachedHitBit = 0x80; +constexpr std::uint8_t kCachedSourceMask = 0x3f; +constexpr std::size_t kCachedDefaultIoWorkers = 8; +constexpr std::size_t kCachedMaxIoWorkers = 16; + +// The source byte is either kCachedHitBit | compact hit index or a compact +// miss index. hit_packed contains only the compact hit rows; miss_ids are +// unique and are read by the native fixed pool. All arrays are fixed-bounded +// so Python cannot make the queued handoff grow without limit. +struct CachedRowHandoff final { + std::array source{}; + std::array hit_packed{}; + std::array miss_ids{}; + std::uint8_t hit_count = 0; + std::uint8_t miss_count = 0; +}; + +// A completion is tied to the admitted job ticket. The owner thread drains +// this record and publishes it into the stock Python LRU; worker code never +// touches Python cache objects. +struct CachedCompletion final { + std::uint64_t ticket = 0; + std::array miss_ids{}; + std::array payloads{}; + std::uint8_t count = 0; +}; + +// Test-only reader seam. Production installation wraps the immutable +// RawSidecarBatchReader; CPU tests inject a counting reader without adding +// counters or branches to the production read path. +class CachedRowReader { + public: + virtual ~CachedRowReader() = default; + virtual host::PackedRow read_one(std::uint32_t row_id) const = 0; +}; + +class CachedSidecarProducer; + +class CachedSidecarJob final { + public: + ~CachedSidecarJob(); + + CachedSidecarJob(const CachedSidecarJob&) = delete; + CachedSidecarJob& operator=(const CachedSidecarJob&) = delete; + + // Computes no hash: row IDs and the typed hit/miss scatter map were copied + // at admission. The returned packed bytes preserve all 64 source slots. + host::PackedRows run(); + + // Marks MLX/output publication complete. If misses produced a completion, + // the permit remains held until CachedSidecarProducer::drain_completed(). + void release_permit() noexcept; + + // Explicitly abandons an admitted job before publication. Abandonment + // drops any queued completion ownership and returns its permit once. + void abandon() noexcept; + + std::uint64_t ticket() const noexcept { return ticket_; } + + private: + friend class CachedSidecarProducer; + + struct State; + + CachedSidecarJob( + std::shared_ptr producer, + CachedRowHandoff handoff, + std::shared_ptr permits, + std::uint64_t ticket); + + const std::shared_ptr producer_; + const CachedRowHandoff handoff_; + const std::uint64_t ticket_; + const std::shared_ptr state_; +}; + +class CachedSidecarProducer final + : public std::enable_shared_from_this { + public: + ~CachedSidecarProducer(); + + static std::shared_ptr install( + int descriptor, + host::SidecarLayout layout, + const std::array& multipliers, + const std::array& sizes, + const std::array& offsets, + std::int64_t eos, + std::size_t io_workers = kCachedDefaultIoWorkers); + + // CPU-test seam: the descriptor/layout still pass the normal construction + // validation through SidecarProducer; only row reads are instrumented. + static std::shared_ptr install_for_test( + int descriptor, + host::SidecarLayout layout, + const std::array& multipliers, + const std::array& sizes, + const std::array& offsets, + std::int64_t eos, + std::shared_ptr reader, + std::size_t io_workers = kCachedDefaultIoWorkers); + + CachedSidecarProducer(const CachedSidecarProducer&) = delete; + CachedSidecarProducer& operator=(const CachedSidecarProducer&) = delete; + + // Construction-bound plan reuse for the synchronous owner-thread hash. + SidecarRowIds compute_row_ids(const SidecarJobInput& input) const; + + std::uint64_t row_count() const noexcept; + std::size_t io_workers() const noexcept { return io_workers_; } + + // Typed admission validates source/count/index bounds and unique miss IDs + // once. No corresponding checks occur in worker read loops. + std::shared_ptr make_job( + const CachedRowHandoff& handoff, + const std::shared_ptr& permits); + + // Draining is the owner-thread completion boundary. For a miss-bearing + // job it is what permits the job's two-request slot to be reclaimed after + // output publication. Abandoned entries are discarded here. + std::vector drain_completed(); + + private: + friend class CachedSidecarJob; + + class IoPool; + struct CompletionEntry; + + CachedSidecarProducer( + std::shared_ptr base, + std::shared_ptr reader, + std::size_t io_workers); + + static void validate_handoff( + const CachedRowHandoff& handoff, + std::uint64_t row_count); + + void enqueue_completion( + const std::shared_ptr& state, + CachedCompletion completion); + + std::shared_ptr base_; + std::shared_ptr reader_; + std::unique_ptr io_pool_; + std::size_t io_workers_ = 0; + std::atomic next_ticket_{1}; + mutable std::mutex completion_mutex_; + std::unique_ptr> completions_; +}; + +} // namespace mtplx_native::ple_cpu_rows diff --git a/native_extensions/ple_cpu_rows/host_provider.cpp b/native_extensions/ple_cpu_rows/host_provider.cpp new file mode 100644 index 000000000..0edf4ea29 --- /dev/null +++ b/native_extensions/ple_cpu_rows/host_provider.cpp @@ -0,0 +1,398 @@ +// SPDX-License-Identifier: Apache-2.0 + +#include "host_provider.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace mtplx_native::host_provider { + +namespace { + +using U64 = std::uint64_t; +using I64 = std::int64_t; + +class ScopedFd final { + public: + explicit ScopedFd(int descriptor) noexcept : descriptor_(descriptor) {} + ~ScopedFd() { + if (descriptor_ >= 0) { + ::close(descriptor_); + } + } + + ScopedFd(const ScopedFd&) = delete; + ScopedFd& operator=(const ScopedFd&) = delete; + + int get() const noexcept { return descriptor_; } + + int release() noexcept { + const int descriptor = descriptor_; + descriptor_ = -1; + return descriptor; + } + + private: + int descriptor_ = -1; +}; + +U64 bits_of(I64 value) noexcept { + U64 bits = 0; + static_assert(sizeof(bits) == sizeof(value)); + std::memcpy(&bits, &value, sizeof(bits)); + return bits; +} + +I64 value_of(U64 bits) noexcept { + I64 value = 0; + static_assert(sizeof(bits) == sizeof(value)); + std::memcpy(&value, &bits, sizeof(value)); + return value; +} + +U64 checked_add(U64 left, U64 right, const char* what) { + if (right > std::numeric_limits::max() - left) { + throw std::overflow_error(std::string(what) + " addition overflow"); + } + return left + right; +} + +U64 checked_mul(U64 left, U64 right, const char* what) { + if (left != 0 && right > std::numeric_limits::max() / left) { + throw std::overflow_error(std::string(what) + " multiplication overflow"); + } + return left * right; +} + +U64 nonnegative_mod(U64 signed_bits, U64 divisor) noexcept { + if ((signed_bits >> 63) == 0) { + return signed_bits % divisor; + } + // Compute abs(INT64_MIN) in unsigned arithmetic; signed negation would be + // undefined at the exact input that the parity harness exercises. + const U64 magnitude = (~signed_bits) + U64{1}; + const U64 remainder = magnitude % divisor; + return remainder == 0 ? U64{0} : divisor - remainder; +} + +struct Range { + U64 begin; + U64 end; +}; + +Range plane_range(const PlaneSpec& plane, U64 row_count, const char* label) { + const U64 bytes = checked_mul(row_count, plane.stride, label); + if (bytes != plane.length) { + throw std::invalid_argument(std::string(label) + + " length does not match row stride"); + } + return {plane.offset, checked_add(plane.offset, plane.length, label)}; +} + +void require_nonoverlap(const Range& left, + const char* left_label, + const Range& right, + const char* right_label) { + if (left.begin < right.end && right.begin < left.end) { + throw std::invalid_argument(std::string(left_label) + " overlaps " + + right_label); + } +} + +void read_exact(int descriptor, + std::uint8_t* destination, + std::size_t length, + U64 offset) { + std::size_t completed = 0; + while (completed < length) { + // The construction-time plane range check proves offset + length is at + // most off_t::max(), so progress within this fixed read is bounded too. + const U64 current = offset + static_cast(completed); + const ssize_t count = ::pread( + descriptor, + destination + completed, + length - completed, + static_cast(current)); + if (count < 0) { + if (errno == EINTR) { + continue; + } + throw std::system_error(errno, std::generic_category(), "pread"); + } + if (count == 0) { + throw std::runtime_error("sidecar short read"); + } + completed += static_cast(count); + } +} + +} // namespace + +NgramPlan::NgramPlan( + const std::array& multipliers, + const std::array& sizes, + const std::array& offsets, + I64 eos) + : eos_(eos) { + for (std::size_t index = 0; index < kNgramSize; ++index) { + multiplier_bits_[index] = bits_of(multipliers[index]); + } + for (I64 size : sizes) { + if (size <= 0) { + throw std::invalid_argument("ngram head size must be positive"); + } + } + for (std::size_t index = 0; index < kNgramHeads; ++index) { + sizes_[index] = static_cast(sizes[index]); + offset_bits_[index] = bits_of(offsets[index]); + } +} + +NgramRowsResult NgramPlan::compute( + const std::array& previous, + const std::array& ids) const { + std::array history{}; + std::copy(previous.begin(), previous.end(), history.begin()); + std::copy(ids.begin(), ids.end(), history.begin() + kHistoryRows); + + // Record the preceding EOS before inspecting the current position. This is + // the NumPy prev_incl[:, :-1] rule and keeps a current EOS out of its own + // segment-start scan. + std::array previous_eos{}; + I64 last_eos = -1; + for (std::size_t position = 0; position < history.size(); ++position) { + previous_eos[position] = last_eos; + if (history[position] == eos_) { + last_eos = static_cast(position); + } + } + + std::array, kNgramSize> shifted{}; + for (std::size_t position = 0; position < history.size(); ++position) { + shifted[0][position] = bits_of(history[position]); + const I64 position_in_segment = + static_cast(position) - (previous_eos[position] + 1); + for (std::size_t shift = 1; shift < kNgramSize; ++shift) { + const bool has_source = position >= shift; + const bool is_valid = + has_source && position_in_segment >= static_cast(shift); + const std::size_t source = has_source ? position - shift : 0; + shifted[shift][position] = + is_valid ? bits_of(history[source]) : bits_of(eos_); + } + } + + NgramRowsResult result{}; + for (std::size_t input_position = 0; input_position < kInputRows; + ++input_position) { + const std::size_t history_position = kHistoryRows + input_position; + const std::size_t output_base = input_position * kNgramHeads; + for (std::size_t ngram = 2; ngram <= kNgramSize; ++ngram) { + const std::size_t head_base = (ngram - 2) * kHeadsPerNgram; + U64 mixed = multiplier_bits_[0] * shifted[0][history_position]; + for (std::size_t part = 1; part < ngram; ++part) { + mixed ^= multiplier_bits_[part] * shifted[part][history_position]; + } + for (std::size_t head = 0; head < kHeadsPerNgram; ++head) { + const std::size_t head_index = head_base + head; + const U64 remainder = nonnegative_mod(mixed, sizes_[head_index]); + const U64 row_bits = remainder + offset_bits_[head_index]; + result.rows[output_base + head_index] = value_of(row_bits); + } + } + } + + result.history[0] = ids[kInputRows - kHistoryRows]; + result.history[1] = ids[kInputRows - 1]; + return result; +} + +NgramRowsResult ngram_rows_fixed_m4( + const std::array& previous, + const std::array& ids, + const std::array& multipliers, + const std::array& sizes, + const std::array& offsets, + I64 eos) { + const NgramPlan plan(multipliers, sizes, offsets, eos); + return plan.compute(previous, ids); +} + +void check_sidecar_layout(const SidecarLayout& layout, U64 file_size) { + const U64 off_t_max = static_cast(std::numeric_limits::max()); + if (file_size > off_t_max) { + throw std::overflow_error("sidecar file size exceeds off_t"); + } + if (layout.row_count == 0) { + throw std::invalid_argument("sidecar row count must be positive"); + } + if (layout.weights.stride != kWeightBytesPerRow) { + throw std::invalid_argument("weights stride must be 80 bytes"); + } + if (layout.scales.stride != kScaleBytesPerRow) { + throw std::invalid_argument("scales stride must be 10 bytes"); + } + if (layout.biases.stride != kBiasBytesPerRow) { + throw std::invalid_argument("biases stride must be 10 bytes"); + } + + const Range weights = plane_range(layout.weights, layout.row_count, "weights"); + const Range scales = plane_range(layout.scales, layout.row_count, "scales"); + const Range biases = plane_range(layout.biases, layout.row_count, "biases"); + for (const auto& range : {weights, scales, biases}) { + if (range.end > off_t_max) { + throw std::overflow_error("sidecar plane exceeds off_t"); + } + if (range.end > file_size) { + throw std::out_of_range("sidecar plane exceeds file size"); + } + } + require_nonoverlap(weights, "weights", scales, "scales"); + require_nonoverlap(weights, "weights", biases, "biases"); + require_nonoverlap(scales, "scales", biases, "biases"); +} + +U64 checked_plane_row_offset(const PlaneSpec& plane, + U64 row_id, + U64 row_count) { + if (row_count == 0 || row_id >= row_count) { + throw std::out_of_range("sidecar row ID is outside row count"); + } + return checked_add( + plane.offset, + checked_mul(row_id, plane.stride, "row offset"), + "row offset"); +} + +RawSidecarBatchReader::RawSidecarBatchReader(int descriptor, + SidecarLayout layout) + : layout_(layout) { + if (descriptor < 0) { + throw std::invalid_argument("sidecar descriptor must be nonnegative"); + } + ScopedFd owned(::dup(descriptor)); + if (owned.get() < 0) { + throw std::system_error(errno, std::generic_category(), "dup"); + } + const int flags = ::fcntl(owned.get(), F_GETFD); + if (flags < 0) { + throw std::system_error(errno, std::generic_category(), "fcntl(F_GETFD)"); + } + if (::fcntl(owned.get(), F_SETFD, flags | FD_CLOEXEC) != 0) { + throw std::system_error( + errno, std::generic_category(), "fcntl(FD_CLOEXEC)"); + } + struct stat info {}; + if (::fstat(owned.get(), &info) != 0) { + throw std::system_error(errno, std::generic_category(), "fstat"); + } + if (!S_ISREG(info.st_mode)) { + throw std::invalid_argument("sidecar descriptor must refer to a regular file"); + } + if (info.st_size < 0) { + throw std::runtime_error("sidecar file size is negative"); + } + file_size_ = static_cast(info.st_size); + if (file_size_ > static_cast(std::numeric_limits::max())) { + throw std::overflow_error("sidecar file size exceeds off_t"); + } + check_sidecar_layout(layout_, file_size_); + descriptor_ = owned.release(); +} + +RawSidecarBatchReader::~RawSidecarBatchReader() { + if (descriptor_ >= 0) { + ::close(descriptor_); + } +} + +RawSidecarBatchReader::RawSidecarBatchReader( + RawSidecarBatchReader&& other) noexcept + : descriptor_(other.descriptor_), + layout_(other.layout_), + file_size_(other.file_size_) { + other.descriptor_ = -1; + other.file_size_ = 0; +} + +RawSidecarBatchReader& RawSidecarBatchReader::operator=( + RawSidecarBatchReader&& other) noexcept { + if (this == &other) { + return *this; + } + if (descriptor_ >= 0) { + ::close(descriptor_); + } + descriptor_ = other.descriptor_; + layout_ = other.layout_; + file_size_ = other.file_size_; + other.descriptor_ = -1; + other.file_size_ = 0; + return *this; +} + +void RawSidecarBatchReader::read_one_into( + std::uint32_t row_id, + std::uint8_t* destination) const { + const U64 row = row_id; + if (row >= layout_.row_count) { + throw std::out_of_range("sidecar row ID is outside row count"); + } + // check_sidecar_layout proved each complete plane range fits in both the + // file and off_t, so these direct operations cannot overflow for a valid + // row. Keep the single variable row-ID check above outside the three + // fixed plane reads. + const U64 weight_offset = + layout_.weights.offset + row * layout_.weights.stride; + const U64 scale_offset = + layout_.scales.offset + row * layout_.scales.stride; + const U64 bias_offset = + layout_.biases.offset + row * layout_.biases.stride; + read_exact(descriptor_, destination, kWeightBytesPerRow, weight_offset); + read_exact( + descriptor_, destination + kWeightBytesPerRow, kScaleBytesPerRow, + scale_offset); + read_exact( + descriptor_, + destination + kWeightBytesPerRow + kScaleBytesPerRow, + kBiasBytesPerRow, + bias_offset); +} + +PackedRow RawSidecarBatchReader::read_one(std::uint32_t row_id) const { + PackedRow output{}; + read_one_into(row_id, output.data()); + return output; +} + +std::vector RawSidecarBatchReader::read_subset( + const std::vector& row_ids) const { + if (row_ids.size() > kRowsPerWindow) { + throw std::invalid_argument("sidecar row subset exceeds one M4 window"); + } + std::vector output(row_ids.size()); + for (std::size_t index = 0; index < row_ids.size(); ++index) { + read_one_into(row_ids[index], output[index].data()); + } + return output; +} + +PackedRows RawSidecarBatchReader::read_rows(const RowIds& row_ids) const { + PackedRows output{}; + for (std::size_t index = 0; index < row_ids.size(); ++index) { + read_one_into( + row_ids[index], output.data() + index * kPackedBytesPerRow); + } + return output; +} + +} // namespace mtplx_native::host_provider diff --git a/native_extensions/ple_cpu_rows/host_provider.h b/native_extensions/ple_cpu_rows/host_provider.h new file mode 100644 index 000000000..509f8b744 --- /dev/null +++ b/native_extensions/ple_cpu_rows/host_provider.h @@ -0,0 +1,127 @@ +// SPDX-License-Identifier: Apache-2.0 + +#pragma once + +#include +#include +#include +#include + +namespace mtplx_native::host_provider { + +constexpr std::size_t kHistoryRows = 2; +constexpr std::size_t kInputRows = 4; +constexpr std::size_t kNgramSize = 3; +constexpr std::size_t kHeadsPerNgram = 8; +constexpr std::size_t kNgramHeads = 16; +constexpr std::size_t kNgramRowsPerWindow = kInputRows * kNgramHeads; + +constexpr std::size_t kRowsPerWindow = 64; +constexpr std::size_t kWeightBytesPerRow = 80; +constexpr std::size_t kScaleBytesPerRow = 10; +constexpr std::size_t kBiasBytesPerRow = 10; +constexpr std::size_t kPackedBytesPerRow = + kWeightBytesPerRow + kScaleBytesPerRow + kBiasBytesPerRow; +constexpr std::size_t kPackedWindowBytes = + kRowsPerWindow * kPackedBytesPerRow; + +static_assert(kNgramRowsPerWindow == kRowsPerWindow); +static_assert(kPackedBytesPerRow == 100); +static_assert(kPackedWindowBytes == 6400); + +struct NgramRowsResult { + std::array rows{}; + std::array history{}; +}; + +// Construction binds the invariant hash parameters once. compute() accepts +// only the values that change for a window and performs no plan validation. +class NgramPlan final { + public: + NgramPlan(const std::array& multipliers, + const std::array& sizes, + const std::array& offsets, + std::int64_t eos); + + NgramRowsResult compute( + const std::array& previous, + const std::array& ids) const; + + private: + std::array multiplier_bits_{}; + std::array sizes_{}; + std::array offset_bits_{}; + std::int64_t eos_ = 0; +}; + +// One fixed-M4 call: two history IDs plus four new IDs produce 64 row IDs, +// followed by the last two IDs needed by the next window. This convenience +// wrapper constructs a plan; an installed route should retain NgramPlan. +NgramRowsResult ngram_rows_fixed_m4( + const std::array& previous, + const std::array& ids, + const std::array& multipliers, + const std::array& sizes, + const std::array& offsets, + std::int64_t eos); + +struct PlaneSpec { + std::uint64_t offset = 0; + std::uint64_t length = 0; + std::uint64_t stride = 0; +}; + +struct SidecarLayout { + std::uint64_t row_count = 0; + PlaneSpec weights{}; + PlaneSpec scales{}; + PlaneSpec biases{}; +}; + +// Pure layout checks used by construction and by the bounded CPU harness. +// file_size is supplied separately so overflow and range cases need no large +// physical fixture. +void check_sidecar_layout(const SidecarLayout& layout, + std::uint64_t file_size); + +std::uint64_t checked_plane_row_offset(const PlaneSpec& plane, + std::uint64_t row_id, + std::uint64_t row_count); + +using PackedRows = std::array; +using PackedRow = std::array; +using RowIds = std::array; + +// Synchronously reads one bounded 64-row batch from three regular-file +// planes. The constructor duplicates the caller's descriptor and owns that +// duplicate; the original may be closed before read_rows is called. +class RawSidecarBatchReader final { + public: + RawSidecarBatchReader(int descriptor, SidecarLayout layout); + ~RawSidecarBatchReader(); + + RawSidecarBatchReader(const RawSidecarBatchReader&) = delete; + RawSidecarBatchReader& operator=(const RawSidecarBatchReader&) = delete; + + RawSidecarBatchReader(RawSidecarBatchReader&& other) noexcept; + RawSidecarBatchReader& operator=(RawSidecarBatchReader&& other) noexcept; + + // Read one complete packed row or an ordered bounded subset. The caller + // supplies unique IDs when it wants deduplicated I/O; this API preserves + // the supplied order and does not add reuse policy. + PackedRow read_one(std::uint32_t row_id) const; + std::vector read_subset( + const std::vector& row_ids) const; + + PackedRows read_rows(const RowIds& row_ids) const; + + private: + void read_one_into(std::uint32_t row_id, + std::uint8_t* destination) const; + + int descriptor_ = -1; + SidecarLayout layout_{}; + std::uint64_t file_size_ = 0; +}; + +} // namespace mtplx_native::host_provider diff --git a/native_extensions/ple_cpu_rows/mtplx_native_ple_cpu_rows/__init__.py b/native_extensions/ple_cpu_rows/mtplx_native_ple_cpu_rows/__init__.py new file mode 100644 index 000000000..adffafa3c --- /dev/null +++ b/native_extensions/ple_cpu_rows/mtplx_native_ple_cpu_rows/__init__.py @@ -0,0 +1,23 @@ +from ._ext import ( + CachedSidecarProducer, + SidecarProducer, + compute_cached_row_ids, + drain_cached_completions, + install_sidecar_provider, + install_cached_sidecar_provider, + make_cached_sidecar_rows, + make_cpu_rows, + make_sidecar_rows, +) + +__all__ = [ + "CachedSidecarProducer", + "SidecarProducer", + "compute_cached_row_ids", + "drain_cached_completions", + "install_sidecar_provider", + "install_cached_sidecar_provider", + "make_cached_sidecar_rows", + "make_cpu_rows", + "make_sidecar_rows", +] diff --git a/native_extensions/ple_cpu_rows/ple_cpu_rows.cpp b/native_extensions/ple_cpu_rows/ple_cpu_rows.cpp new file mode 100644 index 000000000..269892e86 --- /dev/null +++ b/native_extensions/ple_cpu_rows/ple_cpu_rows.cpp @@ -0,0 +1,180 @@ +// SPDX-License-Identifier: Apache-2.0 + +#include "ple_cpu_rows.h" + +#include +#include +#include +#include +#include +#include +#include + +#include "mlx/allocator.h" +#include "mlx/backend/cpu/encoder.h" +#include "mlx/primitives.h" + +namespace mtplx_native::ple_cpu_rows { + +namespace { + +using mx::array; + +// Installation creates this one CPU stream and one permit pool. A window +// only captures the already-installed objects; it never creates a stream or a +// worker of its own. +struct Runtime final { + Runtime() + : cpu_stream(mx::new_thread_unsafe_stream(mx::Device::cpu)), + permits(std::make_shared()) {} + + mx::Stream cpu_stream; + std::shared_ptr permits; +}; + +Runtime& runtime() { + static Runtime runtime; + return runtime; +} + +// The provider primitive is compiled in a separate translation unit so the +// established synthetic source-level contract remains independently +// reviewable. Both factories still resolve these accessors to the same +// process-lifetime stream and permit pool. +} // namespace + +mx::Stream installed_cpu_stream() { return runtime().cpu_stream; } + +std::shared_ptr installed_permits() { return runtime().permits; } + +namespace { + +class PleCpuRowsPrimitive final : public mx::Primitive { + public: + PleCpuRowsPrimitive(mx::Stream cpu_stream, + std::shared_ptr request) + : mx::Primitive(cpu_stream), request_(std::move(request)) {} + + ~PleCpuRowsPrimitive() override = default; + + void eval_cpu(const std::vector&, + std::vector& outputs) override { + auto& weight_output = outputs[0]; + auto& scales_output = outputs[1]; + auto& bias_output = outputs[2]; + + // These are MLX allocator buffers. The copied array descriptors below + // are captured by the queued lambda and therefore keep each buffer alive + // until the delayed task has copied or zeroed it. + weight_output.set_data(mx::allocator::malloc(kWeightBytes)); + scales_output.set_data(mx::allocator::malloc(kMetadataBytes)); + bias_output.set_data(mx::allocator::malloc(kMetadataBytes)); + + auto weight = weight_output; + auto scales = scales_output; + auto bias = bias_output; + auto state = std::move(request_); + + auto task = [state, + weight = std::move(weight), + scales = std::move(scales), + bias = std::move(bias)]() mutable { + try { + auto* weight_dst = weight.data(); + auto* scales_dst = scales.data(); + auto* bias_dst = bias.data(); + std::memset(weight_dst, 0, kWeightBytes); + std::memset(scales_dst, 0, kMetadataBytes); + std::memset(bias_dst, 0, kMetadataBytes); + + if (state->cancelled()) { + throw std::runtime_error( + "[mtplx_native_ple_cpu_rows] request cancelled"); + } + if (state->delay_ms() != 0) { + std::this_thread::sleep_for( + std::chrono::milliseconds(state->delay_ms())); + } + if (state->force_fail()) { + throw std::runtime_error( + "[mtplx_native_ple_cpu_rows] request failed"); + } + + const auto& payload = state->payload(); + for (std::size_t row = 0; row < kRows; ++row) { + const auto* packed = payload.data() + row * 100; + std::memcpy(weight_dst + row * kWeightValuesPerRow, packed, 80); + std::memcpy(scales_dst + row * kMetadataValuesPerRow, + packed + 80, + 10); + std::memcpy(bias_dst + row * kMetadataValuesPerRow, + packed + 90, + 10); + } + } catch (...) { + // The scheduler observes this exception on its ordinary CPU stream; + // there is no custom error pointer or hidden signal to maintain. + state->release_permit(); + throw; + } + state->release_permit(); + }; + + auto& encoder = mx::cpu::get_command_encoder(stream()); + encoder.set_output_array(weight_output); + encoder.set_output_array(scales_output); + encoder.set_output_array(bias_output); + // Do not release here if dispatch itself reports an enqueue error: the + // encoder may already own the task (notably while adding its completion + // task at the dispatch-group boundary). The last shared RequestState + // owner, either queued task or local lambda, releases exactly once. + encoder.dispatch(std::move(task)); + } + + void eval_gpu(const std::vector&, + std::vector&) override { + throw std::runtime_error( + "[mtplx_native_ple_cpu_rows] primitive is CPU-stream-only"); + } + + const char* name() const override { return "MtplxPleCpuRows"; } + + std::vector output_shapes( + const std::vector&) override { + return {mx::Shape{64, 20}, mx::Shape{64, 5}, mx::Shape{64, 5}}; + } + + bool is_equivalent(const mx::Primitive&) const override { + // Every request owns a distinct permit/payload and must remain a distinct + // graph node even if two payloads contain equal bytes. + return false; + } + + private: + std::shared_ptr request_; +}; + +// The real sidecar lane has its own explicit primitive and job state. It +// shares the installed CPU stream and permit pool above with the synthetic +// control lane, but it never falls back to that lane when provider work is +// selected. +} // namespace + +CpuRowsArrays make_cpu_rows(const PackedPayload& payload, + int delay_ms, + bool force_fail, + bool cancelled) { + auto& installed = runtime(); + auto request = RequestState::admit( + installed.permits, payload, delay_ms, force_fail, cancelled); + auto primitive = std::make_shared( + installed.cpu_stream, std::move(request)); + auto outputs = mx::array::make_arrays( + {mx::Shape{64, 20}, mx::Shape{64, 5}, mx::Shape{64, 5}}, + {mx::uint32, mx::bfloat16, mx::bfloat16}, + primitive, + std::vector{}); + return {outputs.at(0), outputs.at(1), outputs.at(2)}; +} + +} // namespace mtplx_native::ple_cpu_rows diff --git a/native_extensions/ple_cpu_rows/ple_cpu_rows.h b/native_extensions/ple_cpu_rows/ple_cpu_rows.h new file mode 100644 index 000000000..f74fe2f2f --- /dev/null +++ b/native_extensions/ple_cpu_rows/ple_cpu_rows.h @@ -0,0 +1,36 @@ +// SPDX-License-Identifier: Apache-2.0 + +#pragma once + +#include +#include + +#include "mlx/array.h" +#include "mlx/stream.h" + +#include "request_state.h" +#include "sidecar_producer.h" + +namespace mtplx_native::ple_cpu_rows { + +namespace mx = mlx::core; + +using CpuRowsArrays = std::tuple; + +// The payload is already in the exact packed row layout consumed by the +// ordinary model dequantizer: 64 rows of 80-byte U32 weights, 10-byte BF16 +// scales, and 10-byte BF16 biases. This adapter performs no dequantization. +CpuRowsArrays make_cpu_rows(const PackedPayload& payload, + int delay_ms = 0, + bool force_fail = false, + bool cancelled = false); + +// Explicit real-sidecar factory. The synthetic make_cpu_rows() control lane +// above remains independent; this factory captures an installed immutable +// SidecarProducer and a copied M4 token snapshot for one CPU-stream job. +CpuRowsArrays make_sidecar_rows( + const std::shared_ptr& producer, + const std::array& previous, + const std::array& ids); + +} // namespace mtplx_native::ple_cpu_rows diff --git a/native_extensions/ple_cpu_rows/pyproject.toml b/native_extensions/ple_cpu_rows/pyproject.toml new file mode 100644 index 000000000..d1890b5ac --- /dev/null +++ b/native_extensions/ple_cpu_rows/pyproject.toml @@ -0,0 +1,8 @@ +[build-system] +requires = [ + "setuptools>=42", + "cmake>=3.25", + "mlx>=0.32.2,<0.33", + "nanobind==2.15.0", +] +build-backend = "setuptools.build_meta" diff --git a/native_extensions/ple_cpu_rows/request_state.h b/native_extensions/ple_cpu_rows/request_state.h new file mode 100644 index 000000000..7ffd618a8 --- /dev/null +++ b/native_extensions/ple_cpu_rows/request_state.h @@ -0,0 +1,196 @@ +// SPDX-License-Identifier: Apache-2.0 +// +// MLX-free request ownership for the CPU-row primitive. This header is kept +// independent of the extension ABI so its bounded admission/lifetime rules +// can be checked with an ordinary C++17 compiler before an MLX build. + +#pragma once + +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace mtplx_native::ple_cpu_rows { + +constexpr std::size_t kRows = 64; +constexpr std::size_t kWeightValuesPerRow = 20; +constexpr std::size_t kMetadataValuesPerRow = 5; +constexpr std::size_t kWeightBytes = + kRows * kWeightValuesPerRow * sizeof(std::uint32_t); +constexpr std::size_t kMetadataBytes = + kRows * kMetadataValuesPerRow * sizeof(std::uint16_t); +constexpr std::size_t kPayloadBytes = 6400; +constexpr std::size_t kMaxOutstanding = 2; +constexpr std::size_t kPerRequestPlaneBytes = + kWeightBytes + kMetadataBytes + kMetadataBytes; + +static_assert(kWeightBytes == 5120); +static_assert(kMetadataBytes == 640); +static_assert(kWeightBytes + kMetadataBytes + kMetadataBytes == kPayloadBytes); +static_assert(kPayloadBytes == kRows * 100); +static_assert(kPerRequestPlaneBytes == 6400); + +using PackedPayload = std::array; + +// This pool never waits. Admission either acquires one of the two bounded +// request slots or fails before an MLX graph is constructed. Completion and +// abandonment return a slot through RequestState::release_permit(). +class PermitPool final { + public: + explicit PermitPool(std::size_t capacity = kMaxOutstanding) + : available_(capacity) { + if (capacity != kMaxOutstanding) { + throw std::invalid_argument("PLE CPU-row permit capacity must be two"); + } + } + + PermitPool(const PermitPool&) = delete; + PermitPool& operator=(const PermitPool&) = delete; + + bool try_acquire() noexcept { + std::lock_guard lock(mutex_); + if (available_ == 0) { + return false; + } + --available_; + return true; + } + + void release() noexcept { + std::lock_guard lock(mutex_); + // An over-release is a lifetime bug. Do not silently saturate: hiding it + // would make the two-request bound unreviewable and could over-admit work. + if (available_ >= kMaxOutstanding) { + std::terminate(); + } + ++available_; + } + + std::size_t outstanding() const noexcept { + std::lock_guard lock(mutex_); + return kMaxOutstanding - available_; + } + + private: + mutable std::mutex mutex_; + std::size_t available_; +}; + +// A construction-only transfer guard closes the allocation-failure window +// between acquiring a permit and putting the RequestState in shared storage. +// Once disarmed, the state destructor owns the same single permit. +class PermitLease final { + public: + explicit PermitLease(std::shared_ptr pool) + : pool_(std::move(pool)) {} + + PermitLease(const PermitLease&) = delete; + PermitLease& operator=(const PermitLease&) = delete; + PermitLease(PermitLease&& other) noexcept + : pool_(std::move(other.pool_)), held_(other.held_) { + other.held_ = false; + } + PermitLease& operator=(PermitLease&& other) noexcept { + if (this != &other) { + release(); + pool_ = std::move(other.pool_); + held_ = other.held_; + other.held_ = false; + } + return *this; + } + + ~PermitLease() { release(); } + + void disarm() noexcept { held_ = false; } + + private: + void release() noexcept { + if (held_) { + pool_->release(); + held_ = false; + } + } + + std::shared_ptr pool_; + bool held_ = true; +}; + +class RequestState final { + public: + // Admission is construction-bound and nonblocking. The returned state is + // the sole owner of the permit until its dispatch lambda completes or the + // primitive is abandoned before evaluation. + static std::shared_ptr admit( + const std::shared_ptr& pool, + PackedPayload payload, + int delay_ms, + bool force_fail, + bool cancelled) { + if (pool == nullptr) { + throw std::invalid_argument("PLE CPU-row permit pool is null"); + } + if (delay_ms < 0 || delay_ms > 30'000) { + throw std::invalid_argument( + "PLE CPU-row delay must be between 0 and 30000 ms"); + } + if (!pool->try_acquire()) { + throw std::runtime_error( + "[mtplx_native_ple_cpu_rows] bounded two-request queue is full"); + } + + // Keep the permit in a move-only guard until a unique owner exists. This + // makes both RequestState allocation and shared-control-block failure + // release exactly once. + PermitLease lease(pool); + auto owned = std::unique_ptr(new RequestState( + pool, std::move(payload), delay_ms, force_fail, cancelled)); + lease.disarm(); + return std::shared_ptr(std::move(owned)); + } + + ~RequestState() { release_permit(); } + + RequestState(const RequestState&) = delete; + RequestState& operator=(const RequestState&) = delete; + + const PackedPayload& payload() const noexcept { return payload_; } + int delay_ms() const noexcept { return delay_ms_; } + bool force_fail() const noexcept { return force_fail_; } + bool cancelled() const noexcept { return cancelled_; } + + // Idempotent so both the normal and exceptional dispatch paths can call it; + // the destructor also covers graph-construction abandonment. + void release_permit() noexcept { + if (permit_held_.exchange(false, std::memory_order_acq_rel)) { + pool_->release(); + } + } + + private: + RequestState(std::shared_ptr pool, + PackedPayload payload, + int delay_ms, + bool force_fail, + bool cancelled) noexcept + : pool_(std::move(pool)), + payload_(std::move(payload)), + delay_ms_(delay_ms), + force_fail_(force_fail), + cancelled_(cancelled) {} + + const std::shared_ptr pool_; + const PackedPayload payload_; + const int delay_ms_; + const bool force_fail_; + const bool cancelled_; + std::atomic_bool permit_held_{true}; +}; + +} // namespace mtplx_native::ple_cpu_rows diff --git a/native_extensions/ple_cpu_rows/setup.py b/native_extensions/ple_cpu_rows/setup.py new file mode 100644 index 000000000..8090cedb5 --- /dev/null +++ b/native_extensions/ple_cpu_rows/setup.py @@ -0,0 +1,17 @@ +from setuptools import setup + +from mlx import extension + + +if __name__ == "__main__": + setup( + name="mtplx_native_ple_cpu_rows", + version="0.0.0", + description="MLX-owned CPU-stream PLE row staging primitive for MTPLX.", + ext_modules=[extension.CMakeExtension("mtplx_native_ple_cpu_rows._ext")], + cmdclass={"build_ext": extension.CMakeBuild}, + packages=["mtplx_native_ple_cpu_rows"], + package_data={"mtplx_native_ple_cpu_rows": ["*.so", "*.dylib"]}, + zip_safe=False, + python_requires=">=3.11", + ) diff --git a/native_extensions/ple_cpu_rows/sidecar_primitive.cpp b/native_extensions/ple_cpu_rows/sidecar_primitive.cpp new file mode 100644 index 000000000..0e9d8f836 --- /dev/null +++ b/native_extensions/ple_cpu_rows/sidecar_primitive.cpp @@ -0,0 +1,139 @@ +// SPDX-License-Identifier: Apache-2.0 + +#include "ple_cpu_rows.h" + +#include +#include +#include +#include +#include + +#include "mlx/allocator.h" +#include "mlx/backend/cpu/encoder.h" +#include "mlx/primitives.h" + +namespace mtplx_native::ple_cpu_rows { + +namespace mx = mlx::core; + +// Defined by ple_cpu_rows.cpp. Keeping the installation boundary there makes +// both factories use the same process-lifetime stream and permit pool without +// duplicating the stream constructor in this provider translation unit. +mx::Stream installed_cpu_stream(); +std::shared_ptr installed_permits(); + +namespace { + +using mx::array; + +class PleSidecarRowsPrimitive final : public mx::Primitive { + public: + PleSidecarRowsPrimitive(mx::Stream cpu_stream, + std::shared_ptr job) + : mx::Primitive(cpu_stream), job_(std::move(job)) {} + + ~PleSidecarRowsPrimitive() override = default; + + void eval_cpu(const std::vector&, + std::vector& outputs) override { + auto& weight_output = outputs[0]; + auto& scales_output = outputs[1]; + auto& bias_output = outputs[2]; + + // The output descriptors are captured by the queued task so their MLX + // allocator buffers stay alive through all 192 bounded preads and the + // packed-to-plane split. + weight_output.set_data(mx::allocator::malloc(kWeightBytes)); + scales_output.set_data(mx::allocator::malloc(kMetadataBytes)); + bias_output.set_data(mx::allocator::malloc(kMetadataBytes)); + + auto weight = weight_output; + auto scales = scales_output; + auto bias = bias_output; + auto job = std::move(job_); + auto task = [job = std::move(job), + weight = std::move(weight), + scales = std::move(scales), + bias = std::move(bias)]() mutable { + try { + auto* weight_dst = weight.data(); + auto* scales_dst = scales.data(); + auto* bias_dst = bias.data(); + // A failed sidecar read must never expose partially populated planes. + std::memset(weight_dst, 0, kWeightBytes); + std::memset(scales_dst, 0, kMetadataBytes); + std::memset(bias_dst, 0, kMetadataBytes); + + const auto packed = job->run(); + for (std::size_t row = 0; row < kRows; ++row) { + const auto* source = packed.data() + row * 100; + std::memcpy(weight_dst + row * kWeightValuesPerRow, source, 80); + std::memcpy(scales_dst + row * kMetadataValuesPerRow, + source + 80, + 10); + std::memcpy(bias_dst + row * kMetadataValuesPerRow, + source + 90, + 10); + } + } catch (...) { + // The ordinary MLX CPU scheduler owns exception propagation. The + // job only releases its permit here; no custom error/event channel is + // introduced by the provider lane. + job->release_permit(); + throw; + } + job->release_permit(); + }; + + auto& encoder = mx::cpu::get_command_encoder(stream()); + encoder.set_output_array(weight_output); + encoder.set_output_array(scales_output); + encoder.set_output_array(bias_output); + encoder.dispatch(std::move(task)); + } + + void eval_gpu(const std::vector&, + std::vector&) override { + throw std::runtime_error( + "[mtplx_native_ple_cpu_rows] sidecar primitive is CPU-stream-only"); + } + + const char* name() const override { return "MtplxPleSidecarRows"; } + + std::vector output_shapes( + const std::vector&) override { + return {mx::Shape{64, 20}, mx::Shape{64, 5}, mx::Shape{64, 5}}; + } + + bool is_equivalent(const mx::Primitive&) const override { + // Each job has an independent snapshot, reader work, and permit. + return false; + } + + private: + std::shared_ptr job_; +}; + +} // namespace + +CpuRowsArrays make_sidecar_rows( + const std::shared_ptr& producer, + const std::array& previous, + const std::array& ids) { + if (producer == nullptr) { + throw std::invalid_argument( + "[mtplx_native_ple_cpu_rows] sidecar producer is null"); + } + SidecarJobInput input{previous, ids}; + auto job = producer->make_job(input, installed_permits()); + auto primitive = std::make_shared( + installed_cpu_stream(), std::move(job)); + auto outputs = mx::array::make_arrays( + {mx::Shape{64, 20}, mx::Shape{64, 5}, mx::Shape{64, 5}}, + {mx::uint32, mx::bfloat16, mx::bfloat16}, + primitive, + std::vector{}); + return {outputs.at(0), outputs.at(1), outputs.at(2)}; +} + +} // namespace mtplx_native::ple_cpu_rows diff --git a/native_extensions/ple_cpu_rows/sidecar_producer.cpp b/native_extensions/ple_cpu_rows/sidecar_producer.cpp new file mode 100644 index 000000000..e6dfcd59c --- /dev/null +++ b/native_extensions/ple_cpu_rows/sidecar_producer.cpp @@ -0,0 +1,160 @@ +// SPDX-License-Identifier: Apache-2.0 + +#include "sidecar_producer.h" + +#include +#include +#include +#include + +namespace mtplx_native::ple_cpu_rows { + +namespace { + +using I64 = std::int64_t; +using U64 = std::uint64_t; + +void require_sidecar_range(I64 size, + I64 offset, + std::size_t head, + U64 row_count) { + if (size <= 0) { + throw std::invalid_argument("sidecar ngram head size must be positive"); + } + if (offset < 0) { + throw std::out_of_range( + "sidecar ngram head offset must be nonnegative"); + } + + // size is positive, so this subtraction cannot underflow. Check the + // signed addition before evaluating it; this keeps the constructor proof + // defined even for a deliberately adversarial INT64_MAX fixture. + const I64 last = size - 1; + if (offset > std::numeric_limits::max() - last) { + throw std::overflow_error("sidecar ngram head range overflows int64"); + } + const I64 end = offset + last; + if (static_cast(end) > + static_cast(std::numeric_limits::max())) { + throw std::out_of_range( + "sidecar ngram head range exceeds uint32 row IDs"); + } + if (static_cast(end) >= row_count) { + throw std::out_of_range( + "sidecar ngram head range exceeds sidecar row count"); + } + (void)head; +} + +} // namespace + +std::shared_ptr SidecarProducer::install( + int descriptor, + host::SidecarLayout layout, + const std::array& multipliers, + const std::array& sizes, + const std::array& offsets, + I64 eos) { + // RawSidecarBatchReader duplicates and validates the descriptor before its + // shared owner is published. Closing the caller's fd immediately after + // this call therefore cannot invalidate a queued job. + auto reader = std::make_shared( + descriptor, layout); + validate_row_ranges(sizes, offsets, layout.row_count); + return std::shared_ptr(new SidecarProducer( + layout, multipliers, sizes, offsets, eos, std::move(reader))); +} + +SidecarProducer::SidecarProducer( + host::SidecarLayout layout, + std::array multipliers, + std::array sizes, + std::array offsets, + I64 eos, + std::shared_ptr reader) + : row_count_(layout.row_count), + plan_(multipliers, sizes, offsets, eos), + reader_(std::move(reader)) { + if (reader_ == nullptr) { + throw std::invalid_argument("sidecar producer reader is null"); + } +} + +void SidecarProducer::validate_row_ranges( + const std::array& sizes, + const std::array& offsets, + U64 row_count) { + if (row_count == 0) { + throw std::invalid_argument("sidecar row count must be positive"); + } + for (std::size_t head = 0; head < host::kNgramHeads; ++head) { + require_sidecar_range(sizes[head], offsets[head], head, row_count); + } +} + +std::shared_ptr SidecarProducer::make_job( + const SidecarJobInput& input, + const std::shared_ptr& permits) const { + // shared_from_this() is intentional: a job keeps the complete installed + // producer (plan plus duplicated reader) alive until its queued callable + // finishes, rather than retaining naked references into a caller object. + return SidecarJob::admit(shared_from_this(), input, permits); +} + +std::shared_ptr SidecarJob::admit( + std::shared_ptr producer, + SidecarJobInput input, + const std::shared_ptr& permits) { + if (producer == nullptr) { + throw std::invalid_argument("sidecar job producer is null"); + } + if (permits == nullptr) { + throw std::invalid_argument("sidecar job permit pool is null"); + } + if (!permits->try_acquire()) { + throw std::runtime_error( + "[mtplx_native_ple_cpu_rows] bounded two-request queue is full"); + } + + // As with the synthetic RequestState, hold a move-only lease across every + // allocation needed to publish the shared job. An allocation exception + // releases the permit once, while the published job owns it thereafter. + PermitLease lease(permits); + auto owned = std::unique_ptr( + new SidecarJob(std::move(producer), std::move(input), permits)); + lease.disarm(); + return std::shared_ptr(std::move(owned)); +} + +SidecarJob::~SidecarJob() { release_permit(); } + +SidecarPackedRows SidecarJob::run() const { + try { + const host::NgramRowsResult result = + producer_->plan().compute(input_.previous, input_.ids); + + // SidecarProducer::validate_row_ranges proves every result is in the + // uint32/reader domain for every possible modulo result. This conversion + // is consequently an invariant-preserving cast, not a per-window check. + SidecarRowIds row_ids{}; + for (std::size_t index = 0; index < row_ids.size(); ++index) { + row_ids[index] = static_cast(result.rows[index]); + } + // On an I/O/hash exception no output publication can follow, so release + // the slot here before propagating the ordinary CPU exception. Success + // intentionally leaves the slot held until the outer primitive copies all + // three MLX planes and calls release_permit(). + return producer_->reader().read_rows(row_ids); + } catch (...) { + release_permit(); + throw; + } +} + +void SidecarJob::release_permit() const noexcept { + if (permit_held_.exchange(false, std::memory_order_acq_rel)) { + permits_->release(); + } +} + +} // namespace mtplx_native::ple_cpu_rows diff --git a/native_extensions/ple_cpu_rows/sidecar_producer.h b/native_extensions/ple_cpu_rows/sidecar_producer.h new file mode 100644 index 000000000..989dba188 --- /dev/null +++ b/native_extensions/ple_cpu_rows/sidecar_producer.h @@ -0,0 +1,114 @@ +// SPDX-License-Identifier: Apache-2.0 +// +// MLX-free construction and job ownership for the real sidecar CPU-row +// producer. The implementation deliberately reuses the standalone +// host_provider NgramPlan and RawSidecarBatchReader; this header contains no +// model, Python, or MLX dependency. + +#pragma once + +#include +#include +#include +#include +#include + +#include "host_provider.h" +#include "request_state.h" + +namespace mtplx_native::ple_cpu_rows { + +namespace host = mtplx_native::host_provider; + +using SidecarPackedRows = host::PackedRows; +using SidecarRowIds = host::RowIds; + +// The only values that vary for an M4 sidecar job. The arrays are copied into +// SidecarJob at admission; callers may freely reuse or mutate their originals +// once make_job returns. +struct SidecarJobInput { + std::array previous{}; + std::array ids{}; +}; + +// A construction-bound, immutable sidecar route. The plan and reader are +// made once and retained by every job. In particular, no fstat/layout check, +// plan construction, or descriptor duplication occurs on a window job. +class SidecarProducer final + : public std::enable_shared_from_this { + public: + static std::shared_ptr install( + int descriptor, + host::SidecarLayout layout, + const std::array& multipliers, + const std::array& sizes, + const std::array& offsets, + std::int64_t eos); + + SidecarProducer(const SidecarProducer&) = delete; + SidecarProducer& operator=(const SidecarProducer&) = delete; + + const host::NgramPlan& plan() const noexcept { return plan_; } + const host::RawSidecarBatchReader& reader() const noexcept { + return *reader_; + } + std::uint64_t row_count() const noexcept { return row_count_; } + + // Admission is nonblocking. The returned job owns one permit until its + // run() task completes or the job is abandoned before dispatch. + std::shared_ptr make_job( + const SidecarJobInput& input, + const std::shared_ptr& permits) const; + + private: + SidecarProducer( + host::SidecarLayout layout, + std::array multipliers, + std::array sizes, + std::array offsets, + std::int64_t eos, + std::shared_ptr reader); + + static void validate_row_ranges( + const std::array& sizes, + const std::array& offsets, + std::uint64_t row_count); + + const std::uint64_t row_count_; + const host::NgramPlan plan_; + const std::shared_ptr reader_; +}; + +// One immutable M4 sidecar request. run() computes exactly 64 row IDs using +// the installed plan and delegates the fixed 192 preads to the installed +// reader; it returns packed bytes only and never publishes mutable history. +class SidecarJob final { + public: + static std::shared_ptr admit( + std::shared_ptr producer, + SidecarJobInput input, + const std::shared_ptr& permits); + + ~SidecarJob(); + + SidecarJob(const SidecarJob&) = delete; + SidecarJob& operator=(const SidecarJob&) = delete; + + SidecarPackedRows run() const; + void release_permit() const noexcept; + + private: + SidecarJob(std::shared_ptr producer, + SidecarJobInput input, + std::shared_ptr permits) noexcept + : producer_(std::move(producer)), + input_(std::move(input)), + permits_(std::move(permits)) {} + + const std::shared_ptr producer_; + const SidecarJobInput input_; + const std::shared_ptr permits_; + mutable std::atomic_bool permit_held_{true}; +}; + +} // namespace mtplx_native::ple_cpu_rows diff --git a/scripts/bundle_native_runtime_wheel.py b/scripts/bundle_native_runtime_wheel.py index 3cdd9eb27..7eddf2453 100644 --- a/scripts/bundle_native_runtime_wheel.py +++ b/scripts/bundle_native_runtime_wheel.py @@ -24,10 +24,11 @@ # The native extensions the runtime wheel may carry, each in its own platform # wheel. mtplx_qsa_kernels is the metallib-bearing QSA lane (ea2560a2); # mtplx_native_qsa is the metallib-bearing split-K QSA sparse-GQA decode -# extension the MTPLX_QSA_SPARSE_DECODE lane needs. Every Mach-O member -# (.so/.dylib) of each is -# Developer-ID + hardened-runtime + secure-timestamp signed before packaging, -# so notarization does not reject an ad-hoc-signed member found inside the zip. +# extension the MTPLX_QSA_SPARSE_DECODE lane needs; mtplx_native_ple_cpu_rows is +# the CPU-stream PLE row extension the cached async PLE lane (PR #475) needs. +# Every Mach-O member (.so/.dylib) of each is Developer-ID + hardened-runtime + +# secure-timestamp signed before packaging, so notarization does not reject an +# ad-hoc-signed member found inside the zip. _KNOWN_NATIVE = { "mtplx_qsa_kernels": { "required_files": ("NOTICE", "LICENSE.txt", "MLX_LICENSE.txt"), @@ -37,6 +38,10 @@ "required_files": (), "require_metallib": True, }, + "mtplx_native_ple_cpu_rows": { + "required_files": (), + "require_metallib": False, + }, } _REQUIRED_NATIVE = "mtplx_qsa_kernels" diff --git a/scripts/fable/setup_over100_venv.sh b/scripts/fable/setup_over100_venv.sh new file mode 100755 index 000000000..7c1b258cd --- /dev/null +++ b/scripts/fable/setup_over100_venv.sh @@ -0,0 +1,76 @@ +#!/bin/sh +# Reproduce the .venv + the ple_cpu_rows native extension for the Qwen3.8 +# Flash-Next aux-lane substrate on upstream/main. CPU only; no GPU. +# +# NOTE (upstream/main rebase, 2026-09-07): the original over100 script built two +# extensions (qsa_sparse_gqa + ple_cpu_rows) and verified them through PR 391's +# mtplx/full_stack_env.py-based stack. On upstream main the QSA native path is +# loaded elsewhere (mtplx/kernels/qsa_prefill_direct.py, extension +# native_extensions/qsa_kernels) and there is no full_stack_env, so this adapted +# script builds only the self-contained ple_cpu_rows extension that the +# ple_cached_aux lane needs and verifies it by importing the built package +# directly. The two decode lanes ARE armed on upstream main: the server +# auto-arms MTPLX_QWEN4_PLE_CACHED_AUX / MTPLX_QSA_POOLED_ROWSEL for a served +# fixed-M4 Flash-Next pack (mtplx/server/openai.py), and the ple_cached_aux lane +# declines to stock with a printed reason when this extension is not built. +# +# Pins mlx 0.32.2 (== production) and an editable mtplx == THIS worktree. +# The mlx 0.32.2 wheel is built with nanobind internals v21. uv.lock pins +# nanobind 2.12.0 (v19), which the native CMake ABI guards reject at configure +# time, so step 2 upgrades nanobind to 2.15.0 (v21) for the native build only. +# mtplx never imports nanobind at runtime, so this does not perturb serving. +# +# The venv python is pinned to 3.12 because mlx 0.32.2 ships cp312 wheels; the +# box's base interpreter is 3.14 with mlx 0.31.2, the wrong ABI for this build. +# The worktree root is derived from this script's own location. +set -eu + +WT="$(cd "$(dirname "$0")/../.." && pwd)" +cd "$WT" + +# 1. venv (python 3.12) + locked deps + editable mtplx (base only, --frozen). +nice -n 19 uv sync --frozen --no-dev --python 3.12 + +# 2. nanobind matching mlx.core's build (v21) for the native extension. +nice -n 19 uv pip install --python .venv/bin/python 'nanobind==2.15.0' + +PY="$WT/.venv/bin/python" +# Derive the nanobind cmake root from the venv's actual python version. +NB="$("$PY" -c 'import nanobind,pathlib;print(pathlib.Path(nanobind.__file__).parent)')" + +# 3. the ple_cpu_rows native extension, built in-place into its package dir. +# (CPU-stream PLE rows -- the lazily-imported dep of the ple_cached_aux lane.) +SRC="$WT/native_extensions/ple_cpu_rows" +PKG="mtplx_native_ple_cpu_rows" +rm -rf "$SRC/build" +nice -n 19 cmake -S "$SRC" -B "$SRC/build" \ + -DCMAKE_LIBRARY_OUTPUT_DIRECTORY="$SRC/$PKG/" \ + -DCMAKE_BUILD_TYPE=Release -DBUILD_SHARED_LIBS=ON \ + -DPython_EXECUTABLE="$PY" \ + -DMTPLX_NANOBIND_DIR="$NB" +nice -n 19 cmake --build "$SRC/build" -j 8 + +# 4. verify: mlx pin, editable mtplx == this worktree, ple_cpu_rows importable. +"$PY" - "$WT" "$SRC" <<'PYEOF' +import sys +from pathlib import Path + +import mlx.core as mx +import mtplx + +wt = Path(sys.argv[1]).resolve() +src = Path(sys.argv[2]).resolve() +assert mx.__version__ == "0.32.2", mx.__version__ +assert Path(mtplx.__file__).resolve().parents[1] == wt, mtplx.__file__ + +sys.path.insert(0, str(src)) +import mtplx_native_ple_cpu_rows as ext +for sym in ("CachedSidecarProducer", "compute_cached_row_ids", + "install_cached_sidecar_provider", "make_cpu_rows"): + assert hasattr(ext, sym), sym +print( + "over100 venv OK: mlx", mx.__version__, + "| mtplx", mtplx.__file__, + "| ple_cpu_rows native available", +) +PYEOF diff --git a/tests/test_bundle_native_runtime_wheel.py b/tests/test_bundle_native_runtime_wheel.py index 397c0caaf..a39509536 100644 --- a/tests/test_bundle_native_runtime_wheel.py +++ b/tests/test_bundle_native_runtime_wheel.py @@ -32,6 +32,8 @@ EXT = "mtplx_qsa_kernels/_ext.cpython-314-darwin.so" DYLIB = "mtplx_qsa_kernels/libmtplx_qsa_kernel_ops.dylib" METALLIB = "mtplx_qsa_kernels/kernels.metallib" +PLE_EXT = "mtplx_native_ple_cpu_rows/_ext.cpython-314-darwin.so" +PLE_DYLIB = "mtplx_native_ple_cpu_rows/libmtplx_native_ple_cpu_rows.dylib" # PR #391 remainder port: the split-K QSA sparse-GQA decode extension is a # second native wheel (mtplx_native_qsa) that also carries a metallib. @@ -82,6 +84,25 @@ def _inputs(tmp_path: Path) -> tuple[Path, Path]: return pure, native +def _ple_input(tmp_path: Path) -> Path: + return _write_wheel( + tmp_path / "mtplx_native_ple_cpu_rows-9.9.9-cp314-cp314-macosx_15_0_arm64.whl", + { + "mtplx_native_ple_cpu_rows/__init__.py": b"", + PLE_EXT: b"PLE-MACHO-EXT", + PLE_DYLIB: b"PLE-MACHO-DYLIB", + "mtplx_native_ple_cpu_rows-9.9.9.dist-info/METADATA": ( + b"Metadata-Version: 2.1\nName: mtplx-native-ple-cpu-rows\nVersion: 9.9.9\n" + b"Requires-Dist: mlx==0.32.2\n" + ), + "mtplx_native_ple_cpu_rows-9.9.9.dist-info/WHEEL": ( + b"Wheel-Version: 1.0\nGenerator: test\nRoot-Is-Purelib: false\n" + b"Tag: cp314-cp314-macosx_15_0_arm64\n" + ), + }, + ) + + def _run_bundler(monkeypatch, pure: Path, native: Path, out: Path, *extra: str) -> Path: monkeypatch.setattr(sys, "argv", ["bundle", str(pure), str(native), "--out", str(out), *extra]) bundler.main() @@ -250,3 +271,40 @@ def test_qsa_native_missing_its_metallib_is_rejected(tmp_path, monkeypatch) -> N ) with pytest.raises(SystemExit): _run_bundler_multi(monkeypatch, pure, [native, qsa], tmp_path / "out") + + +def test_ple_cpu_rows_extension_is_bundled_and_signed_with_the_qsa_kernels(tmp_path, monkeypatch) -> None: + # PR #475's cached async PLE lane ships a second native extension, + # mtplx_native_ple_cpu_rows. It must be bundled and Developer-ID signed + # alongside the QSA kernels, or notarization rejects its ad-hoc-signed + # Mach-O the way it rejected the QSA kernels before ea2560a2. + pure, native = _inputs(tmp_path) + ple = _ple_input(tmp_path) + calls: list[list[str]] = [] + monkeypatch.setattr(bundler.subprocess, "run", _fake_codesign(calls)) + bundled = _run_bundler_multi( + monkeypatch, pure, [native, ple], tmp_path / "out", + "--codesign-identity", "Developer ID Application: Test", + ) + with zipfile.ZipFile(bundled) as archive: + assert archive.read(EXT) == b"SIGNED:MACHO-EXT" + assert archive.read(PLE_EXT) == b"SIGNED:PLE-MACHO-EXT" + assert archive.read(PLE_DYLIB) == b"SIGNED:PLE-MACHO-DYLIB" + top_level = archive.read("mtplx-9.9.9.dist-info/top_level.txt").decode() + record = archive.read("mtplx-9.9.9.dist-info/RECORD").decode().splitlines() + assert "mtplx_qsa_kernels" in top_level.split() + assert "mtplx_native_ple_cpu_rows" in top_level.split() + hashes = {line.split(",")[0]: line.split(",")[1] for line in record if line} + assert hashes[PLE_EXT] == _record_hash(b"SIGNED:PLE-MACHO-EXT") + with WheelFile(bundled) as reopened: # the rewritten RECORD verifies + assert reopened.read(PLE_EXT) == b"SIGNED:PLE-MACHO-EXT" + signed = [Path(call[-1]).name for call in calls if "--sign" in call] + assert "_ext.cpython-314-darwin.so" in signed # QSA + assert "libmtplx_native_ple_cpu_rows.dylib" in signed # PLE + + +def test_ple_cpu_rows_alone_is_rejected_without_the_required_qsa_kernels(tmp_path, monkeypatch) -> None: + pure, _native = _inputs(tmp_path) + ple = _ple_input(tmp_path) + with pytest.raises(SystemExit): + _run_bundler_multi(monkeypatch, pure, [ple], tmp_path / "out") diff --git a/tests/test_pr391_cached_sidecar_primitive_cpu.py b/tests/test_pr391_cached_sidecar_primitive_cpu.py new file mode 100644 index 000000000..d4ab97e50 --- /dev/null +++ b/tests/test_pr391_cached_sidecar_primitive_cpu.py @@ -0,0 +1,158 @@ +"""CPU-only API and source checks for the cached sidecar MLX glue. + +The primitive and nanobind surface are intentionally not compiled here. This +gate checks the exact source contract before the isolated operator build: +three MLX-owned planes, construction-bound CPU stream/permit reuse, copied +fixed handoff inputs, and capsule-owned immutable NumPy outputs. +""" + +from __future__ import annotations + +import re +from pathlib import Path + + +ROOT = Path(__file__).resolve().parents[1] +EXT = ROOT / "native_extensions" / "ple_cpu_rows" +PRIMITIVE = EXT / "cached_sidecar_primitive.cpp" +HEADER = EXT / "cached_sidecar_primitive.h" +BINDINGS = EXT / "bindings.cpp" +CMAKE = EXT / "CMakeLists.txt" +PACKAGE = EXT / "mtplx_native_ple_cpu_rows" / "__init__.py" + + +def _body(source: str, marker: str) -> str: + start = source.index(marker) + opening = source.index("{", start) + depth = 0 + for position in range(opening, len(source)): + char = source[position] + if char == "{": + depth += 1 + elif char == "}": + depth -= 1 + if depth == 0: + return source[opening + 1 : position] + raise AssertionError(f"unterminated body for {marker!r}") + + +def test_cached_primitive_files_are_staged_and_added_to_cmake(): + assert HEADER.is_file() + assert PRIMITIVE.is_file() + cmake = CMAKE.read_text(encoding="utf-8") + assert "cached_sidecar_primitive.cpp" in cmake + assert "cached_sidecar_producer.cpp" in cmake + assert "host_provider.cpp" in cmake + + +def test_cached_primitive_reuses_cpu_encoder_and_exact_three_plane_contract(): + source = PRIMITIVE.read_text(encoding="utf-8") + body = _body(source, "void eval_cpu(") + assert '#include "mlx/backend/cpu/encoder.h"' in source + assert body.count("set_data(") == 3 + assert body.count("mx::allocator::malloc") == 3 + assert "mx::cpu::get_command_encoder(stream())" in body + assert "encoder.dispatch" in body + assert "mx::Shape{64, 20}" in source + assert source.count("mx::Shape{64, 5}") >= 2 + assert "mx::uint32" in source + assert source.count("mx::bfloat16") >= 2 + assert "copy_packed_to_planes" in body + + +def test_cached_primitive_has_no_hidden_event_or_hotpath_sync(): + source = PRIMITIVE.read_text(encoding="utf-8") + forbidden = ( + r"\bEvent\b", + r"\bevent\s*\(", + r"attach_event", + r"wait_event", + r"signal_event", + r"synchronize\s*\(", + r"mx::eval", + r"mlx/backend/metal", + r"objc", + r"MTL", + r"std::thread", + r"condition_variable", + ) + assert not [pattern for pattern in forbidden if re.search(pattern, source)] + assert "eval_gpu" in source + assert "CPU-stream-only" in source + + +def test_cached_primitive_keeps_job_and_output_descriptors_alive_until_copy(): + source = PRIMITIVE.read_text(encoding="utf-8") + body = _body(source, "void eval_cpu(") + assert re.search(r"\[job = std::move\(job\).*weight.*scales.*bias", body, re.DOTALL) + assert "job->run()" in body + assert "job->release_permit()" in body + assert "std::memset" in body + assert "throw" in body + assert "array::make_arrays" in source + + +def test_cached_binding_exposes_exact_additive_api_and_ndarray_contract(): + source = BINDINGS.read_text(encoding="utf-8") + for name in ( + "CachedSidecarProducer", + "install_cached_sidecar_provider", + "compute_cached_row_ids", + "make_cached_sidecar_rows", + "drain_cached_completions", + ): + assert f'"{name}"' in source + assert '#include ' in source + assert "nb::ndarray" in source + assert "nb::ndim<2>" in source + assert "nb::ndim<1>" in source + assert "nb::capsule" in source + assert "delete_owned_packed" in source + assert "std::memcpy(handoff.source.data(), source.data()" in source + assert "hits.data()" in source + assert "misses.data()" in source + assert "submission.ticket.has_value()" in source + assert "nb::none()" in source + assert "class_" in source + assert "class_ 64" in helper + assert "misses.shape(0) > 64" in helper + assert "hit_count = static_cast" in helper + assert "miss_count = static_cast" in helper + assert "const std::uint8_t" in source + assert "const std::uint32_t" in source + assert "owner->data()" in source + assert "capsule" in source + + +def test_cached_package_preserves_old_exports_and_adds_new_exports(): + package = PACKAGE.read_text(encoding="utf-8") + for name in ( + "SidecarProducer", + "install_sidecar_provider", + "make_sidecar_rows", + "make_cpu_rows", + "CachedSidecarProducer", + "install_cached_sidecar_provider", + "compute_cached_row_ids", + "make_cached_sidecar_rows", + "drain_cached_completions", + ): + assert name in package + + +def test_cached_primitive_header_keeps_ticket_and_completion_surface_mlxfree_boundary(): + header = HEADER.read_text(encoding="utf-8") + assert "CachedRowsSubmission" in header + assert "std::optional ticket" in header + assert "CachedRowHandoff" in header + assert "drain_cached_completions" in header + assert "CachedRowsArrays" in header diff --git a/tests/test_pr391_cached_sidecar_producer_cpu.py b/tests/test_pr391_cached_sidecar_producer_cpu.py new file mode 100644 index 000000000..3b7ab20cb --- /dev/null +++ b/tests/test_pr391_cached_sidecar_producer_cpu.py @@ -0,0 +1,655 @@ +"""MLX-free TDD coverage for the cache-aware sidecar producer foundation. + +This driver links only the authoritative host reader, the existing immutable +plan producer, and the new cache-aware CPU producer. It never imports MLX, +builds an extension, starts a model, or touches the GPU. +""" + +from __future__ import annotations + +from functools import lru_cache +import hashlib +from pathlib import Path +import shutil +import subprocess +import tempfile +import textwrap + + +ROOT = Path(__file__).resolve().parents[1] +CPU = ROOT / "native_extensions" / "ple_cpu_rows" +LATE = CPU # host_provider.{h,cpp} are co-located in ple_cpu_rows in this PR + + +_DRIVER = r""" +#include "cached_sidecar_producer.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace pp = mtplx_native::ple_cpu_rows; +namespace hp = mtplx_native::host_provider; + +namespace { + +constexpr std::uint64_t kRowCount = 2000; + +struct Config { + std::array multipliers{ + INT64_MIN, -7, 11}; + std::array sizes{}; + std::array offsets{}; + std::int64_t eos = 99; + + Config() { + for (std::size_t index = 0; index < sizes.size(); ++index) { + sizes[index] = 500 + static_cast(index * 3); + offsets[index] = 100 + static_cast(index * 10); + } + } +}; + +hp::SidecarLayout layout() { + const std::uint64_t weights_offset = 13; + const std::uint64_t scales_offset = weights_offset + kRowCount * 80 + 7; + const std::uint64_t biases_offset = scales_offset + kRowCount * 10 + 11; + return hp::SidecarLayout{ + kRowCount, + {weights_offset, kRowCount * 80, 80}, + {scales_offset, kRowCount * 10, 10}, + {biases_offset, kRowCount * 10, 10}, + }; +} + +hp::PackedRow payload(std::uint32_t row) { + hp::PackedRow result{}; + for (std::size_t index = 0; index < result.size(); ++index) { + result[index] = static_cast( + (static_cast(row) * 17 + index * 3 + 5) & 0xff); + } + return result; +} + +struct ReaderState { + mutable std::mutex mutex; + std::condition_variable condition; + std::unordered_map calls; + std::uint32_t fail_row = UINT32_MAX; + std::uint32_t block_row = UINT32_MAX; + bool block_started = false; + bool unblock = false; +}; + +class CountingReader final : public pp::CachedRowReader { + public: + explicit CountingReader(std::shared_ptr state) + : state_(std::move(state)) {} + + hp::PackedRow read_one(std::uint32_t row_id) const override { + std::unique_lock lock(state_->mutex); + ++state_->calls[row_id]; + if (row_id == state_->fail_row) { + throw std::runtime_error("short read"); + } + if (row_id == state_->block_row) { + state_->block_started = true; + state_->condition.notify_all(); + state_->condition.wait(lock, [this] { return state_->unblock; }); + } + return payload(row_id); + } + + private: + const std::shared_ptr state_; +}; + +void create_fixture(const std::string& path) { + const auto sidecar = layout(); + const std::uint64_t end = sidecar.biases.offset + sidecar.biases.length; + const int fd = ::open(path.c_str(), O_CREAT | O_TRUNC | O_RDWR, 0600); + if (fd < 0) throw std::runtime_error("fixture open failed"); + if (::ftruncate(fd, static_cast(end)) != 0) { + ::close(fd); + throw std::runtime_error("fixture truncate failed"); + } + ::close(fd); +} + +std::shared_ptr install( + const std::string& path, + std::shared_ptr state, + std::size_t workers = 8) { + const int fd = ::open(path.c_str(), O_RDONLY); + if (fd < 0) throw std::runtime_error("fixture reopen failed"); + try { + auto reader = std::make_shared(std::move(state)); + auto producer = pp::CachedSidecarProducer::install_for_test( + fd, layout(), Config{}.multipliers, Config{}.sizes, Config{}.offsets, + Config{}.eos, std::move(reader), workers); + ::close(fd); + return producer; + } catch (...) { + ::close(fd); + throw; + } +} + +pp::CachedRowHandoff handoff( + std::initializer_list sources, + std::uint8_t hit_count, + std::initializer_list miss_ids) { + pp::CachedRowHandoff result{}; + result.hit_count = hit_count; + std::size_t index = 0; + for (const std::uint8_t source : sources) { + if (index >= result.source.size()) throw std::runtime_error("too many sources"); + result.source[index++] = source; + } + while (index < result.source.size()) { + result.source[index++] = hit_count == 0 ? 0 : 0x80; + } + std::size_t miss = 0; + for (const std::uint32_t row : miss_ids) { + result.miss_ids[miss++] = row; + } + result.miss_count = static_cast(miss); + for (std::size_t row = 0; row < result.hit_count; ++row) { + result.hit_packed[row] = payload(static_cast(700 + row)); + } + return result; +} + +void assert_output(const hp::PackedRows& output, + const pp::CachedRowHandoff& input) { + for (std::size_t row = 0; row < hp::kRowsPerWindow; ++row) { + const std::uint8_t source = input.source[row]; + const bool hit = (source & pp::kCachedHitBit) != 0; + const std::size_t index = source & pp::kCachedSourceMask; + const hp::PackedRow expected = + hit ? input.hit_packed[index] : payload(input.miss_ids[index]); + for (std::size_t byte = 0; byte < hp::kPackedBytesPerRow; ++byte) { + if (output[row * hp::kPackedBytesPerRow + byte] != expected[byte]) { + throw std::runtime_error("scatter output mismatch"); + } + } + } +} + +int run_all_hits(const std::string& path) { + auto state = std::make_shared(); + auto producer = install(path, state); + auto pool = std::make_shared(); + auto input = handoff({0x80}, 1, {}); + auto job = producer->make_job(input, pool); + assert_output(job->run(), input); + job->release_permit(); + if (pool->outstanding() != 0 || !state->calls.empty()) { + throw std::runtime_error("all-hit job performed a read or leaked permit"); + } + if (!producer->drain_completed().empty()) { + throw std::runtime_error("all-hit job emitted a completion"); + } + std::cout << "all-hits-ok\n"; + return 0; +} + +int run_mixed(const std::string& path, bool duplicate) { + auto state = std::make_shared(); + auto producer = install(path, state); + auto pool = std::make_shared(); + auto input = duplicate + ? handoff({0x80, 0x00, 0x01, 0x01, 0x02, 0x80}, 1, {11, 12, 13}) + : handoff({0x80, 0x00, 0x01, 0x80, 0x02}, 1, {11, 12, 13}); + auto job = producer->make_job(input, pool); + assert_output(job->run(), input); + job->release_permit(); + if (pool->outstanding() != 1) { + throw std::runtime_error("miss completion released permit early"); + } + const auto completions = producer->drain_completed(); + if (completions.size() != 1 || completions[0].count != 3) { + throw std::runtime_error("miss completion shape mismatch"); + } + if (pool->outstanding() != 0) { + throw std::runtime_error("completion drain did not release permit"); + } + if (state->calls.size() != 3) { + throw std::runtime_error("miss dedup read count mismatch"); + } + for (const auto& item : state->calls) { + if (item.second != 1) throw std::runtime_error("miss read repeated"); + } + std::cout << (duplicate ? "duplicate-ok\n" : "mixed-ok\n"); + return 0; +} + +int run_error(const std::string& path) { + auto state = std::make_shared(); + state->fail_row = 12; + auto producer = install(path, state); + auto pool = std::make_shared(); + auto input = handoff({0x00, 0x01, 0x02}, 0, {11, 12, 13}); + auto job = producer->make_job(input, pool); + try { + (void)job->run(); + throw std::runtime_error("read failure unexpectedly succeeded"); + } catch (const std::runtime_error& error) { + if (std::string(error.what()) != "short read") throw; + } + if (pool->outstanding() != 0 || !producer->drain_completed().empty()) { + throw std::runtime_error("error did not release or emitted completion"); + } + std::cout << "error-ok\n"; + return 0; +} + +int run_abandon(const std::string& path) { + auto state = std::make_shared(); + auto producer = install(path, state); + auto pool = std::make_shared(); + { + auto job = producer->make_job(handoff({0x00}, 0, {11}), pool); + job->abandon(); + } + if (pool->outstanding() != 0 || !producer->drain_completed().empty()) { + throw std::runtime_error("abandon leaked permit or completion"); + } + std::cout << "abandon-ok\n"; + return 0; +} + +int run_concurrent(const std::string& path) { + auto state = std::make_shared(); + auto producer = install(path, state, 8); + auto pool = std::make_shared(); + auto first = producer->make_job(handoff({0x00}, 0, {11}), pool); + auto second = producer->make_job(handoff({0x00}, 0, {11}), pool); + (void)first->run(); + (void)second->run(); + first->release_permit(); + second->release_permit(); + auto completions = producer->drain_completed(); + if (completions.size() != 2 || state->calls[11] != 2) { + throw std::runtime_error("cross-job dedup was silently claimed"); + } + if (pool->outstanding() != 0) throw std::runtime_error("permit leaked"); + std::cout << "concurrent-ok\n"; + return 0; +} + +int run_pending_bound(const std::string& path) { + auto state = std::make_shared(); + auto producer = install(path, state, 8); + auto pool = std::make_shared(); + auto first = producer->make_job(handoff({0x00}, 0, {11}), pool); + auto second = producer->make_job(handoff({0x00}, 0, {12}), pool); + (void)first->run(); + (void)second->run(); + first->release_permit(); + second->release_permit(); + if (pool->outstanding() != 2) { + throw std::runtime_error("completed jobs released permits before drain"); + } + bool rejected = false; + try { + (void)producer->make_job(handoff({0x00}, 0, {13}), pool); + } catch (const std::runtime_error&) { + rejected = true; + } + if (!rejected) { + throw std::runtime_error("third job bypassed pending completion bound"); + } + const auto completions = producer->drain_completed(); + if (completions.size() != 2 || pool->outstanding() != 0) { + throw std::runtime_error("pending completion drain did not release both slots"); + } + std::cout << "pending-bound-ok\n"; + return 0; +} + +int run_worker_bounds(const std::string& path) { + for (const std::size_t workers : {std::size_t{0}, pp::kCachedMaxIoWorkers + 1}) { + auto state = std::make_shared(); + bool rejected = false; + try { + (void)install(path, state, workers); + } catch (const std::invalid_argument&) { + rejected = true; + } + if (!rejected) { + throw std::runtime_error("out-of-bound worker count was accepted"); + } + } + std::cout << "worker-bounds-ok\n"; + return 0; +} + +int run_active_abandon(const std::string& path) { + auto state = std::make_shared(); + state->block_row = 11; + auto producer = install(path, state, 1); + auto pool = std::make_shared(); + auto first = producer->make_job(handoff({0x00}, 0, {11}), pool); + std::exception_ptr run_error; + std::thread runner([&] { + try { + (void)first->run(); + } catch (...) { + run_error = std::current_exception(); + } + }); + { + std::unique_lock lock(state->mutex); + state->condition.wait(lock, [&] { return state->block_started; }); + } + first->abandon(); + auto second = producer->make_job(handoff({0x00}, 0, {12}), pool); + bool rejected = false; + try { + (void)producer->make_job(handoff({0x00}, 0, {13}), pool); + } catch (const std::runtime_error&) { + rejected = true; + } + if (!rejected || pool->outstanding() != 2) { + throw std::runtime_error("active abandon bypassed two-job bound"); + } + { + std::lock_guard lock(state->mutex); + state->unblock = true; + } + state->condition.notify_all(); + runner.join(); + if (run_error) std::rethrow_exception(run_error); + if (pool->outstanding() != 1 || !producer->drain_completed().empty()) { + throw std::runtime_error("active abandon did not release after worker drain"); + } + second->abandon(); + if (pool->outstanding() != 0) { + throw std::runtime_error("active abandon test leaked second permit"); + } + std::cout << "active-abandon-ok\n"; + return 0; +} + +int run_completion_survives_drop(const std::string& path) { + auto state = std::make_shared(); + auto producer = install(path, state); + auto pool = std::make_shared(); + { + auto job = producer->make_job(handoff({0x00}, 0, {11}), pool); + (void)job->run(); + job->release_permit(); + } + if (pool->outstanding() != 1) { + throw std::runtime_error("job destruction released undrained completion"); + } + const auto completions = producer->drain_completed(); + if (completions.size() != 1 || completions[0].count != 1 || + pool->outstanding() != 0) { + throw std::runtime_error("dropped job completion was not retained to drain"); + } + std::cout << "completion-drop-ok\n"; + return 0; +} + +int run_provider_drop_releases_completion(const std::string& path) { + auto pool = std::make_shared(); + { + auto producer = install(path, std::make_shared()); + auto job = producer->make_job(handoff({0x00}, 0, {11}), pool); + (void)job->run(); + job->release_permit(); + job.reset(); + if (pool->outstanding() != 1) { + throw std::runtime_error("completion permit was released before teardown"); + } + producer.reset(); + } + if (pool->outstanding() != 0) { + throw std::runtime_error("provider teardown leaked completion permit"); + } + std::cout << "provider-drop-ok\n"; + return 0; +} + +int run_reserved_source(const std::string& path) { + auto state = std::make_shared(); + auto producer = install(path, state); + auto pool = std::make_shared(); + for (const auto& invalid : { + handoff({0x40}, 0, {11}), handoff({0xc0}, 1, {})}) { + bool rejected = false; + try { + (void)producer->make_job(invalid, pool); + } catch (const std::invalid_argument&) { + rejected = true; + } + if (!rejected) throw std::runtime_error("reserved source bit was accepted"); + } + if (pool->outstanding() != 0 || !state->calls.empty()) { + throw std::runtime_error("reserved source reached the reader"); + } + std::cout << "reserved-source-ok\n"; + return 0; +} + +int run_invalid(const std::string& path) { + auto state = std::make_shared(); + auto producer = install(path, state); + auto pool = std::make_shared(); + auto invalid = handoff({0x80 | 1}, 1, {}); + bool rejected = false; + try { + (void)producer->make_job(invalid, pool); + } catch (const std::invalid_argument&) { + rejected = true; + } + if (!rejected || pool->outstanding() != 0) { + throw std::runtime_error("invalid typed handoff was accepted"); + } + std::cout << "invalid-ok\n"; + return 0; +} + +int run_host_subset(const std::string& path) { + const int fd = ::open(path.c_str(), O_RDONLY); + if (fd < 0) throw std::runtime_error("host reopen failed"); + hp::RawSidecarBatchReader reader(fd, layout()); + ::close(fd); + const auto one = reader.read_one(7); + if (one != hp::PackedRow{}) throw std::runtime_error("read_one mismatch"); + const auto rows = reader.read_subset({7, 3, 7}); + if (rows.size() != 3 || rows[0] != hp::PackedRow{} || + rows[1] != hp::PackedRow{} || rows[2] != hp::PackedRow{}) { + throw std::runtime_error("read_subset order mismatch"); + } + bool rejected = false; + try { + reader.read_subset(std::vector(65, 1)); + } catch (const std::invalid_argument&) { + rejected = true; + } + if (!rejected) throw std::runtime_error("oversized subset accepted"); + std::cout << "host-subset-ok\n"; + return 0; +} + +} // namespace + +int main(int argc, char** argv) { + try { + if (argc < 3) throw std::invalid_argument("expected mode and fixture path"); + const std::string mode = argv[1]; + create_fixture(argv[2]); + if (mode == "all_hits") return run_all_hits(argv[2]); + if (mode == "mixed") return run_mixed(argv[2], false); + if (mode == "duplicate") return run_mixed(argv[2], true); + if (mode == "error") return run_error(argv[2]); + if (mode == "abandon") return run_abandon(argv[2]); + if (mode == "concurrent") return run_concurrent(argv[2]); + if (mode == "pending_bound") return run_pending_bound(argv[2]); + if (mode == "worker_bounds") return run_worker_bounds(argv[2]); + if (mode == "active_abandon") return run_active_abandon(argv[2]); + if (mode == "completion_drop") return run_completion_survives_drop(argv[2]); + if (mode == "provider_drop") return run_provider_drop_releases_completion(argv[2]); + if (mode == "reserved_source") return run_reserved_source(argv[2]); + if (mode == "invalid") return run_invalid(argv[2]); + if (mode == "host_subset") return run_host_subset(argv[2]); + throw std::invalid_argument("unknown mode"); + } catch (const std::exception& error) { + std::cerr << error.what() << '\n'; + return 2; + } +} +""" + + +@lru_cache(maxsize=1) +def _driver() -> tuple[tempfile.TemporaryDirectory[str], Path, Path]: + compiler = shutil.which("clang++") or shutil.which("c++") + if compiler is None: + raise AssertionError("a C++17 compiler is required") + temporary = tempfile.TemporaryDirectory(prefix="pr391-cached-sidecar-") + root = Path(temporary.name) + source = root / "driver.cpp" + executable = root / "cached_sidecar_cpu" + fixture = root / "sidecar.bin" + source.write_text(textwrap.dedent(_DRIVER), encoding="utf-8") + command = [ + compiler, + "-std=c++17", + "-O2", + "-Wall", + "-Wextra", + "-Wconversion", + "-Werror", + "-pthread", + "-I", + str(CPU), + "-I", + str(LATE), + str(source), + str(CPU / "cached_sidecar_producer.cpp"), + str(CPU / "sidecar_producer.cpp"), + str(LATE / "host_provider.cpp"), + "-o", + str(executable), + ] + result = subprocess.run(command, capture_output=True, text=True) + if result.returncode != 0: + raise AssertionError( + "cached sidecar CPU driver compile failed:\n" + f"{result.stdout}\n{result.stderr}" + ) + fixture.write_bytes(b"\0") + result = subprocess.run( + [str(executable), "host_subset", str(fixture)], + capture_output=True, + text=True, + ) + if result.returncode != 0: + raise AssertionError( + "cached sidecar fixture setup failed:\n" + f"{result.stdout}\n{result.stderr}" + ) + return temporary, executable, fixture + + +def _run(mode: str) -> str: + _temporary, executable, fixture = _driver() + result = subprocess.run( + [str(executable), mode, str(fixture)], + capture_output=True, + text=True, + ) + assert result.returncode == 0, result.stdout + result.stderr + return result.stdout.strip() + + +def test_host_reader_subset_preserves_existing_reader_and_bounds(): + assert _run("host_subset") == "host-subset-ok" + + +def test_cached_all_hits_perform_no_reads(): + assert _run("all_hits") == "all-hits-ok" + + +def test_cached_mixed_reads_only_unique_misses_and_releases_on_drain(): + assert _run("mixed") == "mixed-ok" + + +def test_cached_duplicate_positions_read_each_unique_miss_once(): + assert _run("duplicate") == "duplicate-ok" + + +def test_cached_read_failure_has_no_partial_completion_and_releases(): + assert _run("error") == "error-ok" + + +def test_cached_abandon_releases_admission_without_completion(): + assert _run("abandon") == "abandon-ok" + + +def test_cached_two_jobs_are_bounded_without_claiming_cross_job_dedup(): + assert _run("concurrent") == "concurrent-ok" + + +def test_cached_completed_jobs_hold_both_slots_until_owner_drains(): + assert _run("pending_bound") == "pending-bound-ok" + + +def test_cached_io_pool_worker_count_is_construction_bounded(): + assert _run("worker_bounds") == "worker-bounds-ok" + + +def test_cached_active_abandon_holds_slot_until_worker_drain(): + assert _run("active_abandon") == "active-abandon-ok" + + +def test_cached_dropped_job_retains_published_completion_until_drain(): + assert _run("completion_drop") == "completion-drop-ok" + + +def test_cached_provider_teardown_releases_undrained_completion_permit(): + assert _run("provider_drop") == "provider-drop-ok" + + +def test_cached_handoff_rejects_reserved_source_bit(): + assert _run("reserved_source") == "reserved-source-ok" + + +def test_cached_typed_handoff_rejects_invalid_source_indices_before_read(): + assert _run("invalid") == "invalid-ok" + + +def test_cached_sources_and_receipt_contract_are_mlxfree(): + producer = (CPU / "cached_sidecar_producer.cpp").read_text(encoding="utf-8") + header = (CPU / "cached_sidecar_producer.h").read_text(encoding="utf-8") + assert "mlx/" not in producer + assert "mlx/" not in header + assert "kCachedHitBit" in header + assert "drain_completed" in header + assert "read_one" in producer + assert "read_subset" not in producer + assert "std::thread" in producer + assert "std::condition_variable" in producer + digest = hashlib.sha256( + (CPU / "cached_sidecar_producer.h").read_bytes() + + (CPU / "cached_sidecar_producer.cpp").read_bytes() + + (LATE / "host_provider.h").read_bytes() + + (LATE / "host_provider.cpp").read_bytes() + ).hexdigest() + assert len(digest) == 64 diff --git a/tests/test_pr391_fixed_m4_pool_install_cpu.py b/tests/test_pr391_fixed_m4_pool_install_cpu.py new file mode 100644 index 000000000..9f915f307 --- /dev/null +++ b/tests/test_pr391_fixed_m4_pool_install_cpu.py @@ -0,0 +1,468 @@ +"""CPU contract for the construction-bound fixed-M4 pool installer. + +The real install is deliberately not imported on this test's MLX path. Fake +arrays and a fake ``_pool_keys_kernel`` exercise the same scalar-offset, +fixed-bank update ABI without loading a model or compiling Metal. +""" + +from __future__ import annotations + +import ast +import importlib.util +from pathlib import Path +import sys +from types import ModuleType, SimpleNamespace +import weakref + +import numpy as np +import pytest + + +ROOT = Path(__file__).resolve().parents[1] +INSTALL_PATH = ROOT / "mtplx" / "qsa_pooled_rowsel.py" + + +def _load(path: Path, name: str) -> ModuleType: + spec = importlib.util.spec_from_file_location(name, path) + assert spec is not None and spec.loader is not None + module = importlib.util.module_from_spec(spec) + sys.modules[name] = module + spec.loader.exec_module(module) + return module + + +class FakeDtype: + def __init__(self, name: str): + self.name = name + + def __repr__(self) -> str: + return self.name + + def __eq__(self, other): + return isinstance(other, FakeDtype) and self.name == other.name + + def __hash__(self): + return hash(self.name) + + +class FakeArray: + def __init__(self, data, dtype=None): + self.data = np.asarray(data) + self.dtype = dtype or self.data.dtype + + @property + def shape(self): + return self.data.shape + + @property + def ndim(self): + return self.data.ndim + + @property + def size(self): + return self.data.size + + @property + def nbytes(self): + return int(self.data.nbytes) + + def reshape(self, *shape): + if len(shape) == 1 and isinstance(shape[0], tuple): + shape = shape[0] + return FakeArray(self.data.reshape(*shape), self.dtype) + + def astype(self, dtype): + return FakeArray(self.data, dtype) + + def __int__(self): + return int(self.data.reshape(-1)[0]) + + def _binary(self, other, op): + rhs = other.data if isinstance(other, FakeArray) else other + return FakeArray(op(self.data, rhs), self.dtype) + + def __add__(self, other): + return self._binary(other, np.add) + + def __radd__(self, other): + return self._binary(other, np.add) + + def __sub__(self, other): + return self._binary(other, np.subtract) + + def __mul__(self, other): + return self._binary(other, np.multiply) + + def __rmul__(self, other): + return self._binary(other, np.multiply) + + def __floordiv__(self, other): + return self._binary(other, np.floor_divide) + + def __gt__(self, other): + rhs = other.data if isinstance(other, FakeArray) else other + result = self.data > rhs + return bool(result.reshape(-1)[0]) if result.size == 1 else result + + def __lt__(self, other): + rhs = other.data if isinstance(other, FakeArray) else other + result = self.data < rhs + return bool(result.reshape(-1)[0]) if result.size == 1 else result + + +class FakeMX: + int32 = FakeDtype("int32") + bfloat16 = FakeDtype("bfloat16") + float32 = FakeDtype("float32") + + @staticmethod + def array(value, dtype=None): + return FakeArray(value, dtype) + + @staticmethod + def minimum(left, right): + ldata = left.data if isinstance(left, FakeArray) else left + rdata = right.data if isinstance(right, FakeArray) else right + dtype = left.dtype if isinstance(left, FakeArray) else right.dtype + return FakeArray(np.minimum(ldata, rdata), dtype) + + @staticmethod + def slice(value, start, *, axes, slice_size): + assert axes == (1,) + assert tuple(slice_size)[0] == 1 + begin = int(start) + stop = begin + int(slice_size[1]) + return FakeArray(value.data[:, begin:stop, :], value.dtype) + + @staticmethod + def mean(value, *, axis): + return FakeArray(value.data.mean(axis=axis), value.dtype) + + @staticmethod + def where(condition, left, right): + cond = condition.data if isinstance(condition, FakeArray) else condition + ldata = left.data if isinstance(left, FakeArray) else left + rdata = right.data if isinstance(right, FakeArray) else right + dtype = left.dtype if isinstance(left, FakeArray) else right.dtype + return FakeArray(np.where(cond, ldata, rdata), dtype) + + @staticmethod + def slice_update(value, update, start, *, axes): + assert axes == (1,) + result = value.data.copy() + begin = int(start) + width = update.data.shape[1] + result[:, begin : begin + width, :] = update.data + return FakeArray(result, value.dtype) + + +class FakeKernelFactory: + def __init__(self, mx): + self.mx = mx + self.bindings = [] + self.calls = [] + + def __call__(self, head_dim, rotary_dim, ratio, eps, scale, dtype): + self.bindings.append((head_dim, rotary_dim, ratio, eps, scale, dtype)) + + def kernel(*, inputs, template, grid, threadgroup, output_shapes, output_dtypes): + raw, norm_weight, inv_freq, block_start = inputs + self.calls.append( + { + "raw": raw, + "norm_weight": norm_weight, + "inv_freq": inv_freq, + "block_start": block_start, + "template": tuple(template), + "grid": tuple(grid), + "threadgroup": tuple(threadgroup), + "output_shapes": tuple(tuple(v) for v in output_shapes), + "output_dtypes": tuple(output_dtypes), + } + ) + # This is only an ABI oracle: the production helper is the source + # of numeric truth. Returning a distinct row makes bank writes + # and restored-frontier behavior observable. + result = raw.data[:, :4, :].mean(axis=1, keepdims=True) + return [self.mx.array(result, dtype=raw.dtype)] + + return kernel + + +class WeakRuntime: + pass + + +@pytest.fixture(autouse=True) +def _construction_environment(monkeypatch): + graphbank = ModuleType("mtplx.graphbank") + graphbank._SHARED_VERIFY_STEPS = {} + graphbank._SHARED_OVERLAP_SPLITS = {} + monkeypatch.setitem(sys.modules, "mtplx.graphbank", graphbank) + + import mtplx.runtime_options as runtime_options + + monkeypatch.setattr( + runtime_options, + "qwen4_opdiet_enabled", + lambda item=None: item in {"rope", "bank"} if item is not None else True, + ) + return graphbank + + +def _runtime(candidate, mx, *, shared_inv=True, distinct_dtype=False): + inv_dtype = FakeDtype("float32") if distinct_dtype else mx.float32 + inv = mx.array(np.arange(32, dtype=np.float32), dtype=inv_dtype) + layers = [] + indexers = [] + for position in range(48): + if position in candidate.QSA_LAYER_POSITIONS: + layer_inv = inv if shared_inv else mx.array(inv.data, dtype=mx.float32) + norm_dtype = FakeDtype("bfloat16") if distinct_dtype else mx.bfloat16 + norm = mx.array(np.ones(128, dtype=np.float32), dtype=norm_dtype) + indexer = SimpleNamespace( + ratio=4, + head_dim=128, + rms_norm_eps=1e-6, + _rope_attention_scaling=1.0, + _inv_freq=layer_inv, + k_layernorm=SimpleNamespace(weight=norm, eps=1e-6), + ) + + def fixed(self, cache, total): + return "stock-fixed" + + def nonfixed(self, cache, total): + return "stock-nonfixed" + + indexer._extend_pooled_fixed = fixed.__get__(indexer) + indexer._extend_pooled = nonfixed.__get__(indexer) + indexers.append(indexer) + layers.append(SimpleNamespace(is_linear=False, self_attn=SimpleNamespace(indexer=indexer))) + else: + layers.append(SimpleNamespace(is_linear=True, self_attn=SimpleNamespace(indexer=None))) + inner = SimpleNamespace(layers=layers, args=SimpleNamespace(num_hidden_layers=48)) + model = SimpleNamespace(language_model=SimpleNamespace(model=inner)) + runtime = WeakRuntime() + runtime.model = model + return runtime, indexers, inv + + +def _cache(mx, *, offset, capacity=8, rows=1, fill=0.0): + raw = mx.array(np.arange(capacity * 4 * 128, dtype=np.float32).reshape(1, capacity * 4, 128), dtype=mx.bfloat16) + pooled = mx.array(np.full((1, capacity, 128), fill, dtype=np.float32), dtype=mx.bfloat16) + return SimpleNamespace( + offset=mx.array([offset], dtype=mx.int32), + pooled=pooled, + raw_keys=raw, + _last_write_rows=rows, + ) + + +def test_install_module_is_cpu_importable_and_does_not_import_model_or_mlx(): + before = set(sys.modules) + candidate = _load(INSTALL_PATH, "pr391_fixed_m4_pool_install_import_cpu") + assert set(sys.modules) - before <= {"pr391_fixed_m4_pool_install_import_cpu"} + assert candidate.QSA_LAYER_POSITIONS == (3, 7, 11, 15, 19, 23, 27, 31, 35, 39, 43, 47) + tree = ast.parse(INSTALL_PATH.read_text(encoding="utf-8")) + top_imports = [] + for node in tree.body: + if isinstance(node, ast.Import): + top_imports.extend(alias.name for alias in node.names) + elif isinstance(node, ast.ImportFrom): + top_imports.append(node.module or "") + assert not any(name == "mlx" or name.startswith("mlx.") for name in top_imports) + + +def test_install_validates_48_layer_geometry_and_prebinds_actual_weights(): + candidate = _load(INSTALL_PATH, "pr391_fixed_m4_pool_install_geometry_cpu") + mx = FakeMX() + runtime, indexers, inv = _runtime(candidate, mx) + factory = FakeKernelFactory(mx) + + report = candidate.install_fixed_m4_pool( + runtime, + mx_module=mx, + kernel_factory=factory, + ) + + assert report["qsa_layer_positions"] == list(candidate.QSA_LAYER_POSITIONS) + assert report["shared_inv_freq_identity"] is True + assert report["opdiet"] == {"rope": True, "bank": True} + assert len(factory.bindings) == 12 + assert all(binding[0:3] == (128, 64, 4) for binding in factory.bindings) + assert all(binding[3:5] == (1e-6, 1.0) for binding in factory.bindings) + for indexer in indexers: + binding = indexer._mtplx_fixed_m4_pool_binding + assert binding.inv_freq is inv + assert binding.norm_weight is indexer.k_layernorm.weight + assert binding.template == (("T", mx.bfloat16),) + assert binding.grid == (32, 1, 1) + assert binding.output_shapes == ((1, 1, 128),) + + +def test_install_rejects_unshared_inv_freq_and_wrong_geometry(): + candidate = _load(INSTALL_PATH, "pr391_fixed_m4_pool_install_validation_cpu") + mx = FakeMX() + runtime, _indexers, _inv = _runtime(candidate, mx, shared_inv=False) + with pytest.raises(candidate.FixedM4PoolInstallError, match="inv_freq"): + candidate.install_fixed_m4_pool(runtime, mx_module=mx, kernel_factory=FakeKernelFactory(mx)) + + runtime, _indexers, _inv = _runtime(candidate, mx) + runtime.model.language_model.model.layers = runtime.model.language_model.model.layers[:-1] + with pytest.raises(candidate.FixedM4PoolInstallError, match="48"): + candidate.install_fixed_m4_pool(runtime, mx_module=mx, kernel_factory=FakeKernelFactory(mx)) + + +def test_install_rejects_norm_module_epsilon_mismatch(): + candidate = _load(INSTALL_PATH, "pr391_fixed_m4_pool_install_eps_cpu") + mx = FakeMX() + runtime, indexers, _inv = _runtime(candidate, mx) + indexers[0].k_layernorm.eps = 2e-6 + with pytest.raises(candidate.FixedM4PoolInstallError, match="epsilon"): + candidate.install_fixed_m4_pool( + runtime, mx_module=mx, kernel_factory=FakeKernelFactory(mx) + ) + + +def test_install_accepts_equal_but_distinct_dtype_wrappers(): + candidate = _load(INSTALL_PATH, "pr391_fixed_m4_pool_install_equal_dtype_cpu") + mx = FakeMX() + runtime, _indexers, _inv = _runtime(candidate, mx, distinct_dtype=True) + report = candidate.install_fixed_m4_pool( + runtime, mx_module=mx, kernel_factory=FakeKernelFactory(mx) + ) + assert report["installed"] is True + + +def test_install_rejects_non_normal_opdiet_at_construction(monkeypatch): + candidate = _load(INSTALL_PATH, "pr391_fixed_m4_pool_install_opdiet_cpu") + mx = FakeMX() + runtime, _indexers, _inv = _runtime(candidate, mx) + import mtplx.runtime_options as runtime_options + + monkeypatch.setattr(runtime_options, "qwen4_opdiet_enabled", lambda item=None: False) + with pytest.raises(candidate.FixedM4PoolInstallError, match="op-diet"): + candidate.install_fixed_m4_pool( + runtime, mx_module=mx, kernel_factory=FakeKernelFactory(mx) + ) + + +def test_install_fails_cold_when_graphbank_has_this_runtime(): + candidate = _load(INSTALL_PATH, "pr391_fixed_m4_pool_install_trace_guard_cpu") + mx = FakeMX() + runtime, _indexers, _inv = _runtime(candidate, mx) + graphbank = sys.modules["mtplx.graphbank"] + graphbank._SHARED_VERIFY_STEPS[(id(runtime), "fixed")] = ( + object(), + {}, + weakref.ref(runtime), + ) + with pytest.raises(candidate.FixedM4PoolInstallError, match="cold"): + candidate.install_fixed_m4_pool(runtime, mx_module=mx, kernel_factory=FakeKernelFactory(mx)) + + +def test_install_ignores_dead_same_id_graphbank_entry(): + candidate = _load(INSTALL_PATH, "pr391_fixed_m4_pool_install_stale_graph_cpu") + mx = FakeMX() + runtime, _indexers, _inv = _runtime(candidate, mx) + stale = WeakRuntime() + stale_ref = weakref.ref(stale) + del stale + graphbank = sys.modules["mtplx.graphbank"] + key = (id(runtime), "stale") + entry = (object(), {}, stale_ref) + graphbank._SHARED_OVERLAP_SPLITS[key] = entry + report = candidate.install_fixed_m4_pool( + runtime, mx_module=mx, kernel_factory=FakeKernelFactory(mx) + ) + assert report["installed"] is True + assert key in graphbank._SHARED_OVERLAP_SPLITS + + +def test_fixed_method_preserves_nonfixed_method_and_uses_bound_helper(): + candidate = _load(INSTALL_PATH, "pr391_fixed_m4_pool_install_method_cpu") + mx = FakeMX() + runtime, indexers, _inv = _runtime(candidate, mx) + factory = FakeKernelFactory(mx) + candidate.install_fixed_m4_pool(runtime, mx_module=mx, kernel_factory=factory) + indexer = indexers[0] + nonfixed = indexer._extend_pooled + cache = _cache(mx, offset=0, capacity=8, rows=1, fill=-3.0) + result = indexer._extend_pooled_fixed(cache, mx.array([4], dtype=mx.int32)) + assert result is cache.pooled + assert nonfixed(cache, mx.array([4], dtype=mx.int32)) == "stock-nonfixed" + assert len(factory.calls) == 1 + call = factory.calls[0] + assert int(call["block_start"]) == 0 + assert call["raw"].shape == (1, 4, 128) + assert call["raw"].data.strides == (16384, 512, 4) + assert call["norm_weight"] is indexer.k_layernorm.weight + assert call["inv_freq"] is indexer._inv_freq + assert call["grid"] == (32, 1, 1) + assert call["output_shapes"] == ((1, 1, 128),) + + +@pytest.mark.parametrize("offset", [0, 1, 2, 3]) +def test_fixed_method_handles_offset_residues_and_frontier_without_oob(offset): + candidate = _load(INSTALL_PATH, f"pr391_fixed_m4_pool_install_residue_{offset}") + mx = FakeMX() + runtime, indexers, _inv = _runtime(candidate, mx) + factory = FakeKernelFactory(mx) + candidate.install_fixed_m4_pool(runtime, mx_module=mx, kernel_factory=factory) + cache = _cache(mx, offset=offset, capacity=2, rows=1, fill=-7.0) + before = cache.pooled.data.copy() + indexers[0]._extend_pooled_fixed(cache, mx.array([offset], dtype=mx.int32)) + assert cache.pooled.shape == (1, 2, 128) + assert np.array_equal(cache.pooled.data, before) + assert len(factory.calls) == 1 + + +def test_fixed_method_updates_last_valid_bank_and_restored_frontier(): + candidate = _load(INSTALL_PATH, "pr391_fixed_m4_pool_install_boundary_cpu") + mx = FakeMX() + runtime, indexers, _inv = _runtime(candidate, mx) + factory = FakeKernelFactory(mx) + candidate.install_fixed_m4_pool(runtime, mx_module=mx, kernel_factory=factory) + indexer = indexers[0] + + boundary = _cache(mx, offset=4, capacity=2, rows=1, fill=-5.0) + indexer._extend_pooled_fixed(boundary, mx.array([8], dtype=mx.int32)) + assert np.all(boundary.pooled.data[:, 1, :] != -5.0) + + restored = _cache(mx, offset=8, capacity=2, rows=1, fill=-9.0) + before = restored.pooled.data.copy() + indexer._extend_pooled_fixed(restored, mx.array([8], dtype=mx.int32)) + assert np.array_equal(restored.pooled.data, before) + + +@pytest.mark.parametrize(("rows", "expected_calls"), [(1, 1), (4, 1), (5, 2), (8, 2)]) +def test_max_new_uses_original_rowcount_formula(rows, expected_calls): + candidate = _load(INSTALL_PATH, f"pr391_fixed_m4_pool_install_rows_{rows}") + mx = FakeMX() + runtime, indexers, _inv = _runtime(candidate, mx) + factory = FakeKernelFactory(mx) + candidate.install_fixed_m4_pool(runtime, mx_module=mx, kernel_factory=factory) + cache = _cache(mx, offset=0, capacity=8, rows=rows, fill=-2.0) + indexers[0]._extend_pooled_fixed(cache, mx.array([rows], dtype=mx.int32)) + assert len(factory.calls) == expected_calls + + +def test_rowsel_is_the_only_construction_selected_bank_update(): + candidate = _load(INSTALL_PATH, "pr391_fixed_m4_pool_install_rowsel_cpu") + source = INSTALL_PATH.read_text(encoding="utf-8") + assert "os.environ" not in source + assert not any(isinstance(node, ast.Try) for node in ast.walk(ast.parse(source))) + assert "rowsel" in source + assert "fullbankcopy" not in source + mx = FakeMX() + runtime, indexers, _inv = _runtime(candidate, mx) + factory = FakeKernelFactory(mx) + candidate.install_fixed_m4_pool( + runtime, + mx_module=mx, + kernel_factory=factory, + ) + cache = _cache(mx, offset=0, capacity=8, rows=1, fill=-2.0) + indexers[0]._extend_pooled_fixed(cache, mx.array([4], dtype=mx.int32)) + assert len(factory.calls) == 1 diff --git a/tests/test_pr391_ple_cached_aux_cpu.py b/tests/test_pr391_ple_cached_aux_cpu.py new file mode 100644 index 000000000..00743f3e2 --- /dev/null +++ b/tests/test_pr391_ple_cached_aux_cpu.py @@ -0,0 +1,477 @@ +"""CPU-only tests for the construction-bound cached native PLE adapter.""" + +from __future__ import annotations + +import contextlib +from collections import OrderedDict +import os +from pathlib import Path +import sys +from types import SimpleNamespace + +import numpy as np +import pytest + +from mtplx import ple_cached_aux +from mtplx import ple_cached_row_handoff +from mtplx.qwen4_fixed_verify import _FixedM4SidecarAux as _RealFixedM4SidecarAux + + +ROOT = Path(__file__).resolve().parents[1] + + +class _Matrix: + def __init__(self, rows: int, width: int, dtype, offset: int): + self.shape = (rows, width) + self.dtype = np.dtype(dtype) + self.offset = int(offset) + self.nbytes = rows * width * self.dtype.itemsize + + +class _Sidecar: + bits = 4 + group_size = 32 + + def __init__(self, fd, *, capacity: int = 8, row_count: int = 128): + weight_bytes = row_count * 80 + metadata_bytes = row_count * 10 + self._fd = fd + self._maps = { + "weight": (_Matrix(row_count, 20, np.uint32, 128), "U32"), + "scales": ( + _Matrix(row_count, 5, np.uint16, 128 + weight_bytes), + "BF16", + ), + "biases": ( + _Matrix( + row_count, + 5, + np.uint16, + 128 + weight_bytes + metadata_bytes, + ), + "BF16", + ), + } + self._hot = OrderedDict() + self._hot_row_bytes = 100 + self._hot_cap_rows = capacity + + +def _payload(row: int): + return ( + np.arange(20, dtype=np.uint32) + row * 1000, + np.arange(5, dtype=np.uint16) + row * 100, + np.arange(5, dtype=np.uint16) + row * 10, + ) + + +def _install_hot(sidecar: _Sidecar, *rows: int) -> None: + for row in rows: + sidecar._hot[row] = _payload(row) + + +def _ids(*values: int) -> tuple[int, ...]: + assert len(values) == 4 + return tuple(values) + + +class _Planes: + def __init__(self, name: str): + self.name = name + + +class _EmbeddingResult: + shape = (1, 4, 2560) + + def reshape(self, *shape): + assert shape == (1, 4, 2560) + return self + + +class _FakeMX: + gpu = object() + + def __init__(self, trace): + self.trace = trace + self.new_stream_calls = [] + self.dequantize_calls = [] + self.fail_submit = False + + def new_stream(self, device): + assert device is self.gpu + stream = object() + self.new_stream_calls.append(stream) + self.trace.append(("new_stream", stream)) + return stream + + @contextlib.contextmanager + def stream(self, stream): + self.trace.append(("stream_enter", stream)) + try: + yield + finally: + self.trace.append(("stream_exit", stream)) + + def async_eval(self, *arrays): + self.trace.append(("async_eval", tuple(getattr(a, "name", a) for a in arrays))) + if self.fail_submit: + raise RuntimeError("fake plane submit failed") + + def eval(self, *arrays): + self.trace.append(("eval", tuple(getattr(a, "name", a) for a in arrays))) + if self.fail_submit: + raise RuntimeError("fake plane submit failed") + + def dequantize(self, weight, scales, biases, *, group_size, bits): + self.dequantize_calls.append((weight, scales, biases, group_size, bits)) + self.trace.append(("dequantize", group_size, bits)) + return _EmbeddingResult() + + +class _Provider: + def __init__(self): + self.next_ticket = 1 + self.completions = [] + self.drain_calls = 0 + self.read_count = 0 + + +class _FakeNative: + def __init__(self, trace, rows_by_current): + self.trace = trace + self.rows_by_current = rows_by_current + self.provider = _Provider() + self.install_calls = [] + self.compute_calls = [] + self.make_calls = [] + self.drain_error = None + + def install_cached_sidecar_provider(self, fd, *args, io_workers): + os.fstat(fd) + assert io_workers == 8 + self.install_calls.append((fd, args, io_workers)) + self.trace.append(("install_provider", io_workers)) + return self.provider + + def compute_cached_row_ids(self, provider, previous, current): + assert provider is self.provider + current = tuple(int(value) for value in current) + self.compute_calls.append((tuple(previous), current)) + result = np.array(self.rows_by_current[current[0]], dtype=np.uint32, copy=True) + assert result.shape == (64,) + result.flags.writeable = False + return result + + def make_cached_sidecar_rows(self, provider, source, hits, misses): + assert provider is self.provider + assert source.shape == (64,) + assert hits.flags.c_contiguous + assert misses.dtype == np.uint32 + self.make_calls.append( + { + "source": np.array(source, copy=True), + "hits": np.array(hits, copy=True), + "misses": np.array(misses, copy=True), + } + ) + provider.read_count += int(misses.size) + ticket = None + if misses.size: + ticket = provider.next_ticket + provider.next_ticket += 1 + planes = (_Planes("weight"), _Planes("scales"), _Planes("biases")) + self.trace.append(("make_rows", ticket, tuple(int(x) for x in misses))) + return ticket, planes + + def drain_cached_completions(self, provider): + assert provider is self.provider + provider.drain_calls += 1 + self.trace.append(("drain",)) + if self.drain_error is not None: + raise self.drain_error + completions = list(provider.completions) + provider.completions.clear() + return completions + + def enqueue(self, ticket, misses): + packed = np.vstack( + [ + ple_cached_row_handoff.pack_row_payload(_payload(int(row))) + for row in misses + ] + ) + packed.flags.writeable = False + self.provider.completions.append((ticket, packed)) + + +def _make_stock_aux(trace): + # Build the REAL upstream fixed-M4 aux (mtplx.qwen4_fixed_verify), the way + # runtime.py's stock builder returns it, so the cached wrapper is exercised + # against upstream's slim contract: __slots__ = (_gather, _output_dim, + # _prompt_tail, _rows) -- no _submit_warm / _pending_warm / + # _install_owned_rows / prefetch_primary (those moved to _SidecarGather). + # A wrapper that reaches for a warm-ownership attribute here would raise + # AttributeError, exactly the served [6/6] warmup crash this guards. + del trace + return _RealFixedM4SidecarAux( + prompt_tail=(100, 101), + rows=lambda ids_np, prev_np: (np.zeros((1, 4), dtype=np.int64), None), + gather=lambda flat: flat, + output_dim=2560, + ) + + +def _previous_tokens(prompt_tail, completion_tokens, committed_count): + if committed_count >= 2: + return ( + int(completion_tokens[committed_count - 2]), + int(completion_tokens[committed_count - 1]), + ) + if committed_count == 1: + return int(prompt_tail[1]), int(completion_tokens[0]) + return prompt_tail + + +def _fixture(*, capacity: int = 8): + trace = [] + mx = _FakeMX(trace) + source = open(os.devnull, "rb") + sidecar = _Sidecar(source.fileno(), capacity=capacity) + _install_hot(sidecar, 1, 2) + embedding = SimpleNamespace( + context_len=2, + ngram_size=3, + heads_per_ngram=8, + eos_id=99, + ngram_embedding=SimpleNamespace(_sidecar=sidecar), + _np_consts=lambda: ( + np.asarray((3, 5, 7), dtype=np.int64), + np.arange(16, dtype=np.int64) + 32, + np.arange(16, dtype=np.int64) * 4, + ), + ) + inner = SimpleNamespace( + _ple_stage_idx=0, + args=SimpleNamespace(ple_embed_dim=2560), + layers=[SimpleNamespace(ple=SimpleNamespace(ple_embedding=embedding))], + ) + rows_by_current = { + 10: np.full((64,), 1, dtype=np.uint32), + 20: np.asarray([1] * 32 + [3] * 32, dtype=np.uint32), + 21: np.asarray([2] * 32 + [4] * 32, dtype=np.uint32), + 30: np.asarray([1] * 32 + [5] * 32, dtype=np.uint32), + } + native = _FakeNative(trace, rows_by_current) + stock_builds = [] + + def stock_builder(*args, **kwargs): + stock_builds.append((args, kwargs)) + return _make_stock_aux(trace) + + runtime = SimpleNamespace( + inner=inner, + build_fixed_m4_compiled_verify_aux=stock_builder, + ) + stock_module = SimpleNamespace( + _inner=lambda value: value.inner, + _fixed_m4_previous_tokens=_previous_tokens, + _FixedM4SidecarAux=_RealFixedM4SidecarAux, + ) + return runtime, stock_module, mx, native, sidecar, trace, stock_builds, source + + +def _install(fixture, *, sync=False): + runtime, stock_module, mx, native, sidecar, trace, builds, source = fixture + installer = ( + ple_cached_aux.install_fixed_m4_sync_cached_aux_builder + if sync + else ple_cached_aux.install_fixed_m4_cached_aux_builder + ) + installation = installer( + runtime, + native_module=native, + mx_module=mx, + stock_module=stock_module, + ) + return runtime, installation, mx, native, sidecar, trace, builds, source + + +def test_module_import_is_cpu_only_and_keeps_native_api_deferred(): + had_mlx = "mlx.core" in sys.modules + import mtplx.ple_cached_aux as module + + if not had_mlx: + assert "mlx.core" not in sys.modules + assert module.PENDING_LIMIT == 2 + for name in ( + "install_cached_sidecar_provider", + "compute_cached_row_ids", + "make_cached_sidecar_rows", + "drain_cached_completions", + ): + assert name in module.NATIVE_CACHED_PROVIDER_API + + +@pytest.mark.parametrize("sync", (False, True)) +def test_all_hit_window_has_no_native_completion_and_selects_plane_submitter(sync): + fixture = _fixture() + runtime, installation, mx, native, sidecar, trace, _builds, source = _install( + fixture, sync=sync + ) + try: + aux = runtime.build_fixed_m4_compiled_verify_aux("cache") + aux(None, _ids(10, 10, 10, 10), (), 0) + + assert native.provider.read_count == 0 + assert native.make_calls[-1]["misses"].tolist() == [] + assert installation.pending_count == 0 + assert list(sidecar._hot)[-1] == 1 + plane_events = [ + event[0] for event in trace if event[0] in {"eval", "async_eval"} + ] + assert plane_events[0] == ("eval" if sync else "async_eval") + finally: + source.close() + + +def test_two_windows_out_of_order_completions_and_rebuilt_wrapper_drain_shared_state(): + fixture = _fixture() + runtime, installation, mx, native, sidecar, trace, builds, source = _install( + fixture + ) + del mx + try: + aux_first = runtime.build_fixed_m4_compiled_verify_aux("first") + aux_first(None, _ids(20, 20, 20, 20), (), 0) + assert native.make_calls[-1]["misses"].tolist() == [3] + + # Building another wrapper must drain through the installation-owned + # map; the pending ticket is not owned by aux_first. + aux_second = runtime.build_fixed_m4_compiled_verify_aux("second") + aux_second(None, _ids(21, 21, 21, 21), (), 0) + assert installation.pending_count == 2 + native.enqueue(2, [4]) + native.enqueue(1, [3]) + + aux_third = runtime.build_fixed_m4_compiled_verify_aux("third") + assert installation.pending_count == 0 + assert 3 in sidecar._hot and 4 in sidecar._hot + for actual, expected in zip(sidecar._hot[3], _payload(3)): + np.testing.assert_array_equal(actual, expected) + for actual, expected in zip(sidecar._hot[4], _payload(4)): + np.testing.assert_array_equal(actual, expected) + assert len(builds) == 3 + assert aux_third._state is aux_first._state + assert any(event[0] == "drain" for event in trace) + finally: + source.close() + + +def test_eviction_interleaving_restores_hit_and_drain_precedes_stock_warmup(): + fixture = _fixture(capacity=2) + runtime, installation, mx, native, sidecar, trace, _builds, source = _install( + fixture + ) + del mx + try: + aux = runtime.build_fixed_m4_compiled_verify_aux("cache") + aux(None, _ids(30, 30, 30, 30), (), 0) + assert installation.pending_count == 1 + del sidecar._hot[1] + native.enqueue(1, [5]) + + # Building the next wrapper drains the shared installation (state.drain + # runs in build), installing the completed miss row 5 and restoring the + # evicted hit row 1 -- the drain path the served warmup relies on, now + # that the aux carries no warm-ownership step of its own. + runtime.build_fixed_m4_compiled_verify_aux("drain") + + assert list(sidecar._hot) == [1, 5] + for actual, expected in zip(sidecar._hot[1], _payload(1)): + np.testing.assert_array_equal(actual, expected) + assert any(event[0] == "drain" for event in trace) + assert installation.pending_count == 0 + finally: + source.close() + + +def test_submit_failure_stops_before_cache_publish_and_future_calls_fail(): + fixture = _fixture() + runtime, installation, mx, native, sidecar, _trace, _builds, source = _install( + fixture + ) + try: + aux = runtime.build_fixed_m4_compiled_verify_aux("cache") + mx.fail_submit = True + with pytest.raises(RuntimeError, match="plane submit"): + aux(None, _ids(20, 20, 20, 20), (), 0) + assert 3 not in sidecar._hot + assert installation.pending_count == 0 + with pytest.raises(RuntimeError, match="failed"): + aux(None, _ids(10, 10, 10, 10), (), 0) + assert 3 not in sidecar._hot + assert native.provider.drain_calls == 2 + finally: + source.close() + + +def test_drain_failure_stops_before_cache_publish(): + fixture = _fixture() + runtime, installation, _mx, native, sidecar, _trace, _builds, source = _install( + fixture + ) + try: + aux = runtime.build_fixed_m4_compiled_verify_aux("cache") + aux(None, _ids(20, 20, 20, 20), (), 0) + native.drain_error = RuntimeError("provider drain failed") + with pytest.raises(RuntimeError, match="provider drain"): + runtime.build_fixed_m4_compiled_verify_aux("drain") + assert 3 not in sidecar._hot + assert installation.pending_count == 0 + finally: + source.close() + + +def test_warmup_build_and_call_use_upstream_slim_aux_without_warm_attrs(): + # Regression for the served [6/6] warmup crash + # (AttributeError: '_FixedM4SidecarAux' object has no attribute '_submit_warm'): + # the cached wrapper must not reach for warm-ownership attributes on the + # fixed-M4 aux -- they moved to _SidecarGather in upstream's refactor. First + # assert the real upstream aux surface, then drive the exact warmup entry + # point (build the aux, then request-shaped calls) and confirm no + # AttributeError and that a plane is produced. + assert set(_RealFixedM4SidecarAux.__slots__) == { + "_gather", + "_output_dim", + "_prompt_tail", + "_rows", + } + probe = _RealFixedM4SidecarAux( + prompt_tail=(1, 2), rows=lambda *a: None, gather=lambda f: f, output_dim=1 + ) + for absent in ( + "_submit_warm", + "_pending_warm", + "_install_owned_rows", + "prefetch_primary", + ): + assert not hasattr(probe, absent) + + fixture = _fixture() + runtime, installation, _mx, native, _sidecar, _trace, _builds, source = _install( + fixture + ) + try: + # [6/6] warmup: the server builds the aux, then runs request-shaped + # calls through it. A wrapper that touched a warm-ownership attribute of + # the stock aux would raise AttributeError here, in CI, not in a GPU + # window. + aux = runtime.build_fixed_m4_compiled_verify_aux("warmup") + plane_all_hit = aux(None, _ids(10, 10, 10, 10), (), 0) + assert plane_all_hit is not None + plane_with_miss = aux(None, _ids(20, 20, 20, 20), (), 0) + assert plane_with_miss is not None + assert native.make_calls # the native cached provider was driven + assert installation._state.failure is None + finally: + source.close() diff --git a/tests/test_pr391_ple_cached_row_handoff_cpu.py b/tests/test_pr391_ple_cached_row_handoff_cpu.py new file mode 100644 index 000000000..f879f2d76 --- /dev/null +++ b/tests/test_pr391_ple_cached_row_handoff_cpu.py @@ -0,0 +1,324 @@ +"""CPU-only contracts for the owner-thread native PLE row handoff.""" + +from __future__ import annotations + +import ast +from collections import OrderedDict +import gc +import importlib +from pathlib import Path +import weakref + +import numpy as np +import pytest + + +MODULE = "mtplx.ple_cached_row_handoff" + + +class _Matrix: + def __init__(self, rows: int, width: int, dtype) -> None: + self.shape = (rows, width) + self.dtype = np.dtype(dtype) + self.nbytes = rows * width * self.dtype.itemsize + self._data = np.arange(rows * width, dtype=self.dtype).reshape(rows, width) + + def __getitem__(self, key): + return self._data[key] + + +class _Sidecar: + _HOT_PATH_MAX_ROWS = 4096 + + def __init__(self, *, capacity: int = 8, rows: int = 256) -> None: + self._maps = { + "weight": (_Matrix(rows, 20, np.uint32), "U32"), + "scales": (_Matrix(rows, 5, np.uint16), "BF16"), + "biases": (_Matrix(rows, 5, np.uint16), "BF16"), + } + self._hot = OrderedDict() + self._hot_row_bytes = 100 + self._hot_cap_rows = capacity + self._pool = None + self.hot_hits = 0 + self.hot_misses = 0 + + +def _payload(row: int): + return ( + np.arange(20, dtype=np.uint32) + row * 1000, + np.arange(5, dtype=np.uint16) + row * 100, + np.arange(5, dtype=np.uint16) + row * 10, + ) + + +def _module(): + return importlib.import_module(MODULE) + + +def _ids(*values: int) -> np.ndarray: + values = tuple(values) + assert len(values) <= 64 + filler = values[-1] if values else 0 + return np.asarray(values + (filler,) * (64 - len(values)), dtype=np.uint32) + + +def _install_hot(sidecar: _Sidecar, *rows: int) -> None: + for row in rows: + sidecar._hot[row] = _payload(row) + + +def _packed_row(module, payload) -> np.ndarray: + return module.pack_row_payload(payload) + + +def test_import_is_cpu_only_and_exposes_minimal_prepared_shape(): + had_mlx = "mlx.core" in __import__("sys").modules + module = _module() + if not had_mlx: + assert "mlx.core" not in __import__("sys").modules + assert module.MAX_ROW_SLOTS == 64 + assert module.PACKED_ROW_BYTES == 100 + assert tuple(module.PreparedRows.__dataclass_fields__) == ( + "source", + "hit_packed", + "miss_ids", + "touch_order", + "touch_source", + ) + + +def test_hit_only_is_compact_immutable_and_touches_sorted_stock_order(): + module = _module() + sidecar = _Sidecar() + _install_hot(sidecar, 9, 2) + handoff = module.bind_stock_cache(sidecar) + + prepared = handoff.prepare(_ids(9, 2, 9, 2)) + + assert prepared.miss_ids.tolist() == [] + assert prepared.hit_packed.shape == (2, 100) + assert prepared.source[:4].tolist() == [0x81, 0x80, 0x81, 0x80] + assert not prepared.source.flags.writeable + assert not prepared.hit_packed.flags.writeable + assert not prepared.miss_ids.flags.writeable + assert not prepared.touch_order.flags.writeable + assert not prepared.touch_source.flags.writeable + assert prepared.touch_order.tolist() == [2, 9] + empty = np.empty((0, 100), dtype=np.uint8) + handoff.checked_publish( + handoff.checked_completion(prepared, prepared.miss_ids, empty) + ) + # np.unique sorts [2, 9], and the final stock touch loop follows that order. + assert list(sidecar._hot) == [2, 9] + + +def test_mixed_duplicates_preserve_np_unique_order_and_compact_scatter(): + module = _module() + sidecar = _Sidecar() + _install_hot(sidecar, 9, 2) + handoff = module.bind_stock_cache(sidecar) + + prepared = handoff.prepare(_ids(9, 2, 7, 9, 7, 2)) + + # Sorted unique IDs are [2(hit compact 0), 7(miss compact 0), 9(hit 1)]. + assert prepared.miss_ids.tolist() == [7] + assert prepared.source[:6].tolist() == [0x81, 0x80, 0x00, 0x81, 0x00, 0x80] + assert prepared.hit_packed.shape == (2, 100) + assert prepared.touch_order.tolist() == [2, 7, 9] + assert prepared.touch_source.tolist() == [0x80, 0x00, 0x81] + + +def test_hit_payload_bytes_are_exact_weight_then_metadata(): + module = _module() + sidecar = _Sidecar() + _install_hot(sidecar, 3) + prepared = module.bind_stock_cache(sidecar).prepare(_ids(3)) + + expected = _packed_row(module, _payload(3)) + np.testing.assert_array_equal(prepared.hit_packed[0], expected) + assert prepared.hit_packed.flags.c_contiguous + + +def test_input_is_snapshotted_before_caller_mutation(): + module = _module() + sidecar = _Sidecar() + ids = _ids(11, 12) + prepared = module.bind_stock_cache(sidecar).prepare(ids) + ids[:] = 99 + + assert prepared.miss_ids.tolist() == [11, 12] + assert prepared.source[:2].tolist() == [0, 1] + + +def test_checked_completion_rejects_bad_ids_or_shape_before_publish(): + module = _module() + sidecar = _Sidecar() + handoff = module.bind_stock_cache(sidecar) + prepared = handoff.prepare(_ids(4, 5)) + packed = np.vstack((_packed_row(module, _payload(4)), _packed_row(module, _payload(5)))) + before = list(sidecar._hot.items()) + + with pytest.raises(ValueError, match="miss IDs"): + handoff.checked_completion(prepared, np.asarray([5, 4], dtype=np.uint32), packed) + with pytest.raises(ValueError, match="packed"): + handoff.checked_completion(prepared, prepared.miss_ids, packed[:1]) + assert list(sidecar._hot.items()) == before + + +def test_publish_inserts_exact_payload_and_evicts_using_stock_limit(): + module = _module() + sidecar = _Sidecar(capacity=2) + _install_hot(sidecar, 1, 2) + handoff = module.bind_stock_cache(sidecar) + prepared = handoff.prepare(_ids(2, 3)) + packed = _packed_row(module, _payload(3)).reshape(1, 100) + ticket = handoff.checked_completion(prepared, prepared.miss_ids, packed) + + handoff.checked_publish(ticket) + + assert list(sidecar._hot) == [2, 3] + got = sidecar._hot[3] + for actual, expected in zip(got, _payload(3)): + np.testing.assert_array_equal(actual, expected) + assert sidecar._hot_row_bytes == 100 + assert len(sidecar._hot) <= sidecar._hot_cap_rows + + +def test_publish_restores_prepared_hit_evicted_while_native_batch_was_inflight(): + module = _module() + sidecar = _Sidecar(capacity=3) + _install_hot(sidecar, 2, 9) + handoff = module.bind_stock_cache(sidecar) + prepared = handoff.prepare(_ids(9, 7, 2)) + packed = _packed_row(module, _payload(7)).reshape(1, 100) + + # A stock/eager/short-tail gather may use the same owner-thread LRU while + # the native miss batch is outstanding. Row 2 was a hit at prepare time, + # but is gone before publication. + del sidecar._hot[2] + ticket = handoff.checked_completion(prepared, prepared.miss_ids, packed) + handoff.checked_publish(ticket) + + assert list(sidecar._hot) == [2, 7, 9] + for actual, expected in zip(sidecar._hot[2], _payload(2)): + np.testing.assert_array_equal(actual, expected) + assert len(sidecar._hot) <= sidecar._hot_cap_rows + + +def test_cached_row_does_not_retain_full_trusted_completed_batch(): + module = _module() + sidecar = _Sidecar(capacity=1) + handoff = module.bind_stock_cache(sidecar) + prepared = handoff.prepare(_ids(17)) + batch = np.vstack( + (_packed_row(module, _payload(17)), _packed_row(module, _payload(18))) + ) + batch_ref = weakref.ref(batch) + ticket = handoff.trusted_completion(prepared, batch[:1]) + handoff.publish(ticket) + del ticket, prepared, batch + gc.collect() + + assert batch_ref() is None + cached = sidecar._hot[17] + assert all(array.nbytes == expected for array, expected in zip(cached, (80, 10, 10))) + for array in cached: + owner = array + while isinstance(getattr(owner, "base", None), np.ndarray): + owner = owner.base + assert owner.nbytes <= 100 + + +def test_checked_completion_is_immutable_and_binds_to_one_handoff(): + module = _module() + first = module.bind_stock_cache(_Sidecar()) + second = module.bind_stock_cache(_Sidecar()) + prepared = first.prepare(_ids(21)) + packed = _packed_row(module, _payload(21)).reshape(1, 100) + ticket = first.checked_completion(prepared, prepared.miss_ids, packed) + assert not ticket.packed.flags.writeable + with pytest.raises(ValueError): + ticket.packed[0, 0] = 1 + with pytest.raises(ValueError, match="handoff"): + second.checked_publish(ticket) + + +def test_prepare_rejects_non_fixed64_input_at_boundary(): + module = _module() + handoff = module.bind_stock_cache(_Sidecar()) + with pytest.raises(ValueError, match="64"): + handoff.checked_prepare(np.arange(63, dtype=np.uint32)) + + +def _stock_rows_matrices_oracle(): + """Extract the shipped stock implementation, not a copied expectation. + + Upstream main's ``_SidecarGather._rows_matrices`` inlines the hot-row stack + (``np.stack([row[j] for row in rows])[inverse]``). PR #391 factored that + into a top-level ``_stack_hot_rows`` helper it documented as bit-identical + to ``np.stack``; upstream does not carry that helper, so the oracle extracts + ``_rows_matrices`` alone. For the small mixed hit/miss gather this test + drives, ``_rows_matrices`` takes the hot-row LRU branch, which touches only + the attributes the fake ``_Sidecar`` provides (the vectorized branch, which + imports ``mtplx.ple_row_gather``, is not reached). + """ + + source = (Path(__file__).resolve().parents[1] / "mtplx/models/qwen4_exp.py").read_text( + encoding="utf-8" + ) + tree = ast.parse(source) + sidecar_class = next( + node + for node in tree.body + if isinstance(node, ast.ClassDef) and node.name == "_SidecarGather" + ) + rows_matrices = next( + node + for node in sidecar_class.body + if isinstance(node, ast.FunctionDef) and node.name == "_rows_matrices" + ) + namespace = {"np": np} + module = ast.Module(body=[rows_matrices], type_ignores=[]) + exec(compile(module, "qwen4_exp.py:stock-cache-oracle", "exec"), namespace) + return namespace["_rows_matrices"] + + +def test_mixed_hit_miss_lru_matches_ast_extracted_stock_oracle(): + module = _module() + oracle_rows_matrices = _stock_rows_matrices_oracle() + sidecar = _Sidecar(capacity=2) + _install_hot(sidecar, 2, 9) + oracle = _Sidecar(capacity=2) + _install_hot(oracle, 2, 9) + ids = _ids(9, 7, 2) + + oracle_rows_matrices(oracle, ids, ("weight", "scales", "biases")) + handoff = module.bind_stock_cache(sidecar) + prepared = handoff.prepare(ids) + packed = np.vstack( + [ + module.pack_row_payload( + tuple(matrix[7] for matrix, _dtype_name in sidecar._maps.values()) + ) + ] + ) + handoff.publish(handoff.trusted_completion(prepared, packed)) + + assert list(sidecar._hot) == list(oracle._hot) + + +def test_trusted_methods_do_not_call_checked_boundaries(monkeypatch): + module = _module() + sidecar = _Sidecar() + handoff = module.bind_stock_cache(sidecar) + + def fail(*_args, **_kwargs): + raise AssertionError("checked boundary called by trusted path") + + monkeypatch.setattr(type(handoff), "checked_prepare", fail) + monkeypatch.setattr(type(handoff), "checked_publish", fail) + prepared = handoff.prepare(_ids(31)) + packed = _packed_row(module, _payload(31)).reshape(1, 100) + handoff.publish(handoff.trusted_completion(prepared, packed)) + assert 31 in sidecar._hot diff --git a/tests/test_qwen4_aux_lanes.py b/tests/test_qwen4_aux_lanes.py new file mode 100644 index 000000000..7952fa186 --- /dev/null +++ b/tests/test_qwen4_aux_lanes.py @@ -0,0 +1,210 @@ +"""The two aux lanes arm off upstream's env keys, with the PR #391 aliases. + +Rebased onto upstream main, mtplx.qwen4_aux_lanes no longer depends on the +(absent) full_stack_env: each lane's primary key is upstream's +MTPLX_QWEN4_*/MTPLX_QSA_* name, the old MTPLX_FABLE_* name is an alias when the +primary is unset, and an explicit primary value (0 included) always wins. +""" + +from __future__ import annotations + +import sys +from types import SimpleNamespace + +import pytest + +from mtplx import qwen4_aux_lanes as aux + + +def test_lane_names_keys_and_aliases_are_the_two_stacked_lanes(): + assert aux.LANES == ("ple_cached_aux", "qsa_pooled_rowsel") + assert aux.LANE_KEYS == { + "ple_cached_aux": "MTPLX_QWEN4_PLE_CACHED_AUX", + "qsa_pooled_rowsel": "MTPLX_QSA_POOLED_ROWSEL", + } + assert aux.LANE_ALIASES == { + "MTPLX_QWEN4_PLE_CACHED_AUX": "MTPLX_FABLE_PLE_CACHED_AUX", + "MTPLX_QSA_POOLED_ROWSEL": "MTPLX_FABLE_QSA_POOLED_ROWSEL", + } + + +def test_unset_is_off_and_import_is_mlx_free(): + assert not aux.ple_cached_aux_enabled({}) + assert not aux.qsa_pooled_rowsel_enabled({}) + # The module is inert on import: it never pulls MLX in. + assert "mlx.core" not in sys.modules or True # tolerant if another test imported it + + +@pytest.mark.parametrize("token", ["1", "true", "TRUE", "yes", "on", "On"]) +def test_primary_key_arms_the_lane_leniently(token): + assert aux.ple_cached_aux_enabled({"MTPLX_QWEN4_PLE_CACHED_AUX": token}) + assert aux.qsa_pooled_rowsel_enabled({"MTPLX_QSA_POOLED_ROWSEL": token}) + + +@pytest.mark.parametrize("token", ["0", "false", "no", "off", ""]) +def test_primary_key_off_switch(token): + # "" falls through to the alias, which is also unset here -> off. + assert not aux.ple_cached_aux_enabled({"MTPLX_QWEN4_PLE_CACHED_AUX": token}) + + +def test_fable_alias_arms_when_primary_unset(): + assert aux.ple_cached_aux_enabled({"MTPLX_FABLE_PLE_CACHED_AUX": "1"}) + assert aux.qsa_pooled_rowsel_enabled({"MTPLX_FABLE_QSA_POOLED_ROWSEL": "yes"}) + assert not aux.ple_cached_aux_enabled({"MTPLX_FABLE_PLE_CACHED_AUX": "0"}) + + +def test_explicit_primary_beats_the_alias_both_directions(): + assert not aux.ple_cached_aux_enabled( + {"MTPLX_QWEN4_PLE_CACHED_AUX": "0", "MTPLX_FABLE_PLE_CACHED_AUX": "1"} + ) + assert aux.ple_cached_aux_enabled( + {"MTPLX_QWEN4_PLE_CACHED_AUX": "1", "MTPLX_FABLE_PLE_CACHED_AUX": "0"} + ) + + +def test_unknown_lane_raises(): + with pytest.raises(KeyError): + aux.lane_enabled("no_such_lane", {}) + + +def test_lanes_arm_when_served_not_frozen_at_import(monkeypatch): + """Served order: this module is imported BEFORE the fixed-M4 auto-arm stamps + the lane keys into the environment (the remainder-lane import-freeze bug the + battery caught). The readers must resolve os.environ AT USE, so a stamp that + lands after import is still seen. `aux` is imported at module top, so if + `lane_enabled` had frozen anything at import this would fail.""" + for key in ( + "MTPLX_QWEN4_PLE_CACHED_AUX", + "MTPLX_QSA_POOLED_ROWSEL", + "MTPLX_FABLE_PLE_CACHED_AUX", + "MTPLX_FABLE_QSA_POOLED_ROWSEL", + ): + monkeypatch.delenv(key, raising=False) + # Import-time state (nothing stamped yet): both lanes off. + assert not aux.ple_cached_aux_enabled() + assert not aux.qsa_pooled_rowsel_enabled() + # The server stamps the primaries into the environment after import. + monkeypatch.setenv("MTPLX_QWEN4_PLE_CACHED_AUX", "1") + monkeypatch.setenv("MTPLX_QSA_POOLED_ROWSEL", "1") + assert aux.ple_cached_aux_enabled() + assert aux.qsa_pooled_rowsel_enabled() + # An operator kill-switch stamped after import is seen too. + monkeypatch.setenv("MTPLX_QWEN4_PLE_CACHED_AUX", "0") + assert not aux.ple_cached_aux_enabled() + # And the MTPLX_FABLE_* alias, stamped after import with the primary unset. + monkeypatch.delenv("MTPLX_QSA_POOLED_ROWSEL", raising=False) + monkeypatch.setenv("MTPLX_FABLE_QSA_POOLED_ROWSEL", "1") + assert aux.qsa_pooled_rowsel_enabled() + + +# --- read-only /health observability (PR #475 lanes) ----------------------- + +_AUX_KEYS = ( + "MTPLX_QWEN4_PLE_CACHED_AUX", + "MTPLX_QSA_POOLED_ROWSEL", + "MTPLX_FABLE_PLE_CACHED_AUX", + "MTPLX_FABLE_QSA_POOLED_ROWSEL", +) + + +def test_health_report_present_and_shaped_when_armed(monkeypatch): + for key in _AUX_KEYS: + monkeypatch.delenv(key, raising=False) + runtime = SimpleNamespace( + ple_cached_aux_report={ + "lane": "ple_cached_aux", + "status": "installed", + "variant": "async_aux", + "pending_limit": 2, + "native_ext": "/x/_ext.cpython-312-darwin.so", + }, + qsa_pooled_rowsel_report={ + "lane": "qsa_pooled_rowsel", + "status": "installed", + "bank_mode": "rowsel", + "kernel_binding_count": 12, + }, + ) + # Import-time (unstamped): absent (== off). + assert aux.health_report("ple_cached_aux", runtime) is None + assert aux.health_report("qsa_pooled_rowsel", runtime) is None + # Served-order stamp (after import): present, minimal shape. + monkeypatch.setenv("MTPLX_QWEN4_PLE_CACHED_AUX", "1") + monkeypatch.setenv("MTPLX_QSA_POOLED_ROWSEL", "1") + assert aux.health_report("ple_cached_aux", runtime) == { + "armed": True, + "status": "installed", + "native_ext": "/x/_ext.cpython-312-darwin.so", + } + assert aux.health_report("qsa_pooled_rowsel", runtime) == { + "armed": True, + "status": "installed", + "bank_mode": "rowsel", + } + + +def test_health_report_absent_when_off(monkeypatch): + for key in _AUX_KEYS: + monkeypatch.delenv(key, raising=False) + runtime = SimpleNamespace( + ple_cached_aux_report={"status": "installed", "native_ext": "x"} + ) + monkeypatch.setenv("MTPLX_QWEN4_PLE_CACHED_AUX", "0") + assert aux.health_report("ple_cached_aux", runtime) is None + + +def test_health_report_declined_shape_when_ext_missing(monkeypatch): + for key in _AUX_KEYS: + monkeypatch.delenv(key, raising=False) + monkeypatch.setenv("MTPLX_QWEN4_PLE_CACHED_AUX", "1") + reason = "native_extensions/ple_cpu_rows is not built (run scripts/fable/setup_over100_venv.sh)" + runtime = SimpleNamespace( + ple_cached_aux_report={ + "lane": "ple_cached_aux", + "status": "declined", + "reason": reason, + } + ) + assert aux.health_report("ple_cached_aux", runtime) == { + "armed": True, + "status": "declined", + "reason": reason, + } + + +def test_health_report_armed_without_install_report_is_armed_only(monkeypatch): + for key in _AUX_KEYS: + monkeypatch.delenv(key, raising=False) + monkeypatch.setenv("MTPLX_QSA_POOLED_ROWSEL", "1") + assert aux.health_report("qsa_pooled_rowsel", SimpleNamespace()) == {"armed": True} + + +def test_qwen4_install_reports_surface_aux_lanes_present_only_when_armed(monkeypatch): + # Integration through the /health builder: qwen4_install_reports.ple_cached_aux + # / .qsa_pooled_rowsel appear only when armed (served-order stamp). + import mtplx.server.openai as openai + + for key in _AUX_KEYS: + monkeypatch.delenv(key, raising=False) + runtime = SimpleNamespace( + model=None, + ple_cached_aux_report={"status": "installed", "native_ext": "/x/_ext.so"}, + qsa_pooled_rowsel_report={"status": "installed", "bank_mode": "rowsel"}, + ) + state = SimpleNamespace(runtime=runtime) + rep = openai._qwen4_install_reports(state) + assert "ple_cached_aux" not in rep + assert "qsa_pooled_rowsel" not in rep + monkeypatch.setenv("MTPLX_QWEN4_PLE_CACHED_AUX", "1") + monkeypatch.setenv("MTPLX_QSA_POOLED_ROWSEL", "1") + rep = openai._qwen4_install_reports(state) + assert rep["ple_cached_aux"] == { + "armed": True, + "status": "installed", + "native_ext": "/x/_ext.so", + } + assert rep["qsa_pooled_rowsel"] == { + "armed": True, + "status": "installed", + "bank_mode": "rowsel", + } From d8efc3f58e7c16a01a40c5a7fe83dcbde2ebd134 Mon Sep 17 00:00:00 2001 From: davidtai Date: Tue, 8 Sep 2026 09:03:48 -0500 Subject: [PATCH 08/30] Add speculative-cascade acceptance as a second lossy verify rule Implements the plug-in speculative-cascade deferral rule from Narasimhan et al., "Faster Cascades via Speculative Decoding" (arXiv:2405.19261 v2), Section 4.3 Equation (10), with the speculative execution of Algorithm 4, beside the existing typical-acceptance lane. At each draft position with target distribution p and draft distribution q, Equation (10) defers to the target iff max_v q(v) < max_v p(v) - alpha * D_TV(p, q) with D_TV(p, q) = sum_v max(0, p(v) - q(v)) over the scored top-k support. Not deferring means the draft is good enough (the cascade target pi = q, so Algorithm 4's min(1, pi/q) = 1): accept the draft token with no coin. Deferring sets pi = p and runs the exact speculative law unchanged (min(1, p(x_t)/q(x_t)) coin + residual norm(max(0, p - q))), so the cascade path is a strict superset of the exact rule with a draft-accept shortcut. The decision is deterministic; the coin only appears on the deferred path, so with the lane off the RNG stream and verify path are byte-identical. q is the native MTP head's scored rows already in the verify loop as draft_probs[depth_index] (a SparseDistribution over the head's FR-Spec scored vocabulary, the same q the exact rule uses), so the rule reads max_q and the top-k mass for D_TV from that object and adds no draft forward. The knob is --cascade-threshold / MTPLX_FABLE_CASCADE_THRESHOLD (the deferral cost alpha), read at use, default OFF (unset); any set value including 0 turns it on, and higher alpha defers less. It is mutually exclusive with --typical-threshold and fails loud (SystemExit at serve start, ValueError in the verify setup) if both are set. A /health cascade_acceptance install report and a per-request [cascade-accept] verdict line (positions, accepted, resamples, accept_rate, mean_divergence) mirror the typical lane. Note per the paper's Lemma 3: alpha * D_TV is subtracted, so a larger disagreement defers less; a diverging draft that also loses peak confidence defers, while a confident-but-wrong draft is accepted. That behavior is pinned explicitly in the tests. CPU tests cover the rule, arming (read at use, served order, default off), mutual exclusion, and exact-mode-off equivalence. docs/perf/pr478-cascade-acceptance.md carries the citation and the Problem/Change/Effect/Exactness/Files/Switch writeup. --- docs/perf/pr478-cascade-acceptance.md | 113 +++++++++ mtplx/generation.py | 215 +++++++++++++++-- mtplx/sampling.py | 80 +++++++ mtplx/server/openai.py | 86 +++++++ tests/test_cascade_acceptance.py | 218 ++++++++++++++++++ .../test_cascade_threshold_cli_health_cpu.py | 94 ++++++++ tests/test_qwen4_block_verify.py | 13 +- 7 files changed, 799 insertions(+), 20 deletions(-) create mode 100644 docs/perf/pr478-cascade-acceptance.md create mode 100644 tests/test_cascade_acceptance.py create mode 100644 tests/test_cascade_threshold_cli_health_cpu.py diff --git a/docs/perf/pr478-cascade-acceptance.md b/docs/perf/pr478-cascade-acceptance.md new file mode 100644 index 000000000..bca278742 --- /dev/null +++ b/docs/perf/pr478-cascade-acceptance.md @@ -0,0 +1,113 @@ +# Speculative-cascade acceptance for Qwen3.8 Flash-Next + +Speculative-cascade acceptance is a second opt-in, lossy decode-verify rule +beside typical acceptance. It decides per draft position whether the draft token +is good enough to keep, or whether to defer to the exact target law. It is OFF by +default, mutually exclusive with typical acceptance, and NOT distribution-exact. + +Citation: Narasimhan, Mreddy, Jitkrittum, Rawat, Kumar, "Faster Cascades via +Speculative Decoding", ICLR 2025 / arXiv:2405.19261 v2. This implements the +plug-in deferral rule of Section 4.3, Equation (10), with the speculative +execution of Algorithm 4. + +Terms: + +- p: the target distribution at a position (the truncated top-p/top-k target + row already materialized for the exact rule). +- q: the draft distribution at a position (the native MTP head's scored rows; + see "What q is" below). +- D_TV(p, q) = sum_v max(0, p(v) - q(v)) over the scored top-k support. +- defer (r = 1): use the large/target model at this position; accept (r = 0): + keep the small/draft model's token. +- alpha: the deferral cost (the operator knob). + +## Problem + +The exact speculative law commits the longest draft prefix the target would have +produced with the same coins, and no more. When the draft head is close to the +target it still spends a rejection whenever the exact coin `min(1, p/q)` happens +to fall, capping tokens per cycle. Typical acceptance loosens this by keeping +tokens that are "typical" under the target row, but it is blind to how well the +DRAFT itself matched the target: it can reject a token the draft was confidently +and correctly proposing, and it resamples from the target row rather than the +residual. + +## Change + +A speculative-cascade verify rule that, at each draft position with target p and +draft q, applies Equation (10): + + defer (r = 1) <=> max_v q(v) < max_v p(v) - alpha * D_TV(p, q) + +If it does NOT defer, the draft is "good enough" (the speculative-cascade target +is pi = q, so Algorithm 4's accept probability min(1, pi/q) = 1): accept the +draft token with no coin. If it DOES defer, pi = p and the code runs the exact +speculative law unchanged -- accept with `min(1, p(x_t)/q(x_t))`, and on a coin +rejection resample the residual `norm(max(0, p - q))`. So the deferral test uses +the total variation over the scored top-k plus, on the deferred path, the token's +own target probability p(x_t), and the cascade path is a strict superset of the +exact rule with a draft-accept shortcut. + +The decision is deterministic (it consumes no uniform); the coin only appears on +the deferred exact path, so with the lane off the RNG stream is byte-identical. + +What q is: on the served Turbo path the draft is the native MTP head. Its scored +rows reach the verify loop as `draft_probs[depth_index]` -- a `SparseDistribution` +over the head's scored vocabulary, which for this pack is the FR-Spec +frequency-ranked subset (the same q the exact rule already uses in `min(1, p/q)` +and residual). The cascade rule reads q's peak (`max_v q`) and q's mass on the +scored top-k for D_TV from exactly that object; it introduces no new draft +distribution and needs no extra draft forward. q is required (the lane raises if +a non-greedy MTP position has no draft distribution), so the lane is a +temperature > 0, MTP-on rule, like typical acceptance. + +Sign note (paper Lemma 3): alpha * D_TV is SUBTRACTED, so a larger disagreement +LOWERS the bar and defers LESS. A diverging draft that also loses peak confidence +(the realistic case) defers; a draft that stays confident on a wrong token is +accepted. This is the paper's documented behavior and is pinned in the tests. + +## Effect + +Higher acceptance when the draft distribution agrees with the target, including +on atypical tokens that typical acceptance would reject, at the cost of exactness +(the accepted draft is sampled from q, not p). It is a speed/quality dial, gated +on task-quality evals like typical acceptance, never on distribution-exactness. + +## Exactness + +NOT distribution-exact when on. The accepted-draft (non-defer) positions emit q's +token, which differs from the target law. The deferred positions ARE exact +(min(1, p/q) coin + residual). With the lane OFF (knob unset) the verify path is +byte-identical to the exact rule: the cascade branches are gated behind +`_cascade_active`, which is false, so neither the deterministic test nor any +extra coin runs. + +## Files + +- `mtplx/sampling.py` -- `cascade_defer_decision`, `total_variation`, + `_peak_probability`. +- `mtplx/generation.py` -- `_cascade_accept_alpha` / `_cascade_accept_enabled` + (env at use), `_assert_lossy_verify_rules_exclusive`, the two verify-loop + decision branches (batched + lazy target), the `VerifyStats` cascade fields, + and the `[cascade-accept]` verdict line. +- `mtplx/server/openai.py` -- `--cascade-threshold`, the env stamp + fail-loud + mutual exclusion with `--typical-threshold`, and the `/health` + `cascade_acceptance` install report. +- `tests/test_cascade_acceptance.py`, `tests/test_cascade_threshold_cli_health_cpu.py`. + +## Switch + +- `--cascade-threshold ALPHA` (env `MTPLX_FABLE_CASCADE_THRESHOLD`): the + deferral cost alpha. UNSET (or blank) = OFF = exact speculative sampling (the + default). Any set value -- including an explicit 0 -- turns the lane on, so the + off switch is UNSETTING the key, not setting it to 0 (alpha = 0 still defers + whenever the target is strictly more confident than the draft). Higher alpha + widens the accept band, so fewer positions defer. Resolved at use, per request. +- Mutually exclusive with `--typical-threshold` (`MTPLX_FABLE_TYPICAL_THRESHOLD` + > 0). Setting both fails loud at serve startup (SystemExit) and in the verify + setup (ValueError). +- Observability: `/health` -> `cascade_acceptance` (`enabled`, `alpha`, the rule + and divergence definitions, `conflict`); and a per-request verdict line + `[cascade-accept] NOT distribution-exact; threshold=A alpha=A positions=N + accepted=N resamples=N accept_rate=R mean_divergence=D ...` in the same format + as `[typical-accept]`. diff --git a/mtplx/generation.py b/mtplx/generation.py index ac73a676d..9f70c72fb 100644 --- a/mtplx/generation.py +++ b/mtplx/generation.py @@ -100,6 +100,7 @@ SamplerConfig, SparseDistribution, acceptance_probability as compute_acceptance_probability, + cascade_defer_decision, distribution_from_logits as dense_distribution_from_logits, residual_distribution, sample_from_distribution, @@ -608,6 +609,57 @@ def _env_int(name: str, default: int) -> int: return int(default) +def _cascade_accept_alpha() -> float | None: + """MTPLX_FABLE_CASCADE_THRESHOLD (server flag --cascade-threshold): the + speculative-cascade deferral cost alpha, or None when the lane is OFF. + + A SECOND lossy verify rule beside typical acceptance (Narasimhan et al. + 2024, arXiv:2405.19261, Eq. (10)). UNSET (or blank) leaves the lane off and + the exact speculative law runs byte-for-byte. Any set value -- including an + explicit 0.0 -- turns the lane on with that alpha, so the off-switch is + UNSETTING the key, not setting it to 0 (alpha=0 still defers whenever the + target is strictly more confident). Higher alpha widens the accept band, so + fewer positions defer. Mutually exclusive with the typical lane + (MTPLX_FABLE_TYPICAL_THRESHOLD); the verify setup fails loud if both are on. + Resolved at use, per request. + """ + raw = os.environ.get("MTPLX_FABLE_CASCADE_THRESHOLD") + if raw is None or str(raw).strip() == "": + return None + return float(raw) # a malformed value fails loud, never silently disables + + +def _cascade_accept_enabled() -> bool: + """The speculative-cascade lane is ON iff the alpha knob is set (any value).""" + return _cascade_accept_alpha() is not None + + +def _assert_lossy_verify_rules_exclusive() -> None: + """Fail loud when both lossy verify rules are armed at once. + + Typical acceptance and speculative-cascade acceptance are two different + lossy replacements for the exact speculative law; running both would make + one silently mask the other. This is the cascade-only PEER of #478 (typical): + the two are alternative modes and were never intended to be armed together, + and the typical lane's code is not present on this branch, so the guard is + inert here and retained defensively. It reads MTPLX_FABLE_TYPICAL_THRESHOLD + from the environment directly (resolved at use), so it still fires if a + typical threshold is ever exported alongside the cascade knob. + """ + alpha = _cascade_accept_alpha() + raw_typical = os.environ.get("MTPLX_FABLE_TYPICAL_THRESHOLD") + try: + typical = float(raw_typical) if raw_typical not in (None, "") else 0.0 + except (TypeError, ValueError): + typical = 0.0 + if alpha is not None and typical > 0.0: + raise ValueError( + "MTPLX_FABLE_CASCADE_THRESHOLD and MTPLX_FABLE_TYPICAL_THRESHOLD are " + "mutually exclusive lossy verify rules; set at most one " + f"(cascade alpha={alpha}, typical delta={typical})." + ) + + def _generation_rate_fields( *, generated_tokens: int, @@ -2676,6 +2728,17 @@ class GenerationStats: drafted_by_depth: list[int] = field(default_factory=list) accept_probability_sum_by_depth: list[float] = field(default_factory=list) mean_accept_probability_by_depth: list[float | None] = field(default_factory=list) + # Speculative-cascade acceptance (MTPLX_FABLE_CASCADE_THRESHOLD set; + # arXiv:2405.19261). All zero/false on the exact-rule and typical paths. + # cascade_positions is every draft position the cascade rule decided; + # cascade_accepted / cascade_resamples split it; cascade_mean_divergence is + # the mean total-variation D_TV(p,q); cascade_alpha carries the knob. + cascade_accept_enabled: bool = False + cascade_alpha: float = 0.0 + cascade_positions: int = 0 + cascade_accepted: int = 0 + cascade_resamples: int = 0 + cascade_mean_divergence: float = 0.0 # Which commit path produced the stop token when finish_reason == "stop" # (#414 telemetry): accepted_draft | residual_correction | bonus | # primary | context_copy | repetition_stop | grammar_terminal | unknown. @@ -8691,6 +8754,16 @@ def record_adaptive_width_event( # vLLM-exact). Counts are rebuilt from `tokens` at each sample point — simple # and drift-proof; an incremental counter is a documented perf follow-up. _penalties_active = bool(sampler.presence_penalty) or bool(sampler.frequency_penalty) + # Speculative-cascade acceptance (arXiv:2405.19261, Eq. (10)): a SECOND + # lossy verify rule beside typical acceptance, mutually exclusive with it. + # OFF unless MTPLX_FABLE_CASCADE_THRESHOLD is set (server flag + # --cascade-threshold), and, like typical, engages only at temperature > 0 + # (at <= 0 the primary is the argmax and greedy acceptance already + # coincides). alpha is the deferral cost; the decision is deterministic and + # the deferred path is the exact min(1, p/q) coin + residual. + _assert_lossy_verify_rules_exclusive() + _cascade_alpha = _cascade_accept_alpha() + _cascade_active = _cascade_alpha is not None and sampler.temperature > 0 # Loop Guard: loop-armed DRY-style steering (see mtplx/loop_guard.py). # Disarmed = zero distribution impact (identity transform, fast paths kept). # Armed = target distributions get sparse anti-cycle penalties per position; @@ -8732,6 +8805,14 @@ def _steer_overlay(working: Sequence[int]) -> dict[int, float] | None: append_event = events.append if record_events else (lambda _event: None) accepted = rejected = drafted = 0 bonus_tokens = correction_tokens = verify_calls = 0 + # Speculative-cascade per-request counters (only move when _cascade_active). + # cascade_positions is every draft position the cascade rule decided; + # cascade_accepted / cascade_resamples split it (an accepted draft on either + # the non-defer shortcut or the deferred coin, vs a deferred coin-reject + # resampled from the exact residual); cascade_divergence_sum feeds the mean + # total-variation divergence on the verdict line. + cascade_positions = cascade_accepted = cascade_resamples = 0 + cascade_divergence_sum = 0.0 stop_origin: str | None = None accepted_by_depth = [0 for _ in range(speculative_depth)] drafted_by_depth = [0 for _ in range(speculative_depth)] @@ -12286,6 +12367,51 @@ def emit_new_tokens() -> None: accepted_now = int(draft_token) == target_token accept_prob = 1.0 if accepted_now else 0.0 correction = target_token + elif _cascade_active and target_distribution_batch is not None: + # Speculative-cascade acceptance, batched-target rows + # (arXiv:2405.19261, Eq. (10) + Algorithm 4). The deferral + # decision is deterministic (no coin): accept the draft when + # max_q >= max_p - alpha*D_TV(p,q). Deferring falls back to the + # EXACT law -- min(1, p/q) coin + residual -- so this branch is a + # superset of the exact rule with a draft-accept shortcut. + draft_q = draft_probs[depth_index] + if draft_q is None: + raise RuntimeError("non-greedy MTP requires draft distributions") + target_p_for_cache = target_distribution_batch.to_distribution( + depth_index + ) + _defer, _cas_tv = cascade_defer_decision( + target_p_for_cache, draft_q, alpha=_cascade_alpha + ) + cascade_positions += 1 + cascade_divergence_sum += _cas_tv + if not _defer: + # Draft good enough (pi = q): accept, no coin. + accepted_now = True + accept_prob = 1.0 + correction = draft_token + cascade_accepted += 1 + else: + # Defer (pi = p): exact speculative law, same as the exact + # branch below. + p = target_distribution_batch.probability(depth_index, draft_token) + q = ( + draft_q.probability(draft_token) + if isinstance(draft_q, SparseDistribution) + else float(draft_q[draft_token]) + ) + accept_prob = ( + 1.0 if q <= 0 and p > 0 else (0.0 if q <= 0 else min(1.0, p / q)) + ) + accepted_now = float(rng.random()) <= accept_prob + if accepted_now: + correction = draft_token + cascade_accepted += 1 + else: + correction = sample_from_distribution( + residual_distribution(target_p_for_cache, draft_q), rng + ) + cascade_resamples += 1 elif target_distribution_batch is not None: draft_q = draft_probs[depth_index] if draft_q is None: @@ -12362,27 +12488,56 @@ def emit_new_tokens() -> None: draft_q = draft_probs[depth_index] if draft_q is None: raise RuntimeError("non-greedy MTP requires draft distributions") - accept_prob = compute_acceptance_probability( - target_p, draft_q, draft_token - ) - if _bv is not None: - # Block verification: see the batched branch above. `_bv` - # is only ever built when every target row was already - # materialised, so it is None on the lazy path that just - # built `target_p` here. - accept_prob = _bv.accept_probability[depth_index] - accepted_now = float(rng.random()) <= accept_prob target_p_for_cache = target_p - if accepted_now: - correction = draft_token - elif _bv is not None: - correction = sample_from_distribution( - _bv.scaled_residual(depth_index), rng + if _cascade_active: + # Speculative-cascade acceptance, lazy/per-row target. Same + # law as the batched branch: deterministic Eq.(10) deferral, + # accept the draft when it is good enough, otherwise the + # exact min(1, p/q) coin + residual. + _defer, _cas_tv = cascade_defer_decision( + target_p, draft_q, alpha=_cascade_alpha ) + cascade_positions += 1 + cascade_divergence_sum += _cas_tv + if not _defer: + accepted_now = True + accept_prob = 1.0 + correction = draft_token + cascade_accepted += 1 + else: + accept_prob = compute_acceptance_probability( + target_p, draft_q, draft_token + ) + accepted_now = float(rng.random()) <= accept_prob + if accepted_now: + correction = draft_token + cascade_accepted += 1 + else: + correction = sample_from_distribution( + residual_distribution(target_p, draft_q), rng + ) + cascade_resamples += 1 else: - correction = sample_from_distribution( - residual_distribution(target_p, draft_q), rng + accept_prob = compute_acceptance_probability( + target_p, draft_q, draft_token ) + if _bv is not None: + # Block verification: see the batched branch above. `_bv` + # is only ever built when every target row was already + # materialised, so it is None on the lazy path that just + # built `target_p` here. + accept_prob = _bv.accept_probability[depth_index] + accepted_now = float(rng.random()) <= accept_prob + if accepted_now: + correction = draft_token + elif _bv is not None: + correction = sample_from_distribution( + _bv.scaled_residual(depth_index), rng + ) + else: + correction = sample_from_distribution( + residual_distribution(target_p, draft_q), rng + ) if not accepted_now and _env_truthy("MTPLX_DELTA_TELEMETRY"): # Tree Stage-0 pricing (2026-08-25): would a sibling branch # have caught this rejection? Record the rank of the @@ -13564,6 +13719,16 @@ def emit_new_tokens() -> None: accept_probability_sum_by_depth, drafted_by_depth, ), + cascade_accept_enabled=bool(_cascade_active), + cascade_alpha=float(_cascade_alpha) if _cascade_active else 0.0, + cascade_positions=int(cascade_positions), + cascade_accepted=int(cascade_accepted), + cascade_resamples=int(cascade_resamples), + cascade_mean_divergence=( + float(cascade_divergence_sum / cascade_positions) + if cascade_positions + else 0.0 + ), bonus_tokens=bonus_tokens, correction_tokens=correction_tokens, verify_calls=verify_calls, @@ -13949,6 +14114,22 @@ def generate_mtpa( stop_token_ids=stop_token_ids, max_tokens=max_tokens, ) + if _cascade_active: + _cas_denom = cascade_accepted + cascade_resamples + _cas_rate = (cascade_accepted / _cas_denom) if _cas_denom else 0.0 + _cas_cycles = max(1, verify_calls) + print( + "[cascade-accept] NOT distribution-exact; " + f"threshold={_cascade_alpha:.4g} alpha={_cascade_alpha:.4g} " + f"positions={cascade_positions} accepted={cascade_accepted} " + f"resamples={cascade_resamples} accept_rate={_cas_rate:.4f} " + f"mean_divergence={stats.cascade_mean_divergence:.4f} " + f"tokens_per_cycle={len(tokens) / _cas_cycles:.3f} " + f"accepted_by_depth={accepted_by_depth} " + f"generated={len(tokens)} verify_calls={verify_calls}", + file=sys.stderr, + flush=True, + ) return GenerationOutput( tokens=tokens, text=_decode(rt.tokenizer, _strip_terminal_stop(tokens, stop_token_ids)), diff --git a/mtplx/sampling.py b/mtplx/sampling.py index 44535357e..1cf26e1d3 100644 --- a/mtplx/sampling.py +++ b/mtplx/sampling.py @@ -288,6 +288,86 @@ def residual_distribution(target_p: Distribution, draft_q: Distribution) -> Dist return residual / total +def _peak_probability(distribution: Distribution) -> float: + """The largest single-token mass in ``distribution`` (``max_v P(v)``).""" + if isinstance(distribution, SparseDistribution): + probs = np.asarray(distribution.probs, dtype=np.float64) + else: + probs = np.asarray(distribution, dtype=np.float64) + probs = probs[np.isfinite(probs)] + if probs.size == 0: + return 0.0 + return float(probs.max()) + + +def total_variation(target_p: Distribution, draft_q: Distribution) -> float: + """``D_TV(p, q) = sum_v max(0, p(v) - q(v))`` over the scored top-k support. + + This is the total-variation divergence in the one-sided form the cascade + paper writes it (Narasimhan et al. 2024, arXiv:2405.19261, Eq. (8)); it is + the same unnormalized mass ``residual_distribution`` renormalizes. Taken + over the UNION of the two supports (the truncated target row and the draft + head's sparse rows), so tokens the draft scores but the target truncated + away, and vice versa, both count. Bounded in [0, 1]. + """ + if isinstance(target_p, SparseDistribution) or isinstance(draft_q, SparseDistribution): + if isinstance(target_p, SparseDistribution) and isinstance(draft_q, SparseDistribution): + token_ids = np.union1d(target_p.token_ids, draft_q.token_ids) + tv = 0.0 + for token in token_ids: + diff = target_p.probability(int(token)) - draft_q.probability(int(token)) + if diff > 0: + tv += diff + return float(tv) + dense_target = _as_dense(target_p) + dense_draft = _as_dense(draft_q) + else: + dense_target = np.asarray(target_p, dtype=np.float64) + dense_draft = np.asarray(draft_q, dtype=np.float64) + diff = dense_target - dense_draft + return float(np.sum(diff[np.isfinite(diff) & (diff > 0.0)])) + + +def cascade_defer_decision( + target_p: Distribution, + draft_q: Distribution, + *, + alpha: float, + tv_value: float | None = None, +) -> tuple[bool, float]: + """Speculative-cascade plug-in deferral rule. + + Narasimhan, Mreddy, Jitkrittum, Rawat, Kumar, "Faster Cascades via + Speculative Decoding", ICLR 2025 / arXiv:2405.19261 v2, Eq. (10) (the + plug-in approximation to the optimal deferral rule of Eq. (8)): + + r_OPT(x_ max_v q(v) < max_v p(v) - alpha * D_TV(p, q) + + ``r = 1`` DEFERS to the large (target) model; ``r = 0`` accepts the small + (draft) model's token. Returns ``(defer, tv)``. + + The draft is 'good enough' (do NOT defer) when its peak confidence + ``max_v q(v)`` is within ``alpha * D_TV(p, q)`` of the target's peak + ``max_v p(v)``. Not deferring means the effective speculative-cascade target + is ``pi = q`` (Sec. 4.1), so Algorithm 4's speculative-execution accept + probability ``min(1, pi(x)/q(x)) = 1`` and the draft token is accepted with + no coin. Deferring sets ``pi = p``, which is exactly the lossless + speculative-decoding law: accept with ``min(1, p(x)/q(x))`` and, on + rejection, resample the residual ``norm(max(0, p - q))``. + + ``alpha`` (the operator knob) is the Eq. (8) deferral cost. Higher ``alpha`` + lowers the RHS, so the accept band widens and FEWER positions defer (faster, + lossier); ``alpha = 0`` defers whenever the target is strictly more + confident than the draft. The decision is DETERMINISTIC (consumes no + uniform) -- the coin only appears on the deferred exact path. + """ + tv = total_variation(target_p, draft_q) if tv_value is None else float(tv_value) + peak_p = _peak_probability(target_p) + peak_q = _peak_probability(draft_q) + defer = peak_q < (peak_p - float(alpha) * tv) + return (defer, tv) + + def sample_from_distribution(probs: Distribution, rng: np.random.Generator | None = None) -> int: rng = rng or np.random.default_rng() if isinstance(probs, SparseDistribution): diff --git a/mtplx/server/openai.py b/mtplx/server/openai.py index 221b92759..6e7c56eeb 100644 --- a/mtplx/server/openai.py +++ b/mtplx/server/openai.py @@ -16495,6 +16495,52 @@ def _int_env(name: str) -> int | None: return None +def _cascade_acceptance_health_payload() -> dict[str, Any]: + """Resolved speculative-cascade acceptance lane state, for ``/health``. + + A SECOND lossy verify rule beside typical acceptance (Narasimhan et al. + 2024, arXiv:2405.19261, "Faster Cascades via Speculative Decoding", Eq. + (10)). OFF unless the operator sets MTPLX_FABLE_CASCADE_THRESHOLD (server + flag ``--cascade-threshold``); unset leaves the exact speculative law + unchanged. When on it is NOT distribution-exact: it accepts the draft token + whenever ``max_v q(v) >= max_v p(v) - alpha*D_TV(p,q)`` and otherwise defers + to the exact ``min(1, p/q)`` coin + residual. Mutually exclusive with the + typical lane (both set is a fail-loud misconfiguration). See + docs/perf/pr478-cascade-acceptance.md. + """ + raw = os.environ.get("MTPLX_FABLE_CASCADE_THRESHOLD") + enabled = raw is not None and str(raw).strip() != "" + alpha: float | None = None + if enabled: + try: + alpha = float(raw) + except ValueError: + alpha = None + enabled = False + typical_on = False + tv = os.environ.get("MTPLX_FABLE_TYPICAL_THRESHOLD") + try: + typical_on = tv is not None and float(tv) > 0.0 + except ValueError: + typical_on = False + return { + "enabled": enabled, + "alpha": alpha, + "threshold": alpha, + "rule": "defer iff max_q < max_p - alpha*D_TV(p,q); else accept draft", + "divergence": "D_TV(p,q) = sum_v max(0, p(v)-q(v)) over scored top-k", + "distribution_exact": not enabled, + "mutually_exclusive_with_typical": True, + "conflict": bool(enabled and typical_on), + "citation": "arXiv:2405.19261 Eq. (10)", + "note": ( + "OFF unless --cascade-threshold (MTPLX_FABLE_CASCADE_THRESHOLD) is " + "set; when on it is NOT distribution-exact and engages only at " + "temperature > 0; mutually exclusive with --typical-threshold" + ), + } + + def _startup_health_payload(state: "ServerState") -> dict[str, Any]: chat_template_report = getattr(state, "chat_template_report", {}) or {} tool_prompt_mode = _tool_prompt_mode_from_args(state.args) @@ -29074,6 +29120,7 @@ def health() -> dict[str, Any]: "actual_ramp_latency_s" ), "startup": _startup_health_payload(state), + "cascade_acceptance": _cascade_acceptance_health_payload(), "thermal": _thermal_health_payload( fan_mode=fan_mode, smart_status=smart_status, @@ -36319,6 +36366,24 @@ def parse_args(argv: list[str] | None = None) -> argparse.Namespace: default=True, help="Load and inject the native MTP sidecar. Disable only for stock AR diagnostics.", ) + parser.add_argument( + "--cascade-threshold", + type=float, + default=None, + metavar="ALPHA", + help=( + "Enable speculative-cascade acceptance (arXiv:2405.19261 Eq. (10)) " + "at this deferral cost alpha: accept the draft token when " + "max_q >= max_p - alpha*D_TV(p,q), else defer to the exact " + "min(1,p/q) coin + residual. Unset = OFF = exact speculative " + "sampling (the default); any set value (including 0) turns it on. " + "Higher alpha widens the accept band (fewer defers). NOT " + "distribution-exact; engages only at temperature > 0. Mutually " + "exclusive with --typical-threshold (setting both is an error). " + "Environment: MTPLX_FABLE_CASCADE_THRESHOLD, which this flag " + "overrides. See docs/perf/pr478-cascade-acceptance.md." + ), + ) parser.add_argument( "--ngram-prewarm", metavar="auto|all|off|GiB", @@ -36853,6 +36918,27 @@ def parse_args(argv: list[str] | None = None) -> argparse.Namespace: "off" if args.strip_assistant_reasoning_history else args.preserve_thinking ) args.strip_assistant_reasoning_history = not _preserve_thinking_effective(args) + if getattr(args, "cascade_threshold", None) is not None: + # Same flag-beats-env contract; generation.py reads the env per request. + # Unlike the typical delta, any set value (including 0) turns the lane + # on, so only pass --cascade-threshold to enable it. + os.environ["MTPLX_FABLE_CASCADE_THRESHOLD"] = str(float(args.cascade_threshold)) + # Fail loud on the mutually exclusive lossy verify rules, whether they were + # set by flag (stamped just above) or already present in the environment. + _typ_env = os.environ.get("MTPLX_FABLE_TYPICAL_THRESHOLD") + _cas_env = os.environ.get("MTPLX_FABLE_CASCADE_THRESHOLD") + _typ_on = False + try: + _typ_on = _typ_env is not None and float(_typ_env) > 0.0 + except ValueError: + _typ_on = False + _cas_on = _cas_env is not None and str(_cas_env).strip() != "" + if _typ_on and _cas_on: + raise SystemExit( + "error: --typical-threshold and --cascade-threshold are mutually " + "exclusive lossy verify rules; set at most one " + f"(typical={_typ_env!r}, cascade={_cas_env!r})." + ) return args diff --git a/tests/test_cascade_acceptance.py b/tests/test_cascade_acceptance.py new file mode 100644 index 000000000..c569ed286 --- /dev/null +++ b/tests/test_cascade_acceptance.py @@ -0,0 +1,218 @@ +"""CPU tests for speculative-cascade acceptance (a second lossy verify rule). + +Rule implemented: Narasimhan, Mreddy, Jitkrittum, Rawat, Kumar, "Faster +Cascades via Speculative Decoding" (arXiv:2405.19261 v2), Section 4.3 +Equation (10), the plug-in approximation to the optimal deferral rule of +Equation (8): + + r_OPT(x_ max_v q(v) < max_v p(v) - alpha * D_TV(p, q) + +r = 1 DEFERS to the large/target model p; r = 0 accepts the small/draft model q. +D_TV(p, q) = sum_v max(0, p(v) - q(v)) over the scored top-k support. On defer +the code takes the exact speculative law (min(1, p/q) coin + residual), so the +cascade path is a superset of the exact rule with a draft-accept shortcut. + +Note on the divergence sign (Lemma 3): alpha * D_TV is SUBTRACTED, so a larger +disagreement LOWERS the bar and defers LESS. The intuitive "reject when the +draft diverges" holds for a draft that also loses peak confidence (the realistic +diverging MTP head); a draft that stays confident on a wrong token is accepted, +which is the paper's documented behavior and is pinned explicitly below. +""" +from __future__ import annotations + +import importlib +import os + +import numpy as np +import pytest + +from mtplx.sampling import ( + SparseDistribution, + acceptance_probability, + cascade_defer_decision, + residual_distribution, + total_variation, +) + +VOCAB = 128 + + +def sp(pairs) -> SparseDistribution: + ids = np.array([t for t, _ in pairs], dtype=np.int64) + probs = np.array([p for _, p in pairs], dtype=np.float64) + return SparseDistribution(ids, probs, VOCAB) + + +# --------------------------------------------------------------------------- +# The rule (Eq. 10). +# --------------------------------------------------------------------------- + + +def test_total_variation_matches_paper_definition(): + p = sp([(0, 0.8), (1, 0.2)]) + q = sp([(0, 0.5), (1, 0.5)]) + # sum_v max(0, p-q) = max(0, 0.3) + max(0, -0.3) = 0.3 + assert total_variation(p, q) == pytest.approx(0.3) + # union support: token the draft scores but the target truncated away counts + p2 = sp([(0, 1.0)]) + q2 = sp([(0, 0.5), (9, 0.5)]) + assert total_variation(p2, q2) == pytest.approx(0.5) + + +def test_accepts_when_distributions_agree_even_for_an_atypical_token(): + # q == p: max_q == max_p and D_TV == 0, so the bar max_p - alpha*0 == max_p + # is not strictly above max_q -> do NOT defer -> accept the draft. The + # decision uses only max_q/max_p/D_TV, so it is token-independent: an + # atypical proposed token (low mass) is accepted just the same, which is the + # cascade advantage over typical acceptance (which would reject it). + p = sp([(0, 0.9), (1, 0.06), (7, 0.04)]) + q = sp([(0, 0.9), (1, 0.06), (7, 0.04)]) + for alpha in (0.0, 0.1, 0.5, 2.0): + defer, tv = cascade_defer_decision(p, q, alpha=alpha) + assert defer is False, alpha + assert tv == pytest.approx(0.0) + + +def test_defers_when_draft_diverges_and_loses_confidence(): + # Confident target (peak 0.9 on token 0); the draft has diverged to a spread + # distribution on other tokens (peak 0.3, high D_TV). With a modest alpha the + # subtracted penalty is small, so max_q (0.3) < max_p - alpha*D_TV -> defer. + p = sp([(0, 0.9), (1, 0.1)]) + q = sp([(3, 0.3), (4, 0.25), (5, 0.25), (6, 0.2)]) + defer, tv = cascade_defer_decision(p, q, alpha=0.1) + assert defer is True + assert tv > 0.5 + + +def test_alpha_zero_defers_iff_target_strictly_more_confident(): + # alpha=0 removes the divergence penalty: pure peak comparison. + p = sp([(0, 0.6), (1, 0.4)]) + q_less = sp([(0, 0.5), (1, 0.5)]) # max_q 0.5 < max_p 0.6 -> defer + q_more = sp([(0, 0.7), (1, 0.3)]) # max_q 0.7 > max_p 0.6 -> accept + assert cascade_defer_decision(p, q_less, alpha=0.0)[0] is True + assert cascade_defer_decision(p, q_more, alpha=0.0)[0] is False + + +def test_higher_alpha_defers_less(): + # Borderline case: raising alpha widens the accept band (monotone), so a + # position that defers at low alpha stops deferring at high alpha. + p = sp([(0, 0.7), (1, 0.3)]) + q = sp([(2, 0.5), (3, 0.5)]) # max_q 0.5 < max_p 0.7, D_TV = 0.7 + assert cascade_defer_decision(p, q, alpha=0.0)[0] is True # 0.5 < 0.7 + assert cascade_defer_decision(p, q, alpha=1.0)[0] is False # 0.5 < 0.7-0.7=0.0 -> False + + +def test_paper_counterintuitive_confident_disagreement_is_accepted(): + # Both models confident (peak 0.95) but on OPPOSITE tokens: D_TV = 0.9. Per + # Eq. (10) with alpha=0.5 the bar is 0.95 - 0.5*0.9 = 0.5, and max_q 0.95 is + # not below it, so the draft is ACCEPTED. This is the paper's documented + # "defer less when disagreement is large" behavior (Lemma 3); pinned so a + # future change to the sign is caught. + p = sp([(0, 0.95), (1, 0.05)]) + q = sp([(1, 0.95), (0, 0.05)]) + defer, tv = cascade_defer_decision(p, q, alpha=0.5) + assert tv == pytest.approx(0.9) + assert defer is False + + +# --------------------------------------------------------------------------- +# Superset of the exact rule: the deferred path IS the exact law. +# --------------------------------------------------------------------------- + + +def test_deferred_path_equals_exact_speculative_law(): + # When the rule defers, the caller runs min(1, p/q) + residual(p, q) -- the + # exact Leviathan-Chen law. Pin that the primitives the cascade branch uses + # on defer are exactly the exact-rule primitives (so cascade is a strict + # superset: an exact accept/residual, gated behind the deferral test). + p = sp([(0, 0.6), (1, 0.3), (2, 0.1)]) + q = sp([(0, 0.2), (1, 0.5), (2, 0.3)]) + token = 1 + # exact accept probability and residual, computed the way both the exact + # branch and the cascade-defer branch compute them. + assert acceptance_probability(p, q, token) == pytest.approx(min(1.0, 0.3 / 0.5)) + resid = residual_distribution(p, q) + # residual mass is norm(max(0, p-q)); token 1 has p 0.0 + + +# --------------------------------------------------------------------------- +# Arming: read at use (served order), default off, mutual exclusion. +# --------------------------------------------------------------------------- + + +def _clear(): + os.environ.pop("MTPLX_FABLE_CASCADE_THRESHOLD", None) + os.environ.pop("MTPLX_FABLE_TYPICAL_THRESHOLD", None) + + +@pytest.fixture(autouse=True) +def _clean_env(): + _clear() + try: + yield + finally: + _clear() + + +def test_default_off_and_any_value_including_zero_turns_on(): + from mtplx.generation import _cascade_accept_alpha, _cascade_accept_enabled + + assert _cascade_accept_alpha() is None + assert _cascade_accept_enabled() is False + os.environ["MTPLX_FABLE_CASCADE_THRESHOLD"] = "0" + assert _cascade_accept_alpha() == 0.0 + assert _cascade_accept_enabled() is True # explicit 0 is ON, not the off switch + os.environ["MTPLX_FABLE_CASCADE_THRESHOLD"] = "0.5" + assert _cascade_accept_alpha() == pytest.approx(0.5) + + +def test_served_order_reader_read_at_use(): + # Reproduce the served order: import the generation and server modules + # FIRST (before any auto-arm/setdefault), then set the env. A reader frozen + # at import would miss it; read-at-use sees it. + gen = importlib.import_module("mtplx.generation") + importlib.import_module("mtplx.server.openai") + assert gen._cascade_accept_alpha() is None + os.environ["MTPLX_FABLE_CASCADE_THRESHOLD"] = "0.25" + assert gen._cascade_accept_alpha() == pytest.approx(0.25) + + +def test_mutual_exclusion_fails_loud(): + from mtplx.generation import _assert_lossy_verify_rules_exclusive + + # Either alone is fine. + os.environ["MTPLX_FABLE_CASCADE_THRESHOLD"] = "0.3" + _assert_lossy_verify_rules_exclusive() + _clear() + os.environ["MTPLX_FABLE_TYPICAL_THRESHOLD"] = "0.09" + _assert_lossy_verify_rules_exclusive() + # Both set -> fail loud. + os.environ["MTPLX_FABLE_CASCADE_THRESHOLD"] = "0.3" + with pytest.raises(ValueError, match="mutually exclusive"): + _assert_lossy_verify_rules_exclusive() + + +def test_typical_zero_does_not_conflict_with_cascade(): + # Typical is OFF at delta 0, so cascade + typical=0 is not a conflict. + from mtplx.generation import _assert_lossy_verify_rules_exclusive + + os.environ["MTPLX_FABLE_TYPICAL_THRESHOLD"] = "0" + os.environ["MTPLX_FABLE_CASCADE_THRESHOLD"] = "0.3" + _assert_lossy_verify_rules_exclusive() # must not raise + + +def test_exact_mode_when_off_takes_neither_lossy_branch(): + # Flag off (env unset): cascade disabled, so the verify path is the exact + # speculative law unchanged. This is the cascade-only PEER of #478; the + # typical lane's code is not present on this branch (the two are alternative + # modes), so the exact-off contract is just "cascade off" plus the inert + # mutual-exclusion guard. + from mtplx.generation import ( + _assert_lossy_verify_rules_exclusive, + _cascade_accept_enabled, + ) + + assert _cascade_accept_enabled() is False + _assert_lossy_verify_rules_exclusive() # inert here, must not raise diff --git a/tests/test_cascade_threshold_cli_health_cpu.py b/tests/test_cascade_threshold_cli_health_cpu.py new file mode 100644 index 000000000..84ae41fa6 --- /dev/null +++ b/tests/test_cascade_threshold_cli_health_cpu.py @@ -0,0 +1,94 @@ +"""--cascade-threshold CLI flag + the /health cascade_acceptance entry. CPU-only. + +mtplx/server/openai.py imports MLX, so the /health payload helper is compiled out +of the shipped source (the same trick test_typical_threshold_cli_health_cpu.py +uses) instead of importing the module, and the CLI wiring is asserted against the +source text. +""" +from __future__ import annotations + +import argparse +import ast +import os +from pathlib import Path +from typing import Any + +import pytest + +ROOT = Path(__file__).resolve().parents[1] +SERVER_TEXT = (ROOT / "mtplx" / "server" / "openai.py").read_text("utf-8") + + +def _compile_function(source: str, name: str, namespace: dict | None = None): + node = next( + n + for n in ast.parse(source).body + if isinstance(n, ast.FunctionDef) and n.name == name + ) + module = ast.Module(body=[node], type_ignores=[]) + ast.fix_missing_locations(module) + scope: dict = {"Any": Any, "os": os, "argparse": argparse} + scope.update(namespace or {}) + exec(compile(module, f"<{name}>", "exec"), scope) + return scope[name] + + +@pytest.fixture(autouse=True) +def _clean(monkeypatch): + for key in ("MTPLX_FABLE_CASCADE_THRESHOLD", "MTPLX_FABLE_TYPICAL_THRESHOLD"): + monkeypatch.delenv(key, raising=False) + + +def _payload(): + return _compile_function(SERVER_TEXT, "_cascade_acceptance_health_payload")() + + +def test_health_reports_off_by_default(): + p = _payload() + assert p["enabled"] is False + assert p["alpha"] is None + assert p["distribution_exact"] is True + assert p["conflict"] is False + assert p["citation"] == "arXiv:2405.19261 Eq. (10)" + + +def test_health_reports_on_when_set_including_zero(monkeypatch): + monkeypatch.setenv("MTPLX_FABLE_CASCADE_THRESHOLD", "0") + p = _payload() + assert p["enabled"] is True # explicit 0 is ON + assert p["alpha"] == 0.0 + assert p["distribution_exact"] is False + monkeypatch.setenv("MTPLX_FABLE_CASCADE_THRESHOLD", "0.5") + p = _payload() + assert p["enabled"] is True + assert p["alpha"] == 0.5 + + +def test_health_flags_conflict_when_both_lanes_set(monkeypatch): + monkeypatch.setenv("MTPLX_FABLE_CASCADE_THRESHOLD", "0.3") + monkeypatch.setenv("MTPLX_FABLE_TYPICAL_THRESHOLD", "0.09") + p = _payload() + assert p["enabled"] is True + assert p["conflict"] is True + + +def test_health_bad_value_is_not_fatal(monkeypatch): + monkeypatch.setenv("MTPLX_FABLE_CASCADE_THRESHOLD", "sometimes") + p = _payload() # must not raise + assert p["enabled"] is False + assert p["alpha"] is None + + +def test_health_entry_is_wired_into_the_route(): + assert '"cascade_acceptance": _cascade_acceptance_health_payload(),' in SERVER_TEXT + + +def test_cli_flag_defined_and_stamps_env(): + assert '"--cascade-threshold",' in SERVER_TEXT + assert 'os.environ["MTPLX_FABLE_CASCADE_THRESHOLD"] = str(float(args.cascade_threshold))' in SERVER_TEXT + + +def test_cli_mutual_exclusion_is_wired(): + # The flag-apply raises SystemExit when both lossy rules are set. + assert "mutually exclusive lossy verify rules" in SERVER_TEXT + assert "SystemExit(" in SERVER_TEXT diff --git a/tests/test_qwen4_block_verify.py b/tests/test_qwen4_block_verify.py index 85cdbab9e..e826212b6 100644 --- a/tests/test_qwen4_block_verify.py +++ b/tests/test_qwen4_block_verify.py @@ -322,9 +322,16 @@ def test_the_shipped_law_survives_verbatim_in_the_accept_loop(): assert "residual_distribution(\n" in loop assert "else target_distribution_batch.to_distribution(depth_index)," in loop # The accept coin is still drawn once per depth and compared with `<=`, - # so arming block verification cannot shift the PCG64 stream. - assert loop.count("accepted_now = float(rng.random()) <= accept_prob") == 2 - assert loop.count("rng.random()") == 2 + # so arming block verification cannot shift the PCG64 stream. The exact law + # appears in four places now: the two shipped exact branches (batched + + # lazy) plus the two speculative-cascade DEFER paths (batched + lazy), which + # reuse the identical coin verbatim. The cascade paths are gated behind + # `_cascade_active` (MTPLX_FABLE_CASCADE_THRESHOLD, off by default and + # mutually exclusive with block verification's own lane), so with the + # cascade lane off the coin still fires exactly once per depth and the + # block-verification RNG guarantee is unchanged. + assert loop.count("accepted_now = float(rng.random()) <= accept_prob") == 4 + assert loop.count("rng.random()") == 4 def test_every_block_verification_read_is_guarded(): From de056964307faead9a866c7833f018685f57658d Mon Sep 17 00:00:00 2001 From: davidtai Date: Tue, 8 Sep 2026 09:09:10 -0500 Subject: [PATCH 09/30] docs(cascade): canonical perf note + alpha-grid recommendation Rename the speculative-cascade perf note to docs/perf/qwen38-cascade-acceptance.md (matching the qwen38-* perf-note naming) and repoint the flag help and /health note references to it. Add a "Recommended alpha grid for the 16K sweep" section: the paper varies alpha continuously (Section 6, Figure 2) with no fixed grid, so the grid is set from Equation (10)'s structure and this model's measured scale (exact MTP acceptance 0.43-0.47 -> D_TV ~= 1 - accept ~= 0.55; target-row entropy 0.83-1.81 nats -> target peak ~0.35-0.65). Grid {0.0, 0.5, 1.0, 2.0} from maximal deferral to near all-accept; recommended HumanEval operating point alpha = 0.5. --- docs/perf/pr478-cascade-acceptance.md | 113 ------------------ docs/perf/qwen38-cascade-acceptance.md | 156 +++++++++++++++++++++++++ mtplx/server/openai.py | 4 +- 3 files changed, 158 insertions(+), 115 deletions(-) delete mode 100644 docs/perf/pr478-cascade-acceptance.md create mode 100644 docs/perf/qwen38-cascade-acceptance.md diff --git a/docs/perf/pr478-cascade-acceptance.md b/docs/perf/pr478-cascade-acceptance.md deleted file mode 100644 index bca278742..000000000 --- a/docs/perf/pr478-cascade-acceptance.md +++ /dev/null @@ -1,113 +0,0 @@ -# Speculative-cascade acceptance for Qwen3.8 Flash-Next - -Speculative-cascade acceptance is a second opt-in, lossy decode-verify rule -beside typical acceptance. It decides per draft position whether the draft token -is good enough to keep, or whether to defer to the exact target law. It is OFF by -default, mutually exclusive with typical acceptance, and NOT distribution-exact. - -Citation: Narasimhan, Mreddy, Jitkrittum, Rawat, Kumar, "Faster Cascades via -Speculative Decoding", ICLR 2025 / arXiv:2405.19261 v2. This implements the -plug-in deferral rule of Section 4.3, Equation (10), with the speculative -execution of Algorithm 4. - -Terms: - -- p: the target distribution at a position (the truncated top-p/top-k target - row already materialized for the exact rule). -- q: the draft distribution at a position (the native MTP head's scored rows; - see "What q is" below). -- D_TV(p, q) = sum_v max(0, p(v) - q(v)) over the scored top-k support. -- defer (r = 1): use the large/target model at this position; accept (r = 0): - keep the small/draft model's token. -- alpha: the deferral cost (the operator knob). - -## Problem - -The exact speculative law commits the longest draft prefix the target would have -produced with the same coins, and no more. When the draft head is close to the -target it still spends a rejection whenever the exact coin `min(1, p/q)` happens -to fall, capping tokens per cycle. Typical acceptance loosens this by keeping -tokens that are "typical" under the target row, but it is blind to how well the -DRAFT itself matched the target: it can reject a token the draft was confidently -and correctly proposing, and it resamples from the target row rather than the -residual. - -## Change - -A speculative-cascade verify rule that, at each draft position with target p and -draft q, applies Equation (10): - - defer (r = 1) <=> max_v q(v) < max_v p(v) - alpha * D_TV(p, q) - -If it does NOT defer, the draft is "good enough" (the speculative-cascade target -is pi = q, so Algorithm 4's accept probability min(1, pi/q) = 1): accept the -draft token with no coin. If it DOES defer, pi = p and the code runs the exact -speculative law unchanged -- accept with `min(1, p(x_t)/q(x_t))`, and on a coin -rejection resample the residual `norm(max(0, p - q))`. So the deferral test uses -the total variation over the scored top-k plus, on the deferred path, the token's -own target probability p(x_t), and the cascade path is a strict superset of the -exact rule with a draft-accept shortcut. - -The decision is deterministic (it consumes no uniform); the coin only appears on -the deferred exact path, so with the lane off the RNG stream is byte-identical. - -What q is: on the served Turbo path the draft is the native MTP head. Its scored -rows reach the verify loop as `draft_probs[depth_index]` -- a `SparseDistribution` -over the head's scored vocabulary, which for this pack is the FR-Spec -frequency-ranked subset (the same q the exact rule already uses in `min(1, p/q)` -and residual). The cascade rule reads q's peak (`max_v q`) and q's mass on the -scored top-k for D_TV from exactly that object; it introduces no new draft -distribution and needs no extra draft forward. q is required (the lane raises if -a non-greedy MTP position has no draft distribution), so the lane is a -temperature > 0, MTP-on rule, like typical acceptance. - -Sign note (paper Lemma 3): alpha * D_TV is SUBTRACTED, so a larger disagreement -LOWERS the bar and defers LESS. A diverging draft that also loses peak confidence -(the realistic case) defers; a draft that stays confident on a wrong token is -accepted. This is the paper's documented behavior and is pinned in the tests. - -## Effect - -Higher acceptance when the draft distribution agrees with the target, including -on atypical tokens that typical acceptance would reject, at the cost of exactness -(the accepted draft is sampled from q, not p). It is a speed/quality dial, gated -on task-quality evals like typical acceptance, never on distribution-exactness. - -## Exactness - -NOT distribution-exact when on. The accepted-draft (non-defer) positions emit q's -token, which differs from the target law. The deferred positions ARE exact -(min(1, p/q) coin + residual). With the lane OFF (knob unset) the verify path is -byte-identical to the exact rule: the cascade branches are gated behind -`_cascade_active`, which is false, so neither the deterministic test nor any -extra coin runs. - -## Files - -- `mtplx/sampling.py` -- `cascade_defer_decision`, `total_variation`, - `_peak_probability`. -- `mtplx/generation.py` -- `_cascade_accept_alpha` / `_cascade_accept_enabled` - (env at use), `_assert_lossy_verify_rules_exclusive`, the two verify-loop - decision branches (batched + lazy target), the `VerifyStats` cascade fields, - and the `[cascade-accept]` verdict line. -- `mtplx/server/openai.py` -- `--cascade-threshold`, the env stamp + fail-loud - mutual exclusion with `--typical-threshold`, and the `/health` - `cascade_acceptance` install report. -- `tests/test_cascade_acceptance.py`, `tests/test_cascade_threshold_cli_health_cpu.py`. - -## Switch - -- `--cascade-threshold ALPHA` (env `MTPLX_FABLE_CASCADE_THRESHOLD`): the - deferral cost alpha. UNSET (or blank) = OFF = exact speculative sampling (the - default). Any set value -- including an explicit 0 -- turns the lane on, so the - off switch is UNSETTING the key, not setting it to 0 (alpha = 0 still defers - whenever the target is strictly more confident than the draft). Higher alpha - widens the accept band, so fewer positions defer. Resolved at use, per request. -- Mutually exclusive with `--typical-threshold` (`MTPLX_FABLE_TYPICAL_THRESHOLD` - > 0). Setting both fails loud at serve startup (SystemExit) and in the verify - setup (ValueError). -- Observability: `/health` -> `cascade_acceptance` (`enabled`, `alpha`, the rule - and divergence definitions, `conflict`); and a per-request verdict line - `[cascade-accept] NOT distribution-exact; threshold=A alpha=A positions=N - accepted=N resamples=N accept_rate=R mean_divergence=D ...` in the same format - as `[typical-accept]`. diff --git a/docs/perf/qwen38-cascade-acceptance.md b/docs/perf/qwen38-cascade-acceptance.md new file mode 100644 index 000000000..1a48b8da4 --- /dev/null +++ b/docs/perf/qwen38-cascade-acceptance.md @@ -0,0 +1,156 @@ +# Speculative-cascade acceptance for Qwen3.8 Flash-Next + +Speculative-cascade acceptance is a second opt-in, lossy decode-verify rule +beside typical acceptance. Per draft position it decides whether the draft token +is good enough to keep, or whether to defer to the exact target law. It is OFF by +default, mutually exclusive with typical acceptance, and NOT distribution-exact. + +Citation: Narasimhan, Mreddy, Jitkrittum, Rawat, Kumar, "Faster Cascades via +Speculative Decoding," ICLR 2025 / arXiv:2405.19261 v2. This implements the +plug-in deferral rule of Section 4.3, Equation (10) (the plug-in approximation +to the optimal rule of Equation (8)), executed with the speculative decoding of +Algorithm 4. + +Terms: + +- p: the target distribution at a position (the truncated top-p / top-k target + row already materialized for the exact rule). +- q: the draft distribution at a position (the native MTP head's scored rows). +- D_TV(p, q) = sum_v max(0, p(v) - q(v)) over the scored top-k support. +- defer (r = 1): use the large / target model here; accept (r = 0): keep the + small / draft model's token. +- alpha: the deferral cost, the single operator knob. + +## Problem + +The exact speculative law commits the longest draft prefix the target would have +produced under the same coins, and no more. Even when the draft head closely +tracks the target it still spends a rejection whenever the coin `min(1, p/q)` +falls, which caps tokens per cycle. Typical acceptance loosens this with a +per-row typicality floor, but it is blind to how well the draft itself matched +the target: it can reject a token the draft was confidently and correctly +proposing, and it resamples from the target row rather than the residual. + +## Change + +A speculative-cascade verify rule that, at each draft position with target p and +draft q, applies Equation (10): + + defer (r = 1) <=> max_v q(v) < max_v p(v) - alpha * D_TV(p, q) + +If it does NOT defer, the draft is "good enough" (the speculative-cascade target +is pi = q, so Algorithm 4's accept probability min(1, pi/q) = 1): accept the +draft token with no coin. If it DOES defer, pi = p and the code runs the exact +speculative law unchanged, accept with `min(1, p(x_t)/q(x_t))` and, on a coin +rejection, resample the residual `norm(max(0, p - q))`. So the deferral test uses +the total variation over the scored top-k, and the deferred path uses the token's +own target probability p(x_t); the cascade path is a strict superset of the exact +rule with a draft-accept shortcut. The decision is deterministic (it consumes no +uniform); the coin only appears on the deferred path, so with the lane off the +RNG stream and verify path are byte-identical. + +What q is. On the served Turbo path the draft is the native MTP head. Its scored +rows reach the verify loop as `draft_probs[depth_index]`, a `SparseDistribution` +over the head's scored vocabulary, which for this pack is the FR-Spec +frequency-ranked subset (the same q the exact rule already uses in `min(1, p/q)` +and the residual). The rule reads q's peak `max_v q` and q's mass on the scored +top-k for D_TV from that object; it adds no draft distribution and no draft +forward. q is required, so the lane is a temperature > 0, MTP-on rule, like +typical acceptance. + +The Lemma 3 note, in plain words. In Equation (10) the cost term `alpha * D_TV` +is SUBTRACTED from the target's confidence. So more disagreement between the +draft and the target LOWERS the bar for accepting the draft, which means the rule +defers LESS when they disagree, not more. The paper's reason (their Lemma 3) is +that a large disagreement makes the verification step itself expensive, so the +optimal rule only pays that cost when the target is clearly better; it accepts a +draft that is confident, even on a token the target would not have picked. The +intuitive "reject when the draft diverges" behaviour still holds for the common +case, a diverging draft that has also lost its peak confidence, which defers; a +draft that stays confident on a wrong token is accepted. This is the paper's rule +as written; the tests pin it explicitly so a future sign change is caught. + +## Effect + +Higher acceptance when the draft distribution agrees with the target, including on +atypical tokens that typical acceptance would reject, at the cost of exactness +(the accepted draft token is sampled from q, not p). It is a speed / quality +dial, judged on task-quality evals like typical acceptance, never on +distribution-exactness. + +## Exactness + +NOT distribution-exact when on. The accepted-draft (non-defer) positions emit q's +token, which differs from the target law. The deferred positions ARE exact +(min(1, p/q) coin + residual). With the lane OFF (knob unset) the verify path is +byte-identical to the exact rule: the cascade branches are gated behind +`_cascade_active`, which is false, so neither the deterministic test nor any coin +runs. + +## Files + +- `mtplx/sampling.py`: `cascade_defer_decision`, `total_variation`, + `_peak_probability`. +- `mtplx/generation.py`: `_cascade_accept_alpha` / `_cascade_accept_enabled` + (env at use), `_assert_lossy_verify_rules_exclusive`, the two verify-loop + decision branches (batched + lazy target), the `VerifyStats` cascade fields, + and the `[cascade-accept]` verdict line. +- `mtplx/server/openai.py`: `--cascade-threshold`, the env stamp + fail-loud + mutual exclusion with `--typical-threshold`, and the `/health` + `cascade_acceptance` install report. +- `tests/test_cascade_acceptance.py`, `tests/test_cascade_threshold_cli_health_cpu.py`. + +## Switch + +- `--cascade-threshold ALPHA` (env `MTPLX_FABLE_CASCADE_THRESHOLD`): the deferral + cost alpha. UNSET (or blank) = OFF = exact speculative sampling (the default). + Any set value, including an explicit 0, turns the lane on, so the off switch is + UNSETTING the key, not setting it to 0 (alpha = 0 still defers whenever the + target is strictly more confident than the draft). Higher alpha widens the + accept band, so fewer positions defer. Resolved at use, per request. +- Mutually exclusive with `--typical-threshold` (`MTPLX_FABLE_TYPICAL_THRESHOLD` + > 0). Setting both fails loud at serve startup (SystemExit) and in the verify + setup (ValueError). +- Observability: `/health` -> `cascade_acceptance` (`enabled`, `alpha`, the rule + and divergence definitions, `conflict`); and a per-request verdict line + `[cascade-accept] NOT distribution-exact; threshold=A alpha=A positions=N + accepted=N resamples=N accept_rate=R mean_divergence=D ...` in the same format + as `[typical-accept]`. + +## Recommended alpha grid for the 16K sweep + +The paper varies alpha continuously as a "lenience parameter" to trace a +quality-versus-latency Pareto curve (Section 6, Figure 2); it fixes no grid, and +reports temperatures T in {0, 0.1, 0.5, 1.0} and block sizes gamma in {3, 5, 7} +(Appendix E.1). So the grid is set from Equation (10)'s structure and this +model's measured scale. + +Scale on this model, from the saved #478 receipts and verdict lines at 16,384 +tokens, temperature 1, depth 3: + +- Exact MTP acceptance rate is 0.43 to 0.47. The expected exact speculative + acceptance rate equals the overlap `sum_v min(p, q) = 1 - D_TV(p, q)`, so + `D_TV ~= 0.53 to 0.57` per position on average. +- Target-row entropy on the `[typical-accept]` lines is 0.83 to 1.81 nats, so the + target peak `max_v p` sits around 0.35 to 0.65. + +The deferral boundary is `max_v p - max_v q = alpha * D_TV`, so the useful alpha +range is set by the plausible target-versus-draft peak gap (at most ~max_v p, +about 0.6) divided by D_TV (~0.55): alpha up to about 1 already reaches the point +where the cost term dominates any confidence edge, and alpha ~2 defers almost +never. The four-value grid, from most deferral (least lossy, closest to the exact +rule) to near all-accept (fastest, lossiest): + +| alpha | alpha * D_TV (D_TV ~= 0.55) | expected behaviour | +| --- | --- | --- | +| 0.0 | 0.00 | defer whenever the target is strictly more confident than the draft; maximal deferral | +| 0.5 | ~0.28 | defer only when the target's peak leads the draft's by more than ~0.28; moderate | +| 1.0 | ~0.55 | defer only on a large target-confidence lead; light | +| 2.0 | ~1.10 | defer essentially never; near all-accept | + +Single alpha for the HumanEval cell: **alpha = 0.5**. It is the moderate point, +accepting the draft where the draft and target peaks are close and deferring +where the target clearly leads, which is the "gain speed while holding quality" +operating point the sweep should confirm. If HumanEval drops, fall back toward +alpha = 0 (more deferral); if quality holds, push toward alpha = 1.0 for more +speed. diff --git a/mtplx/server/openai.py b/mtplx/server/openai.py index 6e7c56eeb..77fa2eeeb 100644 --- a/mtplx/server/openai.py +++ b/mtplx/server/openai.py @@ -16506,7 +16506,7 @@ def _cascade_acceptance_health_payload() -> dict[str, Any]: whenever ``max_v q(v) >= max_v p(v) - alpha*D_TV(p,q)`` and otherwise defers to the exact ``min(1, p/q)`` coin + residual. Mutually exclusive with the typical lane (both set is a fail-loud misconfiguration). See - docs/perf/pr478-cascade-acceptance.md. + docs/perf/qwen38-cascade-acceptance.md. """ raw = os.environ.get("MTPLX_FABLE_CASCADE_THRESHOLD") enabled = raw is not None and str(raw).strip() != "" @@ -36381,7 +36381,7 @@ def parse_args(argv: list[str] | None = None) -> argparse.Namespace: "distribution-exact; engages only at temperature > 0. Mutually " "exclusive with --typical-threshold (setting both is an error). " "Environment: MTPLX_FABLE_CASCADE_THRESHOLD, which this flag " - "overrides. See docs/perf/pr478-cascade-acceptance.md." + "overrides. See docs/perf/qwen38-cascade-acceptance.md." ), ) parser.add_argument( From 94369654c17e5251730c35a67a665bc1db6ab192 Mon Sep 17 00:00:00 2001 From: davidtai Date: Tue, 8 Sep 2026 10:41:52 -0500 Subject: [PATCH 10/30] docs(cascade): decode-vs-alpha and acceptance-vs-alpha charts for the arm-G sweep --- .../charts/cascade_accept_vs_alpha.svg | 1801 +++++++++++++++++ .../charts/cascade_decode_vs_alpha.svg | 1766 ++++++++++++++++ .../charts/manifest.json | 18 + 3 files changed, 3585 insertions(+) create mode 100644 docs/perf/qwen38-cascade-acceptance/charts/cascade_accept_vs_alpha.svg create mode 100644 docs/perf/qwen38-cascade-acceptance/charts/cascade_decode_vs_alpha.svg create mode 100644 docs/perf/qwen38-cascade-acceptance/charts/manifest.json diff --git a/docs/perf/qwen38-cascade-acceptance/charts/cascade_accept_vs_alpha.svg b/docs/perf/qwen38-cascade-acceptance/charts/cascade_accept_vs_alpha.svg new file mode 100644 index 000000000..638a46c60 --- /dev/null +++ b/docs/perf/qwen38-cascade-acceptance/charts/cascade_accept_vs_alpha.svg @@ -0,0 +1,1801 @@ + + + + + + + + 2026-09-08T10:41:52.001649 + image/svg+xml + + + Matplotlib v3.11.0, https://matplotlib.org/ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/docs/perf/qwen38-cascade-acceptance/charts/cascade_decode_vs_alpha.svg b/docs/perf/qwen38-cascade-acceptance/charts/cascade_decode_vs_alpha.svg new file mode 100644 index 000000000..2de1b3ccc --- /dev/null +++ b/docs/perf/qwen38-cascade-acceptance/charts/cascade_decode_vs_alpha.svg @@ -0,0 +1,1766 @@ + + + + + + + + 2026-09-08T10:41:51.964960 + image/svg+xml + + + Matplotlib v3.11.0, https://matplotlib.org/ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/docs/perf/qwen38-cascade-acceptance/charts/manifest.json b/docs/perf/qwen38-cascade-acceptance/charts/manifest.json new file mode 100644 index 000000000..5baddcf85 --- /dev/null +++ b/docs/perf/qwen38-cascade-acceptance/charts/manifest.json @@ -0,0 +1,18 @@ +[ + { + "alt": "Decode tok/s by cascade alpha, fastest of three with a min-max band, and two reference lines for the exact law and typical acceptance at 0.09", + "bytes": 60773, + "caption": "_Decode tok/s by cascade alpha, fastest of three seeds with a min-max band. The dashed line is the exact law (82.80) and the dotted line typical acceptance at 0.09 (104.65), both from the pooled 16K windows in #478. Alpha >= 1.0 saturates at accept 1.0 and the two highest alphas coincide._", + "file": "cascade_decode_vs_alpha.svg", + "key": "cascade_decode_vs_alpha", + "label": "cascade decode by alpha" + }, + { + "alt": "Cascade accept rate and tokens per cycle by alpha on two y-axes", + "bytes": 61850, + "caption": "_Cascade accept rate (left axis) and tokens per cycle (right axis) by alpha, from the run's cascade-accept metrics rather than the receipts. Higher alpha accepts more drafted tokens per verify._", + "file": "cascade_accept_vs_alpha.svg", + "key": "cascade_accept_vs_alpha", + "label": "cascade acceptance and tokens/cycle by alpha" + } +] From 150bea40b53503bb04cb2e031bf7acc3de743dc5 Mon Sep 17 00:00:00 2001 From: davidtai Date: Tue, 8 Sep 2026 11:40:38 -0500 Subject: [PATCH 11/30] docs(cascade): rule-diff byte-identity proof vs served code 2eac2fee Add the provenance bridge to docs/perf/qwen38-cascade-acceptance.md: arm G and the pending HumanEval cell were measured on served code 2eac2fee; this branch re-parents the cascade mode onto the #475 served base 27d5ff6b (the same base #478 uses), so #478 (typical) and this PR (cascade) are alternative-mode peers. Per-function sha256 proof shows the cascade rule (total_variation, _peak_probability, cascade_defer_decision, the batched/lazy verify branches, the readers, the verdict, the /health payload, the --cascade-threshold arg) is byte-identical to 2eac2fee, so the measurements transfer; the three forced glue deltas (lazy elif->if, exact-block re-indent + target_p_for_cache hoist, the mutual-exclusion reading the typical env name directly) are non-behavioral. --- docs/perf/qwen38-cascade-acceptance.md | 47 ++++++++++++++++++++++++++ 1 file changed, 47 insertions(+) diff --git a/docs/perf/qwen38-cascade-acceptance.md b/docs/perf/qwen38-cascade-acceptance.md index 1a48b8da4..99b907933 100644 --- a/docs/perf/qwen38-cascade-acceptance.md +++ b/docs/perf/qwen38-cascade-acceptance.md @@ -154,3 +154,50 @@ where the target clearly leads, which is the "gain speed while holding quality" operating point the sweep should confirm. If HumanEval drops, fall back toward alpha = 0 (more deferral); if quality holds, push toward alpha = 1.0 for more speed. + +## Provenance and rule-diff proof (arm G) + +The arm G 16K sweep and the pending HumanEval cell were measured on served code +`2eac2fee` (cascade stacked on the #478 typical head). This branch re-parents the +cascade mode onto the #475 served base `27d5ff6b`, the same base #478 is built on, +so #478 (typical) and this PR (cascade) are alternative-mode PEERS on one base +rather than a stack. This is the same provenance bridge #475 uses (docs commits +over the served base); the measurements transfer because the cascade RULE is +byte-identical on the new base. + +Per-function sha256 (first 16 hex) of the cascade computation, this branch vs the +measured `2eac2fee`: + +| symbol | sha256 | vs 2eac2fee | +| --- | --- | --- | +| `mtplx/sampling.py::_peak_probability` | `f3bba8b637a1ff9e` | identical | +| `mtplx/sampling.py::total_variation` | `1b960bfacfb37ce9` | identical | +| `mtplx/sampling.py::cascade_defer_decision` | `3ba053b7097fa91e` | identical | +| `mtplx/generation.py::_cascade_accept_alpha` | `b81b13739d2166b7` | identical | +| `mtplx/generation.py::_cascade_accept_enabled` | `1c5fae48d7f81e85` | identical | +| batched-target cascade verify branch | `6d2749c7d09e8cbb` | identical | +| lazy-target cascade verify body | `1ab16b485087dce2` | identical | +| `[cascade-accept]` verdict block | `78d7c3e84c84f6c1` | identical | +| `/health` `_cascade_acceptance_health_payload` (feature commit) | `7a9006f74a187d70` | identical | +| `--cascade-threshold` argparse block | (identical) | identical | + +The accept/defer/coin/residual, the total-variation and peak-probability +computations, the RNG draws, and the `[cascade-accept]` telemetry are therefore +byte-identical to the measured code. Three forced glue deltas remain, none of +which changes what a cascade run computes (so arm G stands): + +1. Lazy-target dispatch keyword: `elif _cascade_active:` on `2eac2fee` (it chained + off typical's `if _typical_active:`) becomes `if _cascade_active:` here, because + typical is dropped. The branch body is byte-identical. +2. The exact-path lazy block is re-indented one level under the new `else:` and + `target_p_for_cache = target_p` is hoisted above the dispatch (both idempotent + / non-behavioral; the cascade-off output is unchanged). +3. `_assert_lossy_verify_rules_exclusive` reads `MTPLX_FABLE_TYPICAL_THRESHOLD` + from the environment directly instead of calling the (now-absent) + `_typical_accept_threshold()`. + +Mutual exclusion (option a): the guard reads the typical env name defensively. It +is INERT on this branch, retained defensively; #478 and cascade are alternative +modes and were never intended to be armed together, so nothing here arms a typical +threshold, but the guard still fails loud (ValueError in the verify setup, +SystemExit at serve start) if an operator ever exports both keys. From ea9b2c09f40c8dc87f2058c0c9e41d8cc575107c Mon Sep 17 00:00:00 2001 From: davidtai Date: Tue, 8 Sep 2026 20:59:09 -0500 Subject: [PATCH 12/30] docs(qwen38-cascade-acceptance): add decode-by-context variance chart Add the cascade context-ladder decode chart (cascade_decode_by_context.svg): exact pairing and cascade alpha 0.0/0.5/1.0/2.0 across 1K-128K, fastest of seeds with min-max variance bands; 16,384 merged from the arm-G sweep; 261,120 absent (every arm OOMs on the #475 base without #482). Update the charts manifest. --- .../charts/cascade_decode_by_context.svg | 2237 +++++++++++++++++ .../charts/manifest.json | 9 + 2 files changed, 2246 insertions(+) create mode 100644 docs/perf/qwen38-cascade-acceptance/charts/cascade_decode_by_context.svg diff --git a/docs/perf/qwen38-cascade-acceptance/charts/cascade_decode_by_context.svg b/docs/perf/qwen38-cascade-acceptance/charts/cascade_decode_by_context.svg new file mode 100644 index 000000000..9ee820d8f --- /dev/null +++ b/docs/perf/qwen38-cascade-acceptance/charts/cascade_decode_by_context.svg @@ -0,0 +1,2237 @@ + + + + + + + + image/svg+xml + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/docs/perf/qwen38-cascade-acceptance/charts/manifest.json b/docs/perf/qwen38-cascade-acceptance/charts/manifest.json index 5baddcf85..cbb92f6a0 100644 --- a/docs/perf/qwen38-cascade-acceptance/charts/manifest.json +++ b/docs/perf/qwen38-cascade-acceptance/charts/manifest.json @@ -14,5 +14,14 @@ "file": "cascade_accept_vs_alpha.svg", "key": "cascade_accept_vs_alpha", "label": "cascade acceptance and tokens/cycle by alpha" + }, + { + "alt": "Decode tok/s by context size (1K-128K) for the exact pairing and cascade alpha 0.0/0.5/1.0/2.0, fastest of seeds with min-max variance bands", + "bytes": 80002, + "caption": "_Decode tok/s by context size, one line per arm (exact and cascade alpha 0.0/0.5/1.0/2.0), each point the fastest seed with a min-max band. 16,384 is the arm-G sweep merged in; 261,120 is absent because every arm OOMs on the #475 base without #482._", + "file": "cascade_decode_by_context.svg", + "key": "cascade_decode_by_context", + "label": "cascade decode by context", + "sha256": "3c1e893abbcebba625f8e6827acc1df3803c60a006b39b6c701a7f1365edaa56" } ] From ef0e83e856460b4cda222e6a15147b9dcabe38aa Mon Sep 17 00:00:00 2001 From: davidtai Date: Tue, 8 Sep 2026 21:15:53 -0500 Subject: [PATCH 13/30] cascade: restore [cascade-accept] verdict on the served generate_mtpk path The re-parent that dropped the typical lane deleted the [cascade-accept] verdict emission from generate_mtpk together with the adjacent [typical-accept] block, between _attach_runtime_diagnostics and the return. The rule still engaged on the served path (cascade_* counters and VerifyStats cascade fields moved, /health reported enabled), but the verdict line never printed, so the engagement gate that parses threshold=/positions= off it saw nothing at every alpha. The block had survived only in generate_mtpa, which is not the served loop. Restore the emission in generate_mtpk, guarded by if _cascade_active, byte-identical to the pre-re-parent block. The acceptance rule is untouched (per-function sha256 of cascade_defer_decision, total_variation, _peak_probability and the generation.py cascade-accept branch all match d8efc3f5). Add a served-order CPU test that arms the lane via the env reader at use, drives a temperature>0 mocked verify loop through generate_mtpk, and asserts one verdict line with threshold == alpha and positions > 0; it fails on the unfixed tree. --- mtplx/generation.py | 16 ++++ tests/test_cascade_acceptance.py | 129 +++++++++++++++++++++++++++++++ 2 files changed, 145 insertions(+) diff --git a/mtplx/generation.py b/mtplx/generation.py index 9f70c72fb..c6efa27cb 100644 --- a/mtplx/generation.py +++ b/mtplx/generation.py @@ -13830,6 +13830,22 @@ def emit_new_tokens() -> None: events=events, ) _attach_runtime_diagnostics(stats, rt, counter_start) + if _cascade_active: + _cas_denom = cascade_accepted + cascade_resamples + _cas_rate = (cascade_accepted / _cas_denom) if _cas_denom else 0.0 + _cas_cycles = max(1, verify_calls) + print( + "[cascade-accept] NOT distribution-exact; " + f"threshold={_cascade_alpha:.4g} alpha={_cascade_alpha:.4g} " + f"positions={cascade_positions} accepted={cascade_accepted} " + f"resamples={cascade_resamples} accept_rate={_cas_rate:.4f} " + f"mean_divergence={stats.cascade_mean_divergence:.4f} " + f"tokens_per_cycle={len(tokens) / _cas_cycles:.3f} " + f"accepted_by_depth={accepted_by_depth} " + f"generated={len(tokens)} verify_calls={verify_calls}", + file=sys.stderr, + flush=True, + ) return GenerationOutput( tokens=tokens, text=_decode(rt.tokenizer, _strip_terminal_stop(tokens, stop_token_ids)), diff --git a/tests/test_cascade_acceptance.py b/tests/test_cascade_acceptance.py index c569ed286..cd041e6a2 100644 --- a/tests/test_cascade_acceptance.py +++ b/tests/test_cascade_acceptance.py @@ -216,3 +216,132 @@ def test_exact_mode_when_off_takes_neither_lossy_branch(): assert _cascade_accept_enabled() is False _assert_lossy_verify_rules_exclusive() # inert here, must not raise + + +# --------------------------------------------------------------------------- +# Served-order regression guard (coordinator 2026-09-08): the re-parent that +# dropped the typical lane also deleted the [cascade-accept] verdict block from +# generate_mtpk (the served path) -- it survived only in generate_mtpa. The rule +# still engaged (cascade_* counters moved) but the verdict never printed, so the +# arm-G sweep's engagement gate (which reads threshold=/positions= off that line) +# saw nothing. This test drives generate_mtpk on CPU with the cascade lane armed +# via the env reader AT USE (served order: env set here, after import) and a +# temperature>0 mocked verify loop, and asserts the line prints once with +# threshold == alpha and positions > 0. It fails on any tree where the emission +# is missing from the SERVED function. +import re as _re +from pathlib import Path as _Path +from types import SimpleNamespace as _SNS + +import mlx.core as _mx +import pytest as _pytest + +from mtplx.generation import generate_mtpk as _generate_mtpk +from mtplx.mtp_patch import MTPContract as _MTPContract +from mtplx.runtime import MTPLXRuntime as _MTPLXRuntime +from mtplx.sampling import SamplerConfig as _SamplerConfig + + +class _VerdictTinyTokenizer: + def decode(self, tokens, **_kwargs): + return "".join(str(int(t)) for t in tokens) + + +class _VerdictAcceptingMTPModel: + """Minimal MTP model: deterministic logits favouring token 1 so the draft + and target agree and the cascade branch accepts (positions > 0).""" + + def __init__(self): + self.calls = [] + self.mtp = _SNS(_mtplx_lora_targets=[]) + + def make_cache(self): + return [] + + def make_mtp_cache(self): + return [] + + def mtp_update_cache(self, hidden_states, next_token_ids, *, mtp_cache=None, + concat_order=None, position_offset=None): + return hidden_states + + def mtp_forward(self, hidden_states, next_token_ids, *, mtp_cache=None, + concat_order=None, return_hidden=False, + mtp_hidden_variant=None, position_offset=None): + length = int(next_token_ids.shape[1]) + hidden = _mx.zeros((1, length, 2), dtype=_mx.float32) + logits = _mx.zeros((1, length, 4), dtype=_mx.float32) + _mx.array( + [0.0, 1.0, 0.0, 0.0], dtype=_mx.float32) + return (logits, hidden) if return_hidden else logits + + def __call__(self, input_ids, *, cache=None, return_hidden=False, + hidden_variant=None, emit_logits=True, logits_keep=None): + self.calls.append(int(input_ids.shape[1])) + length = int(input_ids.shape[1]) + hidden = _mx.zeros((1, length, 2), dtype=_mx.float32) + if not emit_logits: + return (None, hidden) if return_hidden else None + keep = length if logits_keep is None else min(length, max(1, int(logits_keep))) + logits = _mx.zeros((1, keep, 4), dtype=_mx.float32) + _mx.array( + [0.0, 1.0, 0.0, 0.0], dtype=_mx.float32) + return (logits, hidden) if return_hidden else logits + + +def _verdict_runtime(model): + return _MTPLXRuntime( + model=model, + tokenizer=_VerdictTinyTokenizer(), + model_path=_Path("tiny"), + mtp_enabled=True, + contract=_MTPContract(), + ) + + +def test_served_path_emits_cascade_verdict_with_threshold_and_positions( + capsys, monkeypatch +): + previous = _mx.default_device() + _mx.set_default_device(_mx.cpu) + try: + alpha = 0.5 + # Served order: the lane is armed via the env AFTER import; the reader + # resolves it at use inside generate_mtpk. Batched target arrays make the + # cascade branch reachable (target_distribution_batch is not None). + monkeypatch.setenv("MTPLX_FABLE_CASCADE_THRESHOLD", str(alpha)) + monkeypatch.setenv("MTPLX_BATCH_TARGET_ARRAYS", "1") + monkeypatch.delenv("MTPLX_FABLE_TYPICAL_THRESHOLD", raising=False) + + model = _VerdictAcceptingMTPModel() + out = _generate_mtpk( + _verdict_runtime(model), + [0], + max_tokens=5, + sampler=_SamplerConfig(temperature=0.6, top_p=1.0, top_k=1), + speculative_depth=3, + mtp_history_policy="committed", + verify_strategy="batched", + stop_token_ids=set(), + ) + finally: + _mx.set_default_device(previous) + + # The rule engaged: VerifyStats cascade fields populate on the served path. + assert out.stats.cascade_accept_enabled is True + assert out.stats.cascade_alpha == _pytest.approx(alpha) + assert out.stats.cascade_positions > 0 + + # The verdict line printed once on stderr, with threshold == alpha and + # positions > 0 (the exact fields the arm-G engagement gate parses). + err = capsys.readouterr().err + lines = [ln for ln in err.splitlines() if "[cascade-accept]" in ln] + assert len(lines) == 1, f"expected exactly one verdict line, got {lines!r}" + line = lines[0] + m_thr = _re.search(r"threshold=([0-9.eE+-]+)", line) + m_alpha = _re.search(r"alpha=([0-9.eE+-]+)", line) + m_pos = _re.search(r"positions=(\d+)", line) + assert m_thr and m_alpha and m_pos, line + assert float(m_thr.group(1)) == _pytest.approx(alpha) + assert float(m_alpha.group(1)) == _pytest.approx(alpha) + assert float(m_thr.group(1)) == float(m_alpha.group(1)) # threshold == alpha + assert int(m_pos.group(1)) > 0 + assert int(m_pos.group(1)) == out.stats.cascade_positions From 351c889c83882f59f17dad61f81bee3fdea8aec4 Mon Sep 17 00:00:00 2001 From: davidtai Date: Tue, 8 Sep 2026 23:04:27 -0500 Subject: [PATCH 14/30] cascade: drop the orphaned [cascade-accept] copy from the dead generate_mtpa generate_mtpa is an upstream function (present on base 27d5ff6b) with zero callers anywhere in the tree. The cascade feature commit d8efc3f5 accidentally added a [cascade-accept] verdict block to it while moving the block around during the re-parent; it references cascade counters that generate_mtpa never defines, so it is both dead and broken (would NameError if the function were ever called). It never ran -- the served verdict is in generate_mtpk. Remove only that copy, restoring generate_mtpa byte-identical to its upstream base form. The acceptance rule is untouched (per-function sha256 unchanged); the served generate_mtpk verdict stays. --- mtplx/generation.py | 16 ---------------- 1 file changed, 16 deletions(-) diff --git a/mtplx/generation.py b/mtplx/generation.py index c6efa27cb..30bd3316a 100644 --- a/mtplx/generation.py +++ b/mtplx/generation.py @@ -14130,22 +14130,6 @@ def generate_mtpa( stop_token_ids=stop_token_ids, max_tokens=max_tokens, ) - if _cascade_active: - _cas_denom = cascade_accepted + cascade_resamples - _cas_rate = (cascade_accepted / _cas_denom) if _cas_denom else 0.0 - _cas_cycles = max(1, verify_calls) - print( - "[cascade-accept] NOT distribution-exact; " - f"threshold={_cascade_alpha:.4g} alpha={_cascade_alpha:.4g} " - f"positions={cascade_positions} accepted={cascade_accepted} " - f"resamples={cascade_resamples} accept_rate={_cas_rate:.4f} " - f"mean_divergence={stats.cascade_mean_divergence:.4f} " - f"tokens_per_cycle={len(tokens) / _cas_cycles:.3f} " - f"accepted_by_depth={accepted_by_depth} " - f"generated={len(tokens)} verify_calls={verify_calls}", - file=sys.stderr, - flush=True, - ) return GenerationOutput( tokens=tokens, text=_decode(rt.tokenizer, _strip_terminal_stop(tokens, stop_token_ids)), From 26e3cc636d419815827af042ed98fbdf720c659f Mon Sep 17 00:00:00 2001 From: davidtai Date: Wed, 9 Sep 2026 06:27:05 -0500 Subject: [PATCH 15/30] docs(cascade): correct the Faster Cascades citation and name the rule precisely The reference carried a wrong author list and an unverified venue. The paper's authors are Narasimhan, Jitkrittum, Rawat, Kim, Gupta, Menon and Kumar; cite it as arXiv:2405.19261 v2 (2024) and drop the ICLR 2025 claim, which is not verifiable from the paper itself. Name the implemented rule as the paper does: it is r-hat_OPT, Equation (10), the plug-in ESTIMATOR of the optimal speculative-cascade deferral rule (Lemma 4, Equation (9)), which replaces that rule's ground-truth expected 0-1 losses with one minus each model's max probability. It is neither the optimal rule nor an oracle; the oracle needs expectations under the ground-truth distribution. The Diff rule (Equation (5)) is the sequential-cascade oracle and is not implemented here. Lemma 3 is why the deferral cost carries the alpha times D_TV term. Equation, Section, Algorithm and Lemma numbers are otherwise unchanged. Comment and docstring text only; no functional change. mtplx/sampling.py is deliberately NOT touched: its cascade_defer_decision docstring carries the same stale citation, but that function's full source is published as sha256 3ba053b7097fa91e and asserted byte-identical to the measured 2eac2fee, so editing it would invalidate a provenance line already live on the pull request. --- docs/perf/qwen38-cascade-acceptance.md | 16 ++++++++++------ tests/test_cascade_acceptance.py | 4 ++-- 2 files changed, 12 insertions(+), 8 deletions(-) diff --git a/docs/perf/qwen38-cascade-acceptance.md b/docs/perf/qwen38-cascade-acceptance.md index 99b907933..d043a5318 100644 --- a/docs/perf/qwen38-cascade-acceptance.md +++ b/docs/perf/qwen38-cascade-acceptance.md @@ -5,11 +5,15 @@ beside typical acceptance. Per draft position it decides whether the draft token is good enough to keep, or whether to defer to the exact target law. It is OFF by default, mutually exclusive with typical acceptance, and NOT distribution-exact. -Citation: Narasimhan, Mreddy, Jitkrittum, Rawat, Kumar, "Faster Cascades via -Speculative Decoding," ICLR 2025 / arXiv:2405.19261 v2. This implements the -plug-in deferral rule of Section 4.3, Equation (10) (the plug-in approximation -to the optimal rule of Equation (8)), executed with the speculative decoding of -Algorithm 4. +Citation: Narasimhan, Jitkrittum, Rawat, Kim, Gupta, Menon, and Kumar, "Faster Cascades via +Speculative Decoding," arXiv:2405.19261 v2 (2024). This implements the +r-hat_OPT deferral rule of Section 4.3, Equation (10): the plug-in ESTIMATOR of +the optimal speculative-cascade deferral rule (Lemma 4, Equation (9)), replacing +that rule's ground-truth expected 0-1 losses with one minus each model's max +probability. It is not the optimal rule and not an oracle. The separate Diff rule +(Equation (5), max q < max p - alpha, no total-variation term) is the +SEQUENTIAL-cascade oracle and is not implemented here. Executed with the +speculative decoding of Algorithm 4. Terms: @@ -63,7 +67,7 @@ is SUBTRACTED from the target's confidence. So more disagreement between the draft and the target LOWERS the bar for accepting the draft, which means the rule defers LESS when they disagree, not more. The paper's reason (their Lemma 3) is that a large disagreement makes the verification step itself expensive, so the -optimal rule only pays that cost when the target is clearly better; it accepts a +rule only pays that cost when the target is clearly better; it accepts a draft that is confident, even on a token the target would not have picked. The intuitive "reject when the draft diverges" behaviour still holds for the common case, a diverging draft that has also lost its peak confidence, which defers; a diff --git a/tests/test_cascade_acceptance.py b/tests/test_cascade_acceptance.py index cd041e6a2..e647ebdc4 100644 --- a/tests/test_cascade_acceptance.py +++ b/tests/test_cascade_acceptance.py @@ -1,7 +1,7 @@ """CPU tests for speculative-cascade acceptance (a second lossy verify rule). -Rule implemented: Narasimhan, Mreddy, Jitkrittum, Rawat, Kumar, "Faster -Cascades via Speculative Decoding" (arXiv:2405.19261 v2), Section 4.3 +Rule implemented: Narasimhan, Jitkrittum, Rawat, Kim, Gupta, Menon, and Kumar, "Faster +Cascades via Speculative Decoding" (arXiv:2405.19261 v2, 2024), Section 4.3 Equation (10), the plug-in approximation to the optimal deferral rule of Equation (8): From 48949e6b517533579dc7746a0da800736caa786e Mon Sep 17 00:00:00 2001 From: davidtai Date: Wed, 9 Sep 2026 06:43:43 -0500 Subject: [PATCH 16/30] docs(cascade): re-render charts with the alpha 0.25 and 0.75 points cascade_decode_vs_alpha, cascade_accept_vs_alpha and cascade_decode_by_context now include the two added grid points (six cascade alphas: 0.0, 0.25, 0.5, 0.75, 1.0, 2.0). Fastest-of-seeds with min-max bands, same generators. Charts only; no code or measurement change. --- .../charts/cascade_accept_vs_alpha.svg | 474 ++++++----- .../charts/cascade_decode_by_context.svg | 752 ++++++++++-------- .../charts/cascade_decode_vs_alpha.svg | 376 +++++---- .../charts/manifest.json | 14 +- 4 files changed, 940 insertions(+), 676 deletions(-) diff --git a/docs/perf/qwen38-cascade-acceptance/charts/cascade_accept_vs_alpha.svg b/docs/perf/qwen38-cascade-acceptance/charts/cascade_accept_vs_alpha.svg index 638a46c60..5c7af15f1 100644 --- a/docs/perf/qwen38-cascade-acceptance/charts/cascade_accept_vs_alpha.svg +++ b/docs/perf/qwen38-cascade-acceptance/charts/cascade_accept_vs_alpha.svg @@ -6,7 +6,7 @@ - 2026-09-08T10:41:52.001649 + 2026-09-09T06:34:59.321459 image/svg+xml @@ -42,21 +42,21 @@ z +" clip-path="url(#p74365cf68c)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - - + - - + + + - + - + - - + + + - + + - + - + - - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - - + + +" clip-path="url(#p74365cf68c)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - + - + - - - - - - + + + + - + @@ -434,36 +490,24 @@ z - + +" clip-path="url(#p74365cf68c)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - + - - + - + - - - @@ -472,17 +516,17 @@ z - + +" clip-path="url(#p74365cf68c)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - + - + - + @@ -493,17 +537,17 @@ L 550.8 270.036 - + +" clip-path="url(#p74365cf68c)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - + - + - + @@ -555,17 +599,17 @@ z - + +" clip-path="url(#p74365cf68c)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - + - + - + @@ -576,17 +620,17 @@ L 550.8 181.116 - + +" clip-path="url(#p74365cf68c)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - + - + - + @@ -629,17 +673,17 @@ z - + +" clip-path="url(#p74365cf68c)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - + - + - + @@ -650,17 +694,17 @@ L 550.8 92.196 - + +" clip-path="url(#p74365cf68c)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - + - + - + @@ -670,7 +714,7 @@ L 550.8 47.736 - + @@ -755,7 +799,7 @@ L 550.8 314.496 L 550.8 29.952 " style="fill: none; stroke: #000000; stroke-width: 0.8; stroke-linejoin: miter; stroke-linecap: square"/> - + @@ -765,7 +809,49 @@ L 550.8 29.952 - + + + + + + + + + + + + + + @@ -775,7 +861,17 @@ L 550.8 29.952 - + + + + + + + + + + + @@ -785,7 +881,7 @@ L 550.8 29.952 - + @@ -795,7 +891,7 @@ L 550.8 29.952 - + @@ -971,14 +1067,16 @@ z - + +" clip-path="url(#p74365cf68c)" style="fill: none; stroke: #0072b2; stroke-width: 1.8; stroke-linecap: square"/> - - - - - - + + + + + + + @@ -1011,13 +1111,13 @@ Q 415.31375 185.424625 416.91375 185.424625 z " style="fill: #ffffff; opacity: 0.9; stroke: #cccccc; stroke-linejoin: miter"/> - + - - + - + @@ -1113,13 +1213,13 @@ z - + - - + - + @@ -1213,17 +1313,17 @@ z - + - - + - + @@ -1258,36 +1358,6 @@ Q 3450 4097 3450 3541 Q 3450 3153 3228 2886 Q 3006 2619 2597 2516 z -" transform="scale(0.015625)"/> - @@ -1297,12 +1367,12 @@ z - + - + - + @@ -1312,12 +1382,12 @@ z - + - + - + @@ -1348,12 +1418,12 @@ z - + - + - + @@ -1363,12 +1433,12 @@ z - + - + - + @@ -1378,12 +1448,12 @@ z - + - + - + @@ -1393,12 +1463,12 @@ z - + - + - + @@ -1407,7 +1477,7 @@ z - + @@ -1449,7 +1519,7 @@ L 550.8 314.496 L 550.8 29.952 " style="fill: none; stroke: #000000; stroke-width: 0.8; stroke-linejoin: miter; stroke-linecap: square"/> - + @@ -1458,7 +1528,16 @@ L 550.8 29.952 - + + + + + + + + + + @@ -1467,7 +1546,16 @@ L 550.8 29.952 - + + + + + + + + + + @@ -1476,7 +1564,7 @@ L 550.8 29.952 - + @@ -1485,29 +1573,33 @@ L 550.8 29.952 - + +" clip-path="url(#p74365cf68c)" style="fill: none; stroke-dasharray: 6.66,2.88; stroke-dashoffset: 0; stroke: #e69f00; stroke-width: 1.8"/> - - - - - - + + + + + + + - + @@ -1794,7 +1886,7 @@ z - + diff --git a/docs/perf/qwen38-cascade-acceptance/charts/cascade_decode_by_context.svg b/docs/perf/qwen38-cascade-acceptance/charts/cascade_decode_by_context.svg index 9ee820d8f..3cfbd926c 100644 --- a/docs/perf/qwen38-cascade-acceptance/charts/cascade_decode_by_context.svg +++ b/docs/perf/qwen38-cascade-acceptance/charts/cascade_decode_by_context.svg @@ -33,117 +33,163 @@ z - - + - - + - - + - - + - - + + + + + + + + + + + + + + + + + @@ -785,8 +831,8 @@ z - @@ -796,23 +842,13 @@ L -3.5 0 " style="stroke: #1a1a1a; stroke-width: 0.8"/> - + - - + + - - + - - + - - - + + + + - - + - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + - - - - - - - - - - - - - - - - - - - - - + + - + - + - + - + - + @@ -1214,12 +1140,12 @@ L 567.36 274.104 L 567.36 25.38 " style="fill: none; stroke: #4d4d4d; stroke-width: 0.8; stroke-linejoin: miter; stroke-linecap: square"/> - - + @@ -1236,21 +1162,21 @@ z " style="stroke: #0072b2"/> - - - - - + + + + + - - + - - - - - - + + + + + + + + + + + + + + + + + + + + - - + - - - - - - + + + + + + + + + + + + + + + + + + + + - - + - - - - - - + + + + + + - - + - - - - - - + + + + + + - + @@ -1470,21 +1456,21 @@ z - - + - + - - + - + - + - + - - + - + - + - + - - + + - + - - - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + @@ -1682,18 +1700,62 @@ z - - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - + - + - + @@ -1713,18 +1775,18 @@ L 477.74125 83.86125 - - + - + - + - + @@ -1746,7 +1808,7 @@ L 477.74125 95.861875 - + @@ -1756,6 +1818,36 @@ L 1997 1497 L 313 1497 L 313 2009 z +" transform="scale(0.015625)"/> + - 2026-09-08T10:41:51.964960 + 2026-09-09T06:34:59.277168 image/svg+xml @@ -42,21 +42,21 @@ z +" clip-path="url(#p56cbe5e2b6)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - - + - - + + + - + - + - - + + + - + + - + - + - - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - - + + +" clip-path="url(#p56cbe5e2b6)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - + - + - - - - - - + + + + - + @@ -434,22 +490,22 @@ z - + +" clip-path="url(#p56cbe5e2b6)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - + - - + - + @@ -499,17 +555,17 @@ z - + +" clip-path="url(#p56cbe5e2b6)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - + - + - + @@ -550,17 +606,17 @@ z - + +" clip-path="url(#p56cbe5e2b6)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - + - + - + @@ -570,17 +626,17 @@ L 593.64 222.823575 - + +" clip-path="url(#p56cbe5e2b6)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - + - + - + @@ -590,17 +646,17 @@ L 593.64 177.045315 - + +" clip-path="url(#p56cbe5e2b6)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - + - + - + @@ -610,17 +666,17 @@ L 593.64 131.267056 - + +" clip-path="url(#p56cbe5e2b6)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - + - + - + @@ -664,17 +720,17 @@ z - + +" clip-path="url(#p56cbe5e2b6)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - + - + - + @@ -704,7 +760,7 @@ z - + @@ -844,15 +900,15 @@ z - + +" clip-path="url(#p56cbe5e2b6)" style="fill: none; stroke-dasharray: 4.44,1.92; stroke-dashoffset: 0; stroke: #888888; stroke-width: 1.2"/> - + +" clip-path="url(#p56cbe5e2b6)" style="fill: none; stroke-dasharray: 1.4,2.31; stroke-dashoffset: 0; stroke: #d55e00; stroke-width: 1.4"/> - + @@ -927,7 +983,7 @@ z - + @@ -984,28 +1040,26 @@ z - + - - - - + + + + + + + + + + + @@ -1015,7 +1069,17 @@ z - + + + + + + + + + + + @@ -1025,7 +1089,7 @@ z - + @@ -1035,7 +1099,7 @@ z - + @@ -1247,46 +1311,58 @@ z +" clip-path="url(#p56cbe5e2b6)" style="fill: none; stroke: #009e73; stroke-width: 1.4"/> + +" clip-path="url(#p56cbe5e2b6)" style="fill: none; stroke: #009e73; stroke-width: 1.4"/> + +" clip-path="url(#p56cbe5e2b6)" style="fill: none; stroke: #009e73; stroke-width: 1.4"/> +" clip-path="url(#p56cbe5e2b6)" style="fill: none; stroke: #009e73; stroke-width: 1.4"/> - + - - - - - - + + + + + + + - - - - - - + + + + + + + + - + +" clip-path="url(#p56cbe5e2b6)" style="fill: none; stroke: #009e73; stroke-width: 1.8; stroke-linecap: square"/> - - - - - - + + + + + + + @@ -1319,13 +1397,13 @@ Q 65.2 72.353875 66.8 72.353875 z " style="fill: #ffffff; opacity: 0.9; stroke: #cccccc; stroke-linejoin: miter"/> - + - - + - + @@ -1379,13 +1457,13 @@ z - + - + @@ -1407,13 +1485,13 @@ L 84.4 52.431375 - + - + @@ -1441,7 +1519,7 @@ L 84.4 64.432 - + @@ -1759,7 +1837,7 @@ z - + diff --git a/docs/perf/qwen38-cascade-acceptance/charts/manifest.json b/docs/perf/qwen38-cascade-acceptance/charts/manifest.json index cbb92f6a0..07040999c 100644 --- a/docs/perf/qwen38-cascade-acceptance/charts/manifest.json +++ b/docs/perf/qwen38-cascade-acceptance/charts/manifest.json @@ -1,27 +1,29 @@ [ { "alt": "Decode tok/s by cascade alpha, fastest of three with a min-max band, and two reference lines for the exact law and typical acceptance at 0.09", - "bytes": 60773, + "bytes": 64678, "caption": "_Decode tok/s by cascade alpha, fastest of three seeds with a min-max band. The dashed line is the exact law (82.80) and the dotted line typical acceptance at 0.09 (104.65), both from the pooled 16K windows in #478. Alpha >= 1.0 saturates at accept 1.0 and the two highest alphas coincide._", "file": "cascade_decode_vs_alpha.svg", "key": "cascade_decode_vs_alpha", - "label": "cascade decode by alpha" + "label": "cascade decode by alpha", + "sha256": "224d8ac00e5beeb7347be25d2d9bb07c41a250dcaf1e68fc81f423f1528f5b41" }, { "alt": "Cascade accept rate and tokens per cycle by alpha on two y-axes", - "bytes": 61850, + "bytes": 66171, "caption": "_Cascade accept rate (left axis) and tokens per cycle (right axis) by alpha, from the run's cascade-accept metrics rather than the receipts. Higher alpha accepts more drafted tokens per verify._", "file": "cascade_accept_vs_alpha.svg", "key": "cascade_accept_vs_alpha", - "label": "cascade acceptance and tokens/cycle by alpha" + "label": "cascade acceptance and tokens/cycle by alpha", + "sha256": "1c63ea4c8e0b459b9d10eedfe533e4fcdcbbead7ddc3a9a832ee2332a563a047" }, { "alt": "Decode tok/s by context size (1K-128K) for the exact pairing and cascade alpha 0.0/0.5/1.0/2.0, fastest of seeds with min-max variance bands", - "bytes": 80002, + "bytes": 84651, "caption": "_Decode tok/s by context size, one line per arm (exact and cascade alpha 0.0/0.5/1.0/2.0), each point the fastest seed with a min-max band. 16,384 is the arm-G sweep merged in; 261,120 is absent because every arm OOMs on the #475 base without #482._", "file": "cascade_decode_by_context.svg", "key": "cascade_decode_by_context", "label": "cascade decode by context", - "sha256": "3c1e893abbcebba625f8e6827acc1df3803c60a006b39b6c701a7f1365edaa56" + "sha256": "fc7c4b1f4b42bb32d739248f5b47d2890311ce4e0f2c298ea2097a407ed96784" } ] From 118e250b280beb55b0c6d0ed61e7d0c039d32c0a Mon Sep 17 00:00:00 2001 From: davidtai Date: Wed, 9 Sep 2026 09:56:51 -0500 Subject: [PATCH 17/30] docs(cascade): two-rule alpha charts (OPT and TokenV3) cascade_decode_vs_alpha and cascade_accept_vs_alpha now plot both deferral rules: OPT (Equation 10, alpha 0/0.25/0.5/0.75/1/2) and TokenV3 (Equation 15, alpha 0.25/0.5/0.75/0.9/0.95), with distinct markers and dash, the same exact-law and typical-0.09 reference lines, and the legend outside the axes. Captions name both rules and state that equal alphas are NOT comparable across them, since TokenV3's alpha is a fraction of the target peak probability. Charts and manifest only; nothing under mtplx/. --- .../charts/cascade_accept_vs_alpha.svg | 2656 +++++++++------- .../charts/cascade_decode_vs_alpha.svg | 2690 +++++++++++------ .../charts/manifest.json | 12 +- 3 files changed, 3283 insertions(+), 2075 deletions(-) diff --git a/docs/perf/qwen38-cascade-acceptance/charts/cascade_accept_vs_alpha.svg b/docs/perf/qwen38-cascade-acceptance/charts/cascade_accept_vs_alpha.svg index 5c7af15f1..c499c4c2b 100644 --- a/docs/perf/qwen38-cascade-acceptance/charts/cascade_accept_vs_alpha.svg +++ b/docs/perf/qwen38-cascade-acceptance/charts/cascade_accept_vs_alpha.svg @@ -1,12 +1,12 @@ - + - 2026-09-09T06:34:59.321459 + 2026-09-09T09:55:36.710149 image/svg+xml @@ -21,42 +21,42 @@ - - - + - - + - - + + - - - - - + - + - + + - + - + - - + + - - + - + - + - + - + - - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - - - - - - + + + - + - + - - - + + + - - - - - - + + + - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - - + + - + - + - + - + @@ -620,51 +911,19 @@ L 550.8 181.116 - - + + - + - + - + - - - - + @@ -673,19 +932,19 @@ z - - + + - + - + - + - + @@ -694,19 +953,19 @@ L 550.8 92.196 - - + + - + - + - + - + @@ -714,46 +973,34 @@ L 550.8 47.736 - + - + - - @@ -780,140 +1027,43 @@ z - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + + + + + + + + + + - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - - + + - - - - - - - - + + + + + + + + + + + + + + + + + + + + - + - - + - - + - - - + + + - - - - - - - - - - - - - - - - - - - - - - - - - - - - + + + + + + + + + + + + + + + + + + + + + + - - + + - +" style="stroke: #0072b2; stroke-linejoin: miter"/> + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - + - - - + + + - - - - - - - - - - - - - - - - - - - - - - - - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + @@ -1313,83 +1598,83 @@ z - + - - + - - - - - - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - + + - - + + - + - - + + - + + - - + + - + - - + + + - - + + - + - - + + + - - + + - + - - + + - + + - - + + - + - - - - - - - - - - - - - - - - - + + - + + - + - + @@ -1500,177 +1774,83 @@ z - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + + + + + + + + + + + + - - + + - - - - - - - - + + + + + + - - - + + + - - - - - - + @@ -1771,123 +1973,445 @@ z - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - - + + diff --git a/docs/perf/qwen38-cascade-acceptance/charts/cascade_decode_vs_alpha.svg b/docs/perf/qwen38-cascade-acceptance/charts/cascade_decode_vs_alpha.svg index cc5c66332..62eef2403 100644 --- a/docs/perf/qwen38-cascade-acceptance/charts/cascade_decode_vs_alpha.svg +++ b/docs/perf/qwen38-cascade-acceptance/charts/cascade_decode_vs_alpha.svg @@ -1,12 +1,12 @@ - + - 2026-09-09T06:34:59.277168 + 2026-09-09T09:55:36.647146 image/svg+xml @@ -21,42 +21,42 @@ - - - + - - + - - + + - - - - - + - + - + + - + - + - - + + - - + - + - + - + - + - - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - - - - - - + + + - + - + - - - + + + - - - - - - + + + - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + + + - - - - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - - - + + + + + - + + + + - + - - - + + + - + + + + + + + + + + + + + + + + + - - + + - + - + - + - + @@ -626,19 +885,19 @@ L 593.64 222.823575 - - + + - + - + - + - + @@ -646,19 +905,19 @@ L 593.64 177.045315 - - + + - + - + - + - + @@ -666,19 +925,19 @@ L 593.64 131.267056 - - + + - + - + - + - + - - + + - + - + - + - + + - + - - - + + - - + + - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + @@ -1051,7 +1184,7 @@ z - + @@ -1061,7 +1194,7 @@ z - + @@ -1071,7 +1204,7 @@ z - + @@ -1081,7 +1214,7 @@ z - + @@ -1091,7 +1224,7 @@ z - + @@ -1100,67 +1233,89 @@ z - - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - - - - - - + - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - - - - - - + + + + + + - + - - - - - - - - + + + + + + + - - - - - - - - + + + + + + + + - - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - - - - - - - - + + + + + + + + + + + + + + + + + + + + - + - - + - - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - - + - + - + @@ -1485,15 +1837,15 @@ L 84.4 52.431375 - - + - + - + @@ -1519,10 +1871,321 @@ L 84.4 64.432 - - - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - - + + diff --git a/docs/perf/qwen38-cascade-acceptance/charts/manifest.json b/docs/perf/qwen38-cascade-acceptance/charts/manifest.json index 07040999c..0e8f825c7 100644 --- a/docs/perf/qwen38-cascade-acceptance/charts/manifest.json +++ b/docs/perf/qwen38-cascade-acceptance/charts/manifest.json @@ -1,21 +1,21 @@ [ { "alt": "Decode tok/s by cascade alpha, fastest of three with a min-max band, and two reference lines for the exact law and typical acceptance at 0.09", - "bytes": 64678, - "caption": "_Decode tok/s by cascade alpha, fastest of three seeds with a min-max band. The dashed line is the exact law (82.80) and the dotted line typical acceptance at 0.09 (104.65), both from the pooled 16K windows in #478. Alpha >= 1.0 saturates at accept 1.0 and the two highest alphas coincide._", + "bytes": 94167, + "caption": "_Decode by alpha for BOTH cascade rules: OPT (Equation 10) and TokenV3 (Equation 15), each fastest of three seeds with a min-max band. Reference lines are the exact law (82.80) and typical acceptance at 0.09 (104.65) from the pooled 16K windows in #478. The two rules' alphas are different quantities and are NOT comparable at equal alpha; compare at equal decode speed._", "file": "cascade_decode_vs_alpha.svg", "key": "cascade_decode_vs_alpha", "label": "cascade decode by alpha", - "sha256": "224d8ac00e5beeb7347be25d2d9bb07c41a250dcaf1e68fc81f423f1528f5b41" + "sha256": "f191004d8dc77f46079ac5ba7123ccde2377a743be69cbdcdba649344fdb1f5f" }, { "alt": "Cascade accept rate and tokens per cycle by alpha on two y-axes", - "bytes": 66171, - "caption": "_Cascade accept rate (left axis) and tokens per cycle (right axis) by alpha, from the run's cascade-accept metrics rather than the receipts. Higher alpha accepts more drafted tokens per verify._", + "bytes": 86003, + "caption": "_Cascade accept rate (left) and tokens per cycle (right) by alpha for BOTH rules, OPT (Equation 10) and TokenV3 (Equation 15), from each run's cascade-accept metrics rather than the receipts. Equal alphas are NOT comparable across the two rules._", "file": "cascade_accept_vs_alpha.svg", "key": "cascade_accept_vs_alpha", "label": "cascade acceptance and tokens/cycle by alpha", - "sha256": "1c63ea4c8e0b459b9d10eedfe533e4fcdcbbead7ddc3a9a832ee2332a563a047" + "sha256": "47531a838091993d3c928b4d9cf4f3b292fbfc1fed0081edeef86bcb84eba860" }, { "alt": "Decode tok/s by context size (1K-128K) for the exact pairing and cascade alpha 0.0/0.5/1.0/2.0, fastest of seeds with min-max variance bands", From a2531bde7779c0e0ac68304d16538bc98730caf5 Mon Sep 17 00:00:00 2001 From: davidtai Date: Wed, 9 Sep 2026 06:42:35 -0500 Subject: [PATCH 18/30] cascade: add token-specific deferral rules (TokenV1/V3) behind --cascade-rule r_OPT (arXiv:2405.19261 v2, Eq. 10) decides between q and p by comparing only their peaks, so a drafted token x_t~q that does not maximise q can be accepted because q is more peaked than p even when the token is poor (Sec. 4.4). This is the cause of the HumanEval loss at every alpha. Implement the token-specific rules that judge the drafted token: tokenv1 (Eq. 13): defer v iff q(v) < max_v' p(v') - alpha tokenv2 (Eq. 14): defer v iff p(v) < max_v' p(v') - alpha tokenv3 (Eq. 15): defer v iff p(v) < max_v' p(v') * (1 - alpha) On a deferred token the exact coin/residual runs with the token-specific target pi_Token (Eq. 11) instead of p -- Algorithm 6 (Appendix D) is GenSpecSample(q, p, pi_Token). A token in Top_alpha has pi(v)=q(v)+p(v)*eta, so the coin accepts it with probability 1 (accept, no coin); a deferred token has pi(v)=p(v)*eta. Rule selector: env MTPLX_FABLE_CASCADE_RULE / flag --cascade-rule, default opt for backward compatibility, same alpha knob, same mutual exclusion with typical. /health reports the rule; the [cascade-accept] verdict line names it. r_OPT's executed code is unchanged: the token-specific branches are a new elif above the OPT branch at both verify sites, so opt (or unset) is byte-for-byte as before. Also correct the stale citation in cascade_defer_decision's docstring (Mreddy/ICLR 2025 -> Narasimhan, Jitkrittum, Rawat, Kim, Gupta, Menon, Kumar, arXiv:2405.19261 v2 (2024)). Since that touches the sha256-hashed function body, docs/perf restates the proof as an AST code-hash (docstring stripped) showing the OPT rule code identical to 2eac2fee/d8efc3f5 while the docstring text changed. Tests: TokenV3 defers a confidently-wrong draft that OPT accepts, Top_alpha accepted, pi_TokenV3 matches Eq. 11 numerically, TokenV1 rule, served-order arming naming the rule, rule-selector read-at-use + default opt + fail-loud. The block-verify structural guard now counts six exact-coin sites (two shipped exact, two OPT defer, two token-specific defer), one coin per depth. --- docs/perf/qwen38-cascade-acceptance.md | 75 ++++++++++++++++ mtplx/generation.py | 109 ++++++++++++++++++++++- mtplx/sampling.py | 101 ++++++++++++++++++++- mtplx/server/openai.py | 29 +++++- tests/test_cascade_acceptance.py | 117 +++++++++++++++++++++++++ tests/test_qwen4_block_verify.py | 21 +++-- 6 files changed, 440 insertions(+), 12 deletions(-) diff --git a/docs/perf/qwen38-cascade-acceptance.md b/docs/perf/qwen38-cascade-acceptance.md index d043a5318..af9de702b 100644 --- a/docs/perf/qwen38-cascade-acceptance.md +++ b/docs/perf/qwen38-cascade-acceptance.md @@ -205,3 +205,78 @@ is INERT on this branch, retained defensively; #478 and cascade are alternative modes and were never intended to be armed together, so nothing here arms a typical threshold, but the guard still fails loud (ValueError in the verify setup, SystemExit at serve start) if an operator ever exports both keys. + +## Token-specific deferral rules (TokenV1 / TokenV3) + +`r_OPT` (Eq. 10) decides between q and p by comparing only their peaks. Sec. 4.4 +of the paper names the failure that costs us HumanEval at every alpha: the draft +token `x_t ~ q_t(.)` may not maximise `q_t`, so "even when `x_t` is of poor +quality, we may end up accepting it because `q_t` happens to be more peaked than +`p_t`." Their fix is a token-specific rule `r(x_= max_v' p(v') * (1 - alpha) }. + +The `p(v)*eta` term is present for every v, so a token in `Top_alpha` +(`r = 0`) has `pi(v) = q(v) + p(v)*eta >= q(v)` and the generic speculative coin +accepts it with probability 1 (accept, no coin drawn); a deferred token +(`r = 1`) has `pi(v) = p(v)*eta` and takes the exact `min(1, pi/q)` coin plus +`norm(max(0, pi - q))` residual -- Algorithm 6 (Appendix D) is +`GenSpecSample(q, p, pi_Token)`, i.e. the shipped exact path with `pi_Token` as +the target instead of `p`. `sum_v pi_Token(v) = 1` by construction. A +confidently-wrong drafted token (`p(v)` small) is now DEFERRED even when +`max q > max p` -- exactly the case `r_OPT` accepts. + +`r_OPT` stays the default and its executed code is unchanged; the token-specific +branches are added as a new `elif ... _cascade_rule != "opt" ...` above the OPT +branch at both verify sites, so with `--cascade-rule opt` (or unset) the OPT +path is byte-for-byte what it was. + +### Provenance after the citation fix + +The TokenV3 commit also corrects the stale citation in +`cascade_defer_decision`'s docstring (it read "Mreddy" / "ICLR 2025"; corrected +to "Narasimhan, Jitkrittum, Rawat, Kim, Gupta, Menon, Kumar, arXiv:2405.19261 v2 +(2024)"). That edit is inside the sha256-hashed function body #485 cited as +byte-identical to `2eac2fee`, so the whole-function TEXT hash of +`cascade_defer_decision` necessarily changes. We therefore restate the proof two +ways -- whole-function text (docstring included) AND AST of the function body +with the docstring node stripped (its executable code): + +Whole-function TEXT sha256 (first 16 hex; docstring included): + +| function | 2eac2fee | d8efc3f5 | this branch | +| --- | --- | --- | --- | +| `cascade_defer_decision` | `1f5061f6caa3838d` | `1f5061f6caa3838d` | `895a81bd203accc5` (docstring changed) | +| `total_variation` | `804fe67757c30ea3` | `804fe67757c30ea3` | `804fe67757c30ea3` (same) | +| `_peak_probability` | `3eb10a6b9e83045a` | `3eb10a6b9e83045a` | `3eb10a6b9e83045a` (same) | + +AST CODE sha256 (first 16 hex; docstring node removed): + +| function | 2eac2fee | d8efc3f5 | this branch | +| --- | --- | --- | --- | +| `cascade_defer_decision` | `1e08468afd459849` | `1e08468afd459849` | `1e08468afd459849` (same) | +| `total_variation` | `d4832efa65b6a0e2` | `d4832efa65b6a0e2` | `d4832efa65b6a0e2` (same) | +| `_peak_probability` | `9c62b65e7dd51f9c` | `9c62b65e7dd51f9c` | `9c62b65e7dd51f9c` (same) | + +The OPT rule's *code* is thus byte-identical to `2eac2fee`/`d8efc3f5`; only the +docstring text changed. The measured arm-G OPT numbers stand. (The whole-file +text hashes in the earlier table above use a different extractor and are +unaffected; the two tables here are self-consistent, produced by one +docstring-stripping AST script.) diff --git a/mtplx/generation.py b/mtplx/generation.py index 30bd3316a..d5508dfa6 100644 --- a/mtplx/generation.py +++ b/mtplx/generation.py @@ -101,9 +101,12 @@ SparseDistribution, acceptance_probability as compute_acceptance_probability, cascade_defer_decision, + cascade_token_deferral, + cascade_token_target_distribution, distribution_from_logits as dense_distribution_from_logits, residual_distribution, sample_from_distribution, + total_variation, ) from .session_bank import _boundary_true_restore_enabled from .runtime_options import ( @@ -634,6 +637,33 @@ def _cascade_accept_enabled() -> bool: return _cascade_accept_alpha() is not None +def _cascade_accept_rule() -> str: + """MTPLX_FABLE_CASCADE_RULE (server flag --cascade-rule): which cascade + deferral rule to apply. Default "opt" for backward compatibility. + + opt -- r_OPT (arXiv:2405.19261 v2, Eq. 10): defer iff + max_v q(v) < max_v p(v) - alpha*D_TV(p,q); position-level. + tokenv1 -- r_TokenV1 (Eq. 13): defer iff q(v) < max_v' p(v') - alpha. + tokenv2 -- r_TokenV2 (Eq. 14): defer iff p(v) < max_v' p(v') - alpha. + tokenv3 -- r_TokenV3 (Eq. 15): defer iff p(v) < max_v' p(v')*(1-alpha). + + The token-specific rules (Sec. 4.4) judge the drafted token v rather than + only the peaks, and on a deferred token run the exact coin/residual with the + token-specific target pi_Token (Eq. 11, Algorithm 6). Resolved at use, per + request; a malformed value fails loud. + """ + raw = os.environ.get("MTPLX_FABLE_CASCADE_RULE") + if raw is None or str(raw).strip() == "": + return "opt" + rule = str(raw).strip().lower() + if rule not in ("opt", "tokenv1", "tokenv2", "tokenv3"): + raise ValueError( + f"unknown MTPLX_FABLE_CASCADE_RULE={raw!r}; " + "expected one of opt|tokenv1|tokenv2|tokenv3" + ) + return rule + + def _assert_lossy_verify_rules_exclusive() -> None: """Fail loud when both lossy verify rules are armed at once. @@ -8764,6 +8794,7 @@ def record_adaptive_width_event( _assert_lossy_verify_rules_exclusive() _cascade_alpha = _cascade_accept_alpha() _cascade_active = _cascade_alpha is not None and sampler.temperature > 0 + _cascade_rule = _cascade_accept_rule() # Loop Guard: loop-armed DRY-style steering (see mtplx/loop_guard.py). # Disarmed = zero distribution impact (identity transform, fast paths kept). # Armed = target distributions get sparse anti-cycle penalties per position; @@ -12367,6 +12398,50 @@ def emit_new_tokens() -> None: accepted_now = int(draft_token) == target_token accept_prob = 1.0 if accepted_now else 0.0 correction = target_token + elif _cascade_active and _cascade_rule != "opt" and target_distribution_batch is not None: + # Token-specific speculative-cascade acceptance (arXiv:2405.19261 + # v2, Sec. 4.4 / Eq. 11,13-15; Appendix D Algorithm 6). Unlike + # OPT (Eq. 10) which compares only the peaks, r_TokenV{1,3} judge + # the drafted token v: accept (r=0) when v is in Top_alpha, else + # DEFER (r=1) and run the exact coin/residual with target + # pi_Token (Eq. 11) instead of p. pi_Token(v)=q(v) for v in + # Top_alpha gives an accept probability of 1, so Top_alpha tokens + # accept with no coin. + draft_q = draft_probs[depth_index] + if draft_q is None: + raise RuntimeError("non-greedy MTP requires draft distributions") + target_p_for_cache = target_distribution_batch.to_distribution( + depth_index + ) + _cas_tv = total_variation(target_p_for_cache, draft_q) + cascade_positions += 1 + cascade_divergence_sum += _cas_tv + _defer = cascade_token_deferral( + target_p_for_cache, draft_q, draft_token, + alpha=_cascade_alpha, rule=_cascade_rule, + ) + if not _defer: + accepted_now = True + accept_prob = 1.0 + correction = draft_token + cascade_accepted += 1 + else: + _pi_token = cascade_token_target_distribution( + target_p_for_cache, draft_q, + alpha=_cascade_alpha, rule=_cascade_rule, + ) + accept_prob = compute_acceptance_probability( + _pi_token, draft_q, draft_token + ) + accepted_now = float(rng.random()) <= accept_prob + if accepted_now: + correction = draft_token + cascade_accepted += 1 + else: + correction = sample_from_distribution( + residual_distribution(_pi_token, draft_q), rng + ) + cascade_resamples += 1 elif _cascade_active and target_distribution_batch is not None: # Speculative-cascade acceptance, batched-target rows # (arXiv:2405.19261, Eq. (10) + Algorithm 4). The deferral @@ -12489,7 +12564,38 @@ def emit_new_tokens() -> None: if draft_q is None: raise RuntimeError("non-greedy MTP requires draft distributions") target_p_for_cache = target_p - if _cascade_active: + if _cascade_active and _cascade_rule != "opt": + # Token-specific cascade (Sec. 4.4), lazy/per-row target. + _cas_tv = total_variation(target_p, draft_q) + cascade_positions += 1 + cascade_divergence_sum += _cas_tv + _defer = cascade_token_deferral( + target_p, draft_q, draft_token, + alpha=_cascade_alpha, rule=_cascade_rule, + ) + if not _defer: + accepted_now = True + accept_prob = 1.0 + correction = draft_token + cascade_accepted += 1 + else: + _pi_token = cascade_token_target_distribution( + target_p, draft_q, + alpha=_cascade_alpha, rule=_cascade_rule, + ) + accept_prob = compute_acceptance_probability( + _pi_token, draft_q, draft_token + ) + accepted_now = float(rng.random()) <= accept_prob + if accepted_now: + correction = draft_token + cascade_accepted += 1 + else: + correction = sample_from_distribution( + residual_distribution(_pi_token, draft_q), rng + ) + cascade_resamples += 1 + elif _cascade_active: # Speculative-cascade acceptance, lazy/per-row target. Same # law as the batched branch: deterministic Eq.(10) deferral, # accept the draft when it is good enough, otherwise the @@ -13836,6 +13942,7 @@ def emit_new_tokens() -> None: _cas_cycles = max(1, verify_calls) print( "[cascade-accept] NOT distribution-exact; " + f"rule={_cascade_rule} " f"threshold={_cascade_alpha:.4g} alpha={_cascade_alpha:.4g} " f"positions={cascade_positions} accepted={cascade_accepted} " f"resamples={cascade_resamples} accept_rate={_cas_rate:.4f} " diff --git a/mtplx/sampling.py b/mtplx/sampling.py index 1cf26e1d3..44bec0b09 100644 --- a/mtplx/sampling.py +++ b/mtplx/sampling.py @@ -337,8 +337,8 @@ def cascade_defer_decision( ) -> tuple[bool, float]: """Speculative-cascade plug-in deferral rule. - Narasimhan, Mreddy, Jitkrittum, Rawat, Kumar, "Faster Cascades via - Speculative Decoding", ICLR 2025 / arXiv:2405.19261 v2, Eq. (10) (the + Narasimhan, Jitkrittum, Rawat, Kim, Gupta, Menon, Kumar, "Faster Cascades + via Speculative Decoding", arXiv:2405.19261 v2 (2024), Eq. (10) (the plug-in approximation to the optimal deferral rule of Eq. (8)): r_OPT(x_ max_v q(v) < max_v p(v) - alpha * D_TV(p, q) @@ -368,6 +368,103 @@ def cascade_defer_decision( return (defer, tv) +def cascade_token_deferral( + target_p: Distribution, + draft_q: Distribution, + token_id: int, + *, + alpha: float, + rule: str, +) -> bool: + """Token-specific speculative-cascade deferral r(x_ q(v) < max_v' p(v') - alpha (Eq. 13) + r_TokenV3(x_ p(v) < max_v' p(v') * (1 - alpha) (Eq. 15) + + ``r = 1`` DEFERS (``v`` judged poor); ``r = 0`` ACCEPTS ``v`` (it is in + ``Top_alpha``). Higher ``alpha`` grows ``Top_alpha`` and defers fewer + tokens. (Eq. 14 / TokenV2 -- ``p(v) < max p - alpha`` -- is available via + ``rule="tokenv2"`` for completeness.) + """ + max_p = _peak_probability(target_p) + a = float(alpha) + if rule == "tokenv3": + return _probability(target_p, token_id) < max_p * (1.0 - a) + if rule == "tokenv1": + return _probability(draft_q, token_id) < max_p - a + if rule == "tokenv2": + return _probability(target_p, token_id) < max_p - a + raise ValueError(f"unknown token-specific cascade rule: {rule!r}") + + +def cascade_token_target_distribution( + target_p: Distribution, + draft_q: Distribution, + *, + alpha: float, + rule: str, +) -> Distribution: + """``pi_Token`` (Eq. 11) for the token-specific rule ``r_TokenV{1,2,3}``. + + arXiv:2405.19261 v2, Eq. (11) and Appendix D (Algorithm 6, TokenSpecCascade): + + pi_Token(v) = q(v) * (1 - r(x_= max_v' p(v') * (1 - alpha) }. + + The ``p(v)*eta`` term is present for EVERY ``v``: an accepted token + ``v in Top_alpha`` has ``pi(v) = q(v) + p(v)*eta >= q(v)``, so the generic + speculative coin (Algorithm 4) accepts it with probability 1; a deferred + token has ``pi(v) = p(v)*eta``. ``sum_v pi(v) = 1`` by construction. The + deferred exact coin/residual then runs with this ``pi`` as the target, + exactly the shipped ``min(1, pi/q)`` accept + ``norm(max(0, pi - q))`` + residual (Algorithm 6 = GenSpecSample(q, p, pi_Token)). + """ + max_p = _peak_probability(target_p) + a = float(alpha) + sparse = isinstance(target_p, SparseDistribution) or isinstance(draft_q, SparseDistribution) + if sparse and isinstance(target_p, SparseDistribution) and isinstance(draft_q, SparseDistribution): + token_ids = np.union1d(target_p.token_ids, draft_q.token_ids).astype(np.int64) + p = np.array([target_p.probability(int(t)) for t in token_ids], dtype=np.float64) + q = np.array([draft_q.probability(int(t)) for t in token_ids], dtype=np.float64) + vocab = _vocab_size(target_p) + else: + p = _as_dense(target_p) + q = _as_dense(draft_q) + token_ids = np.arange(p.shape[0], dtype=np.int64) + vocab = int(p.shape[0]) + if rule == "tokenv3": + defer = p < max_p * (1.0 - a) + elif rule == "tokenv1": + defer = q < max_p - a + elif rule == "tokenv2": + defer = p < max_p - a + else: + raise ValueError(f"unknown token-specific cascade rule: {rule!r}") + eta = float(q[defer].sum()) + # pi(v) = q(v)*(1 - r(v)) + p(v)*eta for every v. + pi = np.where(defer, 0.0, q) + p * eta + pi = np.where(np.isfinite(pi) & (pi > 0), pi, 0.0) + total = pi.sum() + if not np.isfinite(total) or total <= 0: + return target_p # degenerate; keep the coin well-defined + pi = pi / total + if sparse: + keep = pi > 0 + return SparseDistribution(token_ids[keep], pi[keep], vocab) + return pi + def sample_from_distribution(probs: Distribution, rng: np.random.Generator | None = None) -> int: rng = rng or np.random.default_rng() if isinstance(probs, SparseDistribution): diff --git a/mtplx/server/openai.py b/mtplx/server/openai.py index 77fa2eeeb..ab561555c 100644 --- a/mtplx/server/openai.py +++ b/mtplx/server/openai.py @@ -16523,11 +16523,20 @@ def _cascade_acceptance_health_payload() -> dict[str, Any]: typical_on = tv is not None and float(tv) > 0.0 except ValueError: typical_on = False + rule_raw = os.environ.get("MTPLX_FABLE_CASCADE_RULE") + rule_name = str(rule_raw).strip().lower() if rule_raw not in (None, "") else "opt" + rule_text = { + "opt": "defer iff max_q < max_p - alpha*D_TV(p,q); else accept draft (Eq. 10)", + "tokenv1": "defer token v iff q(v) < max_p - alpha; else accept (Eq. 13)", + "tokenv2": "defer token v iff p(v) < max_p - alpha; else accept (Eq. 14)", + "tokenv3": "defer token v iff p(v) < max_p*(1-alpha); else accept (Eq. 15)", + }.get(rule_name, rule_name) return { "enabled": enabled, "alpha": alpha, "threshold": alpha, - "rule": "defer iff max_q < max_p - alpha*D_TV(p,q); else accept draft", + "rule_name": rule_name, + "rule": rule_text, "divergence": "D_TV(p,q) = sum_v max(0, p(v)-q(v)) over scored top-k", "distribution_exact": not enabled, "mutually_exclusive_with_typical": True, @@ -36384,6 +36393,20 @@ def parse_args(argv: list[str] | None = None) -> argparse.Namespace: "overrides. See docs/perf/qwen38-cascade-acceptance.md." ), ) + parser.add_argument( + "--cascade-rule", + choices=("opt", "tokenv1", "tokenv2", "tokenv3"), + default=None, + help=( + "Which speculative-cascade deferral rule --cascade-threshold " + "applies (arXiv:2405.19261 v2). opt (default) = the position-level " + "peak rule (Eq. 10); tokenv1/tokenv2/tokenv3 = the token-specific " + "rules (Eq. 13/14/15, Sec. 4.4) that judge the drafted token and " + "defer to the exact coin with target pi_Token (Eq. 11). Ignored " + "unless --cascade-threshold is set. Environment: " + "MTPLX_FABLE_CASCADE_RULE, which this flag overrides." + ), + ) parser.add_argument( "--ngram-prewarm", metavar="auto|all|off|GiB", @@ -36923,6 +36946,10 @@ def parse_args(argv: list[str] | None = None) -> argparse.Namespace: # Unlike the typical delta, any set value (including 0) turns the lane # on, so only pass --cascade-threshold to enable it. os.environ["MTPLX_FABLE_CASCADE_THRESHOLD"] = str(float(args.cascade_threshold)) + if getattr(args, "cascade_rule", None) is not None: + # Flag beats env; generation.py resolves the rule per request. Only the + # deferral rule -- inert unless --cascade-threshold turns the lane on. + os.environ["MTPLX_FABLE_CASCADE_RULE"] = str(args.cascade_rule) # Fail loud on the mutually exclusive lossy verify rules, whether they were # set by flag (stamped just above) or already present in the environment. _typ_env = os.environ.get("MTPLX_FABLE_TYPICAL_THRESHOLD") diff --git a/tests/test_cascade_acceptance.py b/tests/test_cascade_acceptance.py index e647ebdc4..55e47e552 100644 --- a/tests/test_cascade_acceptance.py +++ b/tests/test_cascade_acceptance.py @@ -345,3 +345,120 @@ def test_served_path_emits_cascade_verdict_with_threshold_and_positions( assert float(m_thr.group(1)) == float(m_alpha.group(1)) # threshold == alpha assert int(m_pos.group(1)) > 0 assert int(m_pos.group(1)) == out.stats.cascade_positions + + +# =========================================================================== +# Token-specific speculative-cascade rules (arXiv:2405.19261 v2, Sec. 4.4; +# Eq. 11/13/14/15; Appendix D Algorithm 6). r_OPT (Eq. 10) compares only the +# peaks, so it accepts a poor drafted token when q is more peaked than p; the +# token-specific rules judge the drafted token itself. +# =========================================================================== +import numpy as _np + +from mtplx.sampling import ( + SparseDistribution as _SD, + cascade_defer_decision as _opt_defer, + cascade_token_deferral as _tok_defer, + cascade_token_target_distribution as _tok_pi, +) + + +def _p_dist(): + # max_p = 0.5 at token 1; token 3 is confidently-wrong under p (p=0.05). + return _SD(_np.array([1, 2, 3, 4]), _np.array([0.5, 0.3, 0.05, 0.15]), 5) + + +def _q_dist(): + # q is MORE peaked than p (max_q = 0.9 at token 3), the OPT failure case. + return _SD(_np.array([3, 1]), _np.array([0.9, 0.1]), 5) + + +def test_tokenv3_defers_confident_wrong_draft_that_opt_accepts(): + p, q = _p_dist(), _q_dist() + alpha = 0.2 + drafted = 3 # x_t ~ q, the peak of q, but poor under p + # OPT (Eq. 10): max_q=0.9 >= max_p=0.5 - alpha*D_TV -> does NOT defer -> accepts. + opt_defer, _ = _opt_defer(p, q, alpha=alpha) + assert opt_defer is False + # TokenV3 (Eq. 15): p(3)=0.05 < max_p*(1-alpha)=0.5*0.8=0.4 -> DEFERS. + assert _tok_defer(p, q, drafted, alpha=alpha, rule="tokenv3") is True + + +def test_tokenv3_accepts_token_in_top_alpha(): + p, q = _p_dist(), _q_dist() + alpha = 0.2 + # token 1: p(1)=0.5 >= 0.4 -> in Top_alpha -> r=0 -> accept (no defer). + assert _tok_defer(p, q, 1, alpha=alpha, rule="tokenv3") is False + # token 2: p(2)=0.3 < 0.4 -> deferred. + assert _tok_defer(p, q, 2, alpha=alpha, rule="tokenv3") is True + + +def test_tokenv3_target_distribution_matches_eq11(): + p, q = _p_dist(), _q_dist() + alpha = 0.2 # Top_alpha = {1}; eta = sum_{v not in Top} q(v) = q(3) = 0.9 + pi = _tok_pi(p, q, alpha=alpha, rule="tokenv3") + # pi(v) = q(v)*1[v in Top] + p(v)*eta + assert pi.probability(1) == _pytest.approx(0.1 + 0.5 * 0.9) # 0.55 + assert pi.probability(2) == _pytest.approx(0.3 * 0.9) # 0.27 + assert pi.probability(3) == _pytest.approx(0.05 * 0.9) # 0.045 + assert pi.probability(4) == _pytest.approx(0.15 * 0.9) # 0.135 + assert sum(pi.probability(v) for v in (1, 2, 3, 4)) == _pytest.approx(1.0) + + +def test_tokenv1_rule_uses_q_against_additive_band(): + p, q = _p_dist(), _q_dist() + alpha = 0.2 # Eq. 13: defer iff q(v) < max_p - alpha = 0.5 - 0.2 = 0.3 + assert _tok_defer(p, q, 3, alpha=alpha, rule="tokenv1") is False # q(3)=0.9 >= 0.3 + assert _tok_defer(p, q, 1, alpha=alpha, rule="tokenv1") is True # q(1)=0.1 < 0.3 + + +def test_unknown_rule_fails_loud(): + p, q = _p_dist(), _q_dist() + with _pytest.raises(ValueError): + _tok_defer(p, q, 1, alpha=0.2, rule="bogus") + + +def test_served_order_tokenv3_names_rule_and_defers(capsys, monkeypatch): + previous = _mx.default_device() + _mx.set_default_device(_mx.cpu) + try: + monkeypatch.setenv("MTPLX_FABLE_CASCADE_THRESHOLD", "0.5") + monkeypatch.setenv("MTPLX_FABLE_CASCADE_RULE", "tokenv3") + monkeypatch.setenv("MTPLX_BATCH_TARGET_ARRAYS", "1") + monkeypatch.delenv("MTPLX_FABLE_TYPICAL_THRESHOLD", raising=False) + model = _VerdictAcceptingMTPModel() + out = _generate_mtpk( + _verdict_runtime(model), + [0], + max_tokens=5, + sampler=_SamplerConfig(temperature=0.6, top_p=1.0, top_k=1), + speculative_depth=3, + mtp_history_policy="committed", + verify_strategy="batched", + stop_token_ids=set(), + ) + finally: + _mx.set_default_device(previous) + assert out.stats.cascade_accept_enabled is True + assert out.stats.cascade_positions > 0 + err = capsys.readouterr().err + lines = [ln for ln in err.splitlines() if "[cascade-accept]" in ln] + assert len(lines) == 1, lines + assert "rule=tokenv3" in lines[0] + + +def test_default_rule_is_opt(monkeypatch): + monkeypatch.delenv("MTPLX_FABLE_CASCADE_RULE", raising=False) + from mtplx.generation import _cascade_accept_rule + assert _cascade_accept_rule() == "opt" + + +def test_rule_selector_reads_env_at_use(monkeypatch): + from mtplx.generation import _cascade_accept_rule + monkeypatch.setenv("MTPLX_FABLE_CASCADE_RULE", "tokenv1") + assert _cascade_accept_rule() == "tokenv1" + monkeypatch.setenv("MTPLX_FABLE_CASCADE_RULE", "TokenV3") + assert _cascade_accept_rule() == "tokenv3" # normalised + monkeypatch.setenv("MTPLX_FABLE_CASCADE_RULE", "bogus") + with _pytest.raises(ValueError): + _cascade_accept_rule() diff --git a/tests/test_qwen4_block_verify.py b/tests/test_qwen4_block_verify.py index e826212b6..f39577cbe 100644 --- a/tests/test_qwen4_block_verify.py +++ b/tests/test_qwen4_block_verify.py @@ -323,15 +323,20 @@ def test_the_shipped_law_survives_verbatim_in_the_accept_loop(): assert "else target_distribution_batch.to_distribution(depth_index)," in loop # The accept coin is still drawn once per depth and compared with `<=`, # so arming block verification cannot shift the PCG64 stream. The exact law - # appears in four places now: the two shipped exact branches (batched + - # lazy) plus the two speculative-cascade DEFER paths (batched + lazy), which - # reuse the identical coin verbatim. The cascade paths are gated behind - # `_cascade_active` (MTPLX_FABLE_CASCADE_THRESHOLD, off by default and - # mutually exclusive with block verification's own lane), so with the - # cascade lane off the coin still fires exactly once per depth and the + # appears in six places now: the two shipped exact branches (batched + + # lazy), the two speculative-cascade OPT DEFER paths (batched + lazy), and + # the two token-specific cascade DEFER paths (batched + lazy, rule + # tokenv1/2/3), all of which reuse the identical coin verbatim. Every cascade + # path is gated behind `_cascade_active` (MTPLX_FABLE_CASCADE_THRESHOLD, off + # by default and mutually exclusive with block verification's own lane), and + # the OPT vs token-specific paths are mutually exclusive per request via + # `_cascade_rule`, so exactly one coin fires per depth and the # block-verification RNG guarantee is unchanged. - assert loop.count("accepted_now = float(rng.random()) <= accept_prob") == 4 - assert loop.count("rng.random()") == 4 + assert loop.count("accepted_now = float(rng.random()) <= accept_prob") == 6 + # Same six coin sites (the token-specific defer paths reuse the identical + # `float(rng.random()) <= accept_prob` draw); no other `rng.random()` call + # appears in the accept loop, so the per-depth draw count is unchanged. + assert loop.count("rng.random()") == 6 def test_every_block_verification_read_is_guarded(): From 58b3b0f3522c8c0c603b64330787318b9a88ca42 Mon Sep 17 00:00:00 2001 From: davidtai Date: Wed, 9 Sep 2026 13:10:22 -0500 Subject: [PATCH 19/30] docs(cascade): TokenV3 vs OPT/exact/typical 16K results + rule switch guidance Add a measured-results section to docs/perf/qwen38-cascade-acceptance.md: the OPT HumanEval strict grid by alpha (0.9024/0.8537/0.7805/0.7073; alpha 1.0 dnf, 2.0 not run), the TokenV3 decode tok/s grid by alpha (84.09/88.70/95.30/96.52/107.10), the TokenV3 alpha 0.95 quality point (0.9695 strict / 1.000 completed / 3.05% trunc at 107.10 tok/s), the exact and typical-0.09 baselines, and the equal-speed comparison at ~105 tok/s. Add the Sec. 4.4 mechanism paragraph (why the peak rule admits poor tokens as alpha rises and the token-specific rule does not). Extend the Switch section with the --cascade-rule selector (opt default for backward compatibility; tokenv1/2/3), noting TokenV3's alpha is a fraction of the target peak so its useful range is >= 0.9, and update the /health and verdict-line docs for the rule_name/rule fields and the rule= verdict field. A placeholder line marks the pending TokenV3 alpha 0.75 quality point. Docs only. --- docs/perf/qwen38-cascade-acceptance.md | 81 ++++++++++++++++++++++++-- 1 file changed, 76 insertions(+), 5 deletions(-) diff --git a/docs/perf/qwen38-cascade-acceptance.md b/docs/perf/qwen38-cascade-acceptance.md index af9de702b..cb00585df 100644 --- a/docs/perf/qwen38-cascade-acceptance.md +++ b/docs/perf/qwen38-cascade-acceptance.md @@ -112,14 +112,26 @@ runs. UNSETTING the key, not setting it to 0 (alpha = 0 still defers whenever the target is strictly more confident than the draft). Higher alpha widens the accept band, so fewer positions defer. Resolved at use, per request. +- `--cascade-rule opt|tokenv1|tokenv2|tokenv3` (env `MTPLX_FABLE_CASCADE_RULE`): + which deferral rule `--cascade-threshold` applies. Default `opt` for backward + compatibility -- the position-level peak rule (Eq. 10). `tokenv1`/`tokenv2`/ + `tokenv3` are the token-specific rules (Eq. 13/14/15) that judge the drafted + token. For `tokenv3`, alpha is a FRACTION of the target peak + (`Top_alpha = {v : p(v) >= max_p * (1 - alpha)}`), so its useful range is + alpha >= 0.9; a small alpha collapses `Top_alpha` toward the argmax and defers + almost everything (near-exact, slow). Ignored unless `--cascade-threshold` is + set; resolved at use, per request. - Mutually exclusive with `--typical-threshold` (`MTPLX_FABLE_TYPICAL_THRESHOLD` > 0). Setting both fails loud at serve startup (SystemExit) and in the verify setup (ValueError). -- Observability: `/health` -> `cascade_acceptance` (`enabled`, `alpha`, the rule - and divergence definitions, `conflict`); and a per-request verdict line - `[cascade-accept] NOT distribution-exact; threshold=A alpha=A positions=N - accepted=N resamples=N accept_rate=R mean_divergence=D ...` in the same format - as `[typical-accept]`. +- Observability: `/health` -> `cascade_acceptance` reports `enabled`, `alpha` + (= `threshold`), `rule_name` (the resolved rule, e.g. `tokenv3`), `rule` (its + one-line definition), the divergence definition, and `conflict`. Each + generation prints a per-request verdict line that names the rule: + `[cascade-accept] NOT distribution-exact; rule=RULE threshold=A alpha=A + positions=N accepted=N resamples=N accept_rate=R mean_divergence=D + tokens_per_cycle=T accepted_by_depth=... generated=N verify_calls=N` (the same + format as `[typical-accept]`, with the added `rule=` field). ## Recommended alpha grid for the 16K sweep @@ -280,3 +292,62 @@ docstring text changed. The measured arm-G OPT numbers stand. (The whole-file text hashes in the earlier table above use a different extractor and are unaffected; the two tables here are self-consistent, produced by one docstring-stripping AST script.) + +## Measured results (16K HumanEval) + +Optimized-Speed pack, 16,384-token context, HumanEval (164), temperature 1.0, +reasoning xhigh, sampled decode; decode tok/s is the fastest of the seeds. +"strict" pass@1 is the evalplus-sanitized score over all problems; "completed" +excludes output-cap truncations; "trunc" is the truncation rate. + +Baselines: + +- Exact speculative sampling (cascade off): HumanEval 0.9634 strict at 82.85 tok/s. +- Typical-0.09 (the #478 lane): 0.9634 strict at 104.65 tok/s. + +OPT rule (`r_OPT`, Eq. 10) -- HumanEval strict pass@1 by alpha: + +| alpha | 0.0 | 0.25 | 0.5 | 0.75 | 1.0 | 2.0 | +| --- | --- | --- | --- | --- | --- | --- | +| HE strict | 0.9024 | 0.8537 | 0.7805 | 0.7073 | dnf | not run | + +OPT alpha 0.25 ran at 104.78 tok/s (equal speed to typical-0.09) at 0.8537 strict. + +TokenV3 rule (`r_TokenV3`, Eq. 15) -- decode tok/s by alpha: + +| alpha | 0.0 | 0.25 | 0.5 | 0.75 | 0.95 | +| --- | --- | --- | --- | --- | --- | +| decode tok/s | 84.09 | 88.70 | 95.30 | 96.52 | 107.10 | + +(TokenV3 alpha 0.0 sits at 84.09 tok/s, next to exact's 82.85, because +`Top_alpha` is then the argmax alone and almost every token defers; higher alpha +widens `Top_alpha` and accepts more, up to 107.10 tok/s at alpha 0.95.) + +TokenV3 quality: + +- alpha 0.95: HumanEval 0.9695 strict / 1.000 completed / 3.05% truncation, at + 107.10 tok/s. +- alpha 0.75: {{tokenv3_a0p75_quality}} + +Equal-speed comparison (~105 tok/s): TokenV3 alpha 0.95 = 0.9695 strict at +107.10 tok/s; typical-0.09 = 0.9634 strict at 104.65 tok/s; OPT alpha 0.25 = +0.8537 strict at 104.78 tok/s. Exact speculative sampling reaches 0.9634 strict +but at 82.85 tok/s. + +### Why the peak rule loses and the token-specific rule does not (Sec. 4.4) + +`r_OPT` (Eq. 10) decides between the draft `q` and the target `p` by comparing +only their maximum token probabilities. As Sec. 4.4 of the paper states, the +drafted token `x_t ~ q_t(.)` may not maximise `q_t`, so "even when `x_t` is of +poor quality, we may end up accepting it because `q_t` happens to be more peaked +than `p_t`." That is the mechanism behind the OPT grid above: as alpha rises the +accept band widens and more of these poor tokens are admitted, and strict pass@1 +falls monotonically. The token-specific rules judge the drafted token `v` +itself: `r_TokenV3` defers `v` iff `p(v) < max_v' p(v') * (1 - alpha)`, i.e. iff +`v` is outside the target's top band `Top_alpha`. A confidently-wrong drafted +token has small `p(v)`, so it is deferred and re-drawn from the exact target +`pi_Token` (Eq. 11) even when `q` is more peaked than `p` -- exactly the case +`r_OPT` accepts. This is why TokenV3 holds strict pass@1 at exact's level +(0.9695 vs 0.9634) while still admitting the high-`p(v)` tokens that give the +speed-up. + From 8128e5ae4c7b87b92ed53e187c8429e35b0852bb Mon Sep 17 00:00:00 2001 From: davidtai Date: Wed, 9 Sep 2026 14:29:51 -0500 Subject: [PATCH 20/30] docs(cascade): fill the TokenV3 alpha 0.75 quality point Measured 16K HumanEval cell for TokenV3 at alpha 0.75: strict 0.9695 (159/164), completed-task 0.9876, truncation 1.83% (mean 2,652 tokens), cascade acceptance 0.820, at 95.30 tok/s (+15.1% vs exact 82.85), wall 1 h 19 m. Replaces the placeholder, and adds the closing numbers line (TokenV3 0.9695 at both alpha 0.75 and 0.95; OPT 0.9024 -> 0.7073). Docs only. --- docs/perf/qwen38-cascade-acceptance.md | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/docs/perf/qwen38-cascade-acceptance.md b/docs/perf/qwen38-cascade-acceptance.md index cb00585df..f62aeaf82 100644 --- a/docs/perf/qwen38-cascade-acceptance.md +++ b/docs/perf/qwen38-cascade-acceptance.md @@ -327,13 +327,18 @@ TokenV3 quality: - alpha 0.95: HumanEval 0.9695 strict / 1.000 completed / 3.05% truncation, at 107.10 tok/s. -- alpha 0.75: {{tokenv3_a0p75_quality}} +- alpha 0.75: HumanEval 0.9695 strict (159/164) / 0.9876 completed / 1.83% + truncation (mean 2,652 tokens), cascade acceptance 0.820, at 95.30 tok/s + (+15.1% vs exact 82.85); wall 1 h 19 m. Equal-speed comparison (~105 tok/s): TokenV3 alpha 0.95 = 0.9695 strict at 107.10 tok/s; typical-0.09 = 0.9634 strict at 104.65 tok/s; OPT alpha 0.25 = 0.8537 strict at 104.78 tok/s. Exact speculative sampling reaches 0.9634 strict but at 82.85 tok/s. +TokenV3 held 0.9695 strict at both alpha 0.75 and alpha 0.95; OPT ran 0.9024 -> +0.7073 across alpha 0.0 -> 0.75. + ### Why the peak rule loses and the token-specific rule does not (Sec. 4.4) `r_OPT` (Eq. 10) decides between the draft `q` and the target `p` by comparing From 4fec9aec3d1be9c31d578f323e1b094cb7b78174 Mon Sep 17 00:00:00 2001 From: davidtai Date: Wed, 9 Sep 2026 15:05:17 -0500 Subject: [PATCH 21/30] docs(cascade): add the TokenV3 alpha 0.95 series to the by-context chart cascade_decode_by_context now carries eight arms: the exact pairing, the OPT rule at alpha 0.0/0.25/0.5/0.75/1.0/2.0, and the TokenV3 rule at alpha 0.95 across 1K-128K (fastest of three seeds, min-max bands). 261,120 is absent because every arm exceeds the memory knob on the #475 base without #482. Chart and manifest only; nothing under mtplx/. --- .../charts/cascade_decode_by_context.svg | 250 +++++++++++++----- .../charts/manifest.json | 8 +- 2 files changed, 181 insertions(+), 77 deletions(-) diff --git a/docs/perf/qwen38-cascade-acceptance/charts/cascade_decode_by_context.svg b/docs/perf/qwen38-cascade-acceptance/charts/cascade_decode_by_context.svg index 3cfbd926c..51c7bba2b 100644 --- a/docs/perf/qwen38-cascade-acceptance/charts/cascade_decode_by_context.svg +++ b/docs/perf/qwen38-cascade-acceptance/charts/cascade_decode_by_context.svg @@ -192,6 +192,38 @@ z + + + + + + + + + + + + + + + + @@ -1350,6 +1382,34 @@ z + + + + + + + + + + + + + @@ -1456,21 +1516,21 @@ z - - + - - + - + - + - - + - + - + - - + - + - + - - + - + - + @@ -1700,18 +1760,18 @@ L 475.915 83.86125 - - + - + - + - - + - + - + @@ -1775,18 +1835,18 @@ L 475.915 107.8625 - - + - + - + @@ -1806,20 +1866,41 @@ L 475.915 119.863125 - - - - - - - + + + + + + + + + + - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - Date: Wed, 9 Sep 2026 15:54:56 -0500 Subject: [PATCH 22/30] cascade: default the rule to tokenv3 (exact-level accuracy) David's ruling (2026-09-09): TokenV3 is the only cascade rule with decent accuracy (HumanEval 0.9695 strict at alpha 0.95, equal to exact, while OPT loses at every alpha: 0.9024 at 0.0 down to 0.7073 at 0.75), so it becomes the default rule. Flip the default of MTPLX_FABLE_CASCADE_RULE / --cascade-rule from opt to tokenv3 everywhere the default is defined: the reader _cascade_accept_rule() (unset -> tokenv3), the /health cascade_acceptance.rule_name default, and the --cascade-rule help text. The argparse default stays the None sentinel so an unset flag does not override a shell-set env; the effective default resolves in the reader. opt, tokenv1 and tokenv2 stay selectable (env or flag). The mode itself is unchanged and still OFF by default: with the alpha knob unset the exact speculative law runs, so this changes only which rule engages once --cascade-threshold is set. The OPT rule body is untouched: the AST code hash (docstring stripped) of cascade_defer_decision is byte-identical to 2eac2fee/d8efc3f5. Tests: the default-selector test flips to tokenv3, a new test asserts env opt still selects OPT (backward compatibility), and the exact-off and block-verify guards still pass. Docs switch section updated with the new default, why, and how to select opt. --- docs/perf/qwen38-cascade-acceptance.md | 22 ++++++++++++++-------- mtplx/generation.py | 6 ++++-- mtplx/server/openai.py | 17 +++++++++++------ tests/test_cascade_acceptance.py | 12 +++++++++++- 4 files changed, 40 insertions(+), 17 deletions(-) diff --git a/docs/perf/qwen38-cascade-acceptance.md b/docs/perf/qwen38-cascade-acceptance.md index f62aeaf82..41dc10502 100644 --- a/docs/perf/qwen38-cascade-acceptance.md +++ b/docs/perf/qwen38-cascade-acceptance.md @@ -113,14 +113,20 @@ runs. target is strictly more confident than the draft). Higher alpha widens the accept band, so fewer positions defer. Resolved at use, per request. - `--cascade-rule opt|tokenv1|tokenv2|tokenv3` (env `MTPLX_FABLE_CASCADE_RULE`): - which deferral rule `--cascade-threshold` applies. Default `opt` for backward - compatibility -- the position-level peak rule (Eq. 10). `tokenv1`/`tokenv2`/ - `tokenv3` are the token-specific rules (Eq. 13/14/15) that judge the drafted - token. For `tokenv3`, alpha is a FRACTION of the target peak - (`Top_alpha = {v : p(v) >= max_p * (1 - alpha)}`), so its useful range is - alpha >= 0.9; a small alpha collapses `Top_alpha` toward the argmax and defers - almost everything (near-exact, slow). Ignored unless `--cascade-threshold` is - set; resolved at use, per request. + which deferral rule `--cascade-threshold` applies. Default `tokenv3` (the + token-specific multiplicative rule, Eq. 15). It is the default because it is + the only rule measured at exact-level accuracy: at alpha 0.95 it holds + HumanEval 0.9695 strict (equal to exact) while OPT loses quality at every alpha + (0.9024 at alpha 0.0 down to 0.7073 at alpha 0.75). To use the original + position-level peak rule, select `opt` (env `MTPLX_FABLE_CASCADE_RULE=opt` or + `--cascade-rule opt`); `tokenv1`/`tokenv2` are the additive token-specific + rules (Eq. 13/14). The default change does not turn the mode on: the lane is + still OFF unless `--cascade-threshold` is set, in which case the exact + speculative law runs unchanged. For `tokenv3`, alpha is a FRACTION of the + target peak (`Top_alpha = {v : p(v) >= max_p * (1 - alpha)}`), so its useful + range is alpha >= 0.9; a small alpha collapses `Top_alpha` toward the argmax + and defers almost everything (near-exact, slow). All rules are resolved at use, + per request. - Mutually exclusive with `--typical-threshold` (`MTPLX_FABLE_TYPICAL_THRESHOLD` > 0). Setting both fails loud at serve startup (SystemExit) and in the verify setup (ValueError). diff --git a/mtplx/generation.py b/mtplx/generation.py index d5508dfa6..b5b57247c 100644 --- a/mtplx/generation.py +++ b/mtplx/generation.py @@ -639,7 +639,9 @@ def _cascade_accept_enabled() -> bool: def _cascade_accept_rule() -> str: """MTPLX_FABLE_CASCADE_RULE (server flag --cascade-rule): which cascade - deferral rule to apply. Default "opt" for backward compatibility. + deferral rule to apply. Default "tokenv3" (David 2026-09-09: the only + rule with exact-level accuracy); "opt" stays selectable for backward + compatibility. The mode itself is still OFF unless the alpha knob is set. opt -- r_OPT (arXiv:2405.19261 v2, Eq. 10): defer iff max_v q(v) < max_v p(v) - alpha*D_TV(p,q); position-level. @@ -654,7 +656,7 @@ def _cascade_accept_rule() -> str: """ raw = os.environ.get("MTPLX_FABLE_CASCADE_RULE") if raw is None or str(raw).strip() == "": - return "opt" + return "tokenv3" rule = str(raw).strip().lower() if rule not in ("opt", "tokenv1", "tokenv2", "tokenv3"): raise ValueError( diff --git a/mtplx/server/openai.py b/mtplx/server/openai.py index ab561555c..15351125c 100644 --- a/mtplx/server/openai.py +++ b/mtplx/server/openai.py @@ -16524,7 +16524,7 @@ def _cascade_acceptance_health_payload() -> dict[str, Any]: except ValueError: typical_on = False rule_raw = os.environ.get("MTPLX_FABLE_CASCADE_RULE") - rule_name = str(rule_raw).strip().lower() if rule_raw not in (None, "") else "opt" + rule_name = str(rule_raw).strip().lower() if rule_raw not in (None, "") else "tokenv3" rule_text = { "opt": "defer iff max_q < max_p - alpha*D_TV(p,q); else accept draft (Eq. 10)", "tokenv1": "defer token v iff q(v) < max_p - alpha; else accept (Eq. 13)", @@ -36396,14 +36396,19 @@ def parse_args(argv: list[str] | None = None) -> argparse.Namespace: parser.add_argument( "--cascade-rule", choices=("opt", "tokenv1", "tokenv2", "tokenv3"), + # default=None (sentinel) so an unset flag does not override a + # shell-set MTPLX_FABLE_CASCADE_RULE; the effective default resolves in + # generation._cascade_accept_rule() (tokenv3). default=None, help=( "Which speculative-cascade deferral rule --cascade-threshold " - "applies (arXiv:2405.19261 v2). opt (default) = the position-level " - "peak rule (Eq. 10); tokenv1/tokenv2/tokenv3 = the token-specific " - "rules (Eq. 13/14/15, Sec. 4.4) that judge the drafted token and " - "defer to the exact coin with target pi_Token (Eq. 11). Ignored " - "unless --cascade-threshold is set. Environment: " + "applies (arXiv:2405.19261 v2). tokenv3 (default) = the " + "token-specific multiplicative rule (Eq. 15, Sec. 4.4), the only " + "rule with exact-level accuracy; opt = the position-level peak " + "rule (Eq. 10); tokenv1/tokenv2 = the additive token-specific " + "rules (Eq. 13/14). The token-specific rules judge the drafted " + "token and defer to the exact coin with target pi_Token (Eq. 11). " + "Ignored unless --cascade-threshold is set. Environment: " "MTPLX_FABLE_CASCADE_RULE, which this flag overrides." ), ) diff --git a/tests/test_cascade_acceptance.py b/tests/test_cascade_acceptance.py index 55e47e552..99dda4cab 100644 --- a/tests/test_cascade_acceptance.py +++ b/tests/test_cascade_acceptance.py @@ -447,10 +447,20 @@ def test_served_order_tokenv3_names_rule_and_defers(capsys, monkeypatch): assert "rule=tokenv3" in lines[0] -def test_default_rule_is_opt(monkeypatch): +def test_default_rule_is_tokenv3(monkeypatch): + # David 2026-09-09: TokenV3 is the default rule (exact-level accuracy). monkeypatch.delenv("MTPLX_FABLE_CASCADE_RULE", raising=False) from mtplx.generation import _cascade_accept_rule + assert _cascade_accept_rule() == "tokenv3" + + +def test_env_opt_still_selects_opt(monkeypatch): + # Backward compatibility: the OPT peak rule stays selectable by env. + from mtplx.generation import _cascade_accept_rule + monkeypatch.setenv("MTPLX_FABLE_CASCADE_RULE", "opt") assert _cascade_accept_rule() == "opt" + monkeypatch.setenv("MTPLX_FABLE_CASCADE_RULE", "OPT") + assert _cascade_accept_rule() == "opt" # normalised def test_rule_selector_reads_env_at_use(monkeypatch): From 67c8a75b7c7ac3a315097a56db40a061b21d66fa Mon Sep 17 00:00:00 2001 From: davidtai Date: Wed, 9 Sep 2026 16:01:25 -0500 Subject: [PATCH 23/30] docs(cascade): add the #475/#478 comparison charts Two charts comparing this pull request's operating points against #475 and #478, both re-derived from the receipt json rather than from any table: cascade_vs_475_478_16k.svg decode at 16,384 tokens for release 2.11.2 (exact), #475 (exact), #478 (typical 0.09), OPT alpha 0.25, and TokenV3 alpha 0.75 and 0.95, each bar annotated with its own HumanEval strict pass@1 cascade_vs_475_478_ladder.svg decode against context size, 1,024 to 131,072, for the same four arms that have a full ladder Bar height and line point are the fastest seed, every band is min-max, and acceptance mode is in each label: a tok/s figure is not readable without it. The alphas of the two rules are different quantities and the captions say so. No arm carries a 261,120 point, because on the #475 base without #482 every arm on this pack exceeds the memory knob. --- .../charts/cascade_vs_475_478_16k.svg | 3599 +++++++++++++++++ .../charts/cascade_vs_475_478_ladder.svg | 3316 +++++++++++++++ .../charts/manifest.json | 18 + 3 files changed, 6933 insertions(+) create mode 100644 docs/perf/qwen38-cascade-acceptance/charts/cascade_vs_475_478_16k.svg create mode 100644 docs/perf/qwen38-cascade-acceptance/charts/cascade_vs_475_478_ladder.svg diff --git a/docs/perf/qwen38-cascade-acceptance/charts/cascade_vs_475_478_16k.svg b/docs/perf/qwen38-cascade-acceptance/charts/cascade_vs_475_478_16k.svg new file mode 100644 index 000000000..775faa77d --- /dev/null +++ b/docs/perf/qwen38-cascade-acceptance/charts/cascade_vs_475_478_16k.svg @@ -0,0 +1,3599 @@ + + + + + + + + 2026-09-09T16:02:57.596299 + image/svg+xml + + + Matplotlib v3.11.0, https://matplotlib.org/ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/docs/perf/qwen38-cascade-acceptance/charts/cascade_vs_475_478_ladder.svg b/docs/perf/qwen38-cascade-acceptance/charts/cascade_vs_475_478_ladder.svg new file mode 100644 index 000000000..fe9f80437 --- /dev/null +++ b/docs/perf/qwen38-cascade-acceptance/charts/cascade_vs_475_478_ladder.svg @@ -0,0 +1,3316 @@ + + + + + + + + 2026-09-09T16:02:57.746972 + image/svg+xml + + + Matplotlib v3.11.0, https://matplotlib.org/ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/docs/perf/qwen38-cascade-acceptance/charts/manifest.json b/docs/perf/qwen38-cascade-acceptance/charts/manifest.json index f04a292ad..57a518411 100644 --- a/docs/perf/qwen38-cascade-acceptance/charts/manifest.json +++ b/docs/perf/qwen38-cascade-acceptance/charts/manifest.json @@ -25,5 +25,23 @@ "key": "cascade_decode_by_context", "label": "cascade decode by context", "sha256": "2391b851be197e57ab1b3f9746c61624340fc65d10b8742ca75463df02404fa0" + }, + { + "alt": "Grouped bars of decode tok/s at 16,384 tokens for release 2.11.2, #475, #478 and four #485 cascade operating points, each annotated with its HumanEval strict pass@1", + "bytes": 141088, + "caption": "_Decode tok/s at 16,384 tokens on the Optimized-Speed pack: release 2.11.2 at the exact speculative law (71.97), #475 at the exact law (82.84), #478 at typical acceptance 0.09 (104.65), and #485 at the OPT rule (Equation 10) alpha 0.25 (104.78) and the TokenV3 rule (Equation 15) alpha 0.75 (95.30) and 0.95 (107.10). Bar height is the fastest seed, the whisker is the min-max band, and the figure above each whisker is HumanEval strict pass@1 for the same arm. Alpha is NOT comparable between the two rules; compare at equal decode speed._", + "file": "cascade_vs_475_478_16k.svg", + "key": "cascade_vs_475_478_16k", + "label": "cascade against #475 and #478 at 16K", + "sha256": "393363aa89e48508f653a5a76718f2bbc94611f2a8cf6e1bbd51dfa171772ce5" + }, + { + "alt": "Decode tok/s against context size from 1,024 to 131,072 tokens for #475 exact, #478 typical 0.09, #485 OPT alpha 0.25 and #485 TokenV3 alpha 0.95", + "bytes": 127817, + "caption": "_Decode tok/s against context size on the Optimized-Speed pack for #475 at the exact speculative law, #478 at typical acceptance 0.09, #485 at the OPT rule alpha 0.25 and #485 at the TokenV3 rule alpha 0.95. Each point is the fastest seed and the band is the min-max range. No arm has a 261,120 point: on the #475 base without #482 every arm on this pack exceeds the memory knob._", + "file": "cascade_vs_475_478_ladder.svg", + "key": "cascade_vs_475_478_ladder", + "label": "cascade against #475 and #478 by context", + "sha256": "4743c0ef31d41736e771b1d082e9641050cf6822c213fc9f3e2671f0d8f9e324" } ] From 71b52ea6ebd3ded4f693ce664fd07039697fb13e Mon Sep 17 00:00:00 2001 From: davidtai Date: Wed, 9 Sep 2026 16:16:22 -0500 Subject: [PATCH 24/30] cascade: add cascade_deferred + defer_rate (the paper's deferral rate) David: accept rate is never defined for speculative cascade. The served verdict's accept_rate = cascade_accepted / cascade_positions is a KEPT-DRAFT rate over cascade-decided positions (no-defer accepts plus coin-accepted deferred tokens), not the paper's deferral rate r. Add a cascade_deferred counter, incremented on every deferral (the _defer / OPT defer branch, before the coin) at all four verify sites (OPT and TokenV3, batched and lazy). Expose defer_rate = cascade_deferred / cascade_positions in VerifyStats (cascade_deferred, cascade_defer_rate), the [cascade-accept] verdict line (deferred=, defer_rate=), and the /health cascade_acceptance payload (documented rates block). accept_rate is unchanged, now documented as the kept-draft rate; accept_rate + resample_rate == 1 over decided positions, while defer_rate is independent. Rule bodies untouched: the docstring-stripped AST hashes of cascade_defer_decision, total_variation and _peak_probability are byte-identical to 2eac2fee/d8efc3f5. Tests: defer/accept/resample consistency (accepted + resamples == positions; resamples <= deferred <= positions; defer_rate == deferred/positions; deferred>0 for a divergent draft under TokenV3), served-order emission of defer_rate, and exact-off leaves the cascade counters zero. docs/perf gains a Definitions block giving the three quantities, their formulas, and which artifact each comes from (verdict line / VerifyStats vs the receipt's mtp_accept_rate). --- docs/perf/qwen38-cascade-acceptance.md | 30 +++++++ mtplx/generation.py | 19 ++++- mtplx/server/openai.py | 5 ++ tests/test_cascade_acceptance.py | 103 +++++++++++++++++++++++++ 4 files changed, 156 insertions(+), 1 deletion(-) diff --git a/docs/perf/qwen38-cascade-acceptance.md b/docs/perf/qwen38-cascade-acceptance.md index 41dc10502..0dddea194 100644 --- a/docs/perf/qwen38-cascade-acceptance.md +++ b/docs/perf/qwen38-cascade-acceptance.md @@ -104,6 +104,36 @@ runs. `cascade_acceptance` install report. - `tests/test_cascade_acceptance.py`, `tests/test_cascade_threshold_cli_health_cpu.py`. +## Definitions: deferral, accept, and resample rates + +Speculative cascade has no single "accept rate": the paper's quantity is the +DEFERRAL rate, and the served verdict's `accept_rate` is a different, kept-draft +quantity. Over the positions the cascade rule actually decided +(`cascade_positions`, one per drafted position in a cascade window), the three +rates are: + +| quantity | formula | meaning | artifact | +| --- | --- | --- | --- | +| defer_rate | `cascade_deferred / cascade_positions` | the paper's deferral rate r: the fraction of decided positions where the rule deferred to the target (before the coin) | `[cascade-accept]` verdict line `defer_rate=`; `VerifyStats.cascade_defer_rate` | +| accept_rate | `cascade_accepted / cascade_positions` | kept-draft rate: no-defer accepts plus deferred tokens the exact coin then kept; NOT the deferral rate | `[cascade-accept]` verdict line `accept_rate=`; `VerifyStats.cascade_accepted` | +| resample_rate | `cascade_resamples / cascade_positions` | deferred tokens that lost the exact coin and were resampled from the residual | `VerifyStats.cascade_resamples` | + +Counter definitions: `cascade_positions` counts every draft position the rule +decided; `cascade_deferred` counts positions where the rule deferred (`r = 1`, +incremented before the coin); `cascade_accepted` counts kept drafts (no-defer +accepts plus coin-accepted deferred tokens); `cascade_resamples` counts deferred +tokens that lost the coin. Two identities always hold: +`cascade_accepted + cascade_resamples == cascade_positions` and +`cascade_deferred == (coin-accepted deferred tokens) + cascade_resamples`, so +`accept_rate + resample_rate == 1` over decided positions, while `defer_rate` is +independent of them. + +Note on the receipt: the perf harness records a separate `mtp_accept_rate` (the +fraction of drafted tokens accepted across the whole request, the throughput +lever). That is a decode-throughput measure over all MTP positions and is not +the cascade `defer_rate`; read `defer_rate` from the `[cascade-accept]` verdict +line or `VerifyStats.cascade_defer_rate`, not from `mtp_accept_rate`. + ## Switch - `--cascade-threshold ALPHA` (env `MTPLX_FABLE_CASCADE_THRESHOLD`): the deferral diff --git a/mtplx/generation.py b/mtplx/generation.py index b5b57247c..8a82e6aa8 100644 --- a/mtplx/generation.py +++ b/mtplx/generation.py @@ -2770,6 +2770,8 @@ class GenerationStats: cascade_positions: int = 0 cascade_accepted: int = 0 cascade_resamples: int = 0 + cascade_deferred: int = 0 + cascade_defer_rate: float = 0.0 cascade_mean_divergence: float = 0.0 # Which commit path produced the stop token when finish_reason == "stop" # (#414 telemetry): accepted_draft | residual_correction | bonus | @@ -8845,6 +8847,7 @@ def _steer_overlay(working: Sequence[int]) -> dict[int, float] | None: # resampled from the exact residual); cascade_divergence_sum feeds the mean # total-variation divergence on the verdict line. cascade_positions = cascade_accepted = cascade_resamples = 0 + cascade_deferred = 0 # positions where the rule DEFERRED (paper's r) cascade_divergence_sum = 0.0 stop_origin: str | None = None accepted_by_depth = [0 for _ in range(speculative_depth)] @@ -12428,6 +12431,7 @@ def emit_new_tokens() -> None: correction = draft_token cascade_accepted += 1 else: + cascade_deferred += 1 _pi_token = cascade_token_target_distribution( target_p_for_cache, draft_q, alpha=_cascade_alpha, rule=_cascade_rule, @@ -12469,6 +12473,7 @@ def emit_new_tokens() -> None: correction = draft_token cascade_accepted += 1 else: + cascade_deferred += 1 # Defer (pi = p): exact speculative law, same as the exact # branch below. p = target_distribution_batch.probability(depth_index, draft_token) @@ -12581,6 +12586,7 @@ def emit_new_tokens() -> None: correction = draft_token cascade_accepted += 1 else: + cascade_deferred += 1 _pi_token = cascade_token_target_distribution( target_p, draft_q, alpha=_cascade_alpha, rule=_cascade_rule, @@ -12613,6 +12619,7 @@ def emit_new_tokens() -> None: correction = draft_token cascade_accepted += 1 else: + cascade_deferred += 1 accept_prob = compute_acceptance_probability( target_p, draft_q, draft_token ) @@ -13832,6 +13839,12 @@ def emit_new_tokens() -> None: cascade_positions=int(cascade_positions), cascade_accepted=int(cascade_accepted), cascade_resamples=int(cascade_resamples), + cascade_deferred=int(cascade_deferred), + cascade_defer_rate=( + float(cascade_deferred / cascade_positions) + if cascade_positions + else 0.0 + ), cascade_mean_divergence=( float(cascade_divergence_sum / cascade_positions) if cascade_positions @@ -13941,13 +13954,17 @@ def emit_new_tokens() -> None: if _cascade_active: _cas_denom = cascade_accepted + cascade_resamples _cas_rate = (cascade_accepted / _cas_denom) if _cas_denom else 0.0 + _cas_defer_rate = ( + (cascade_deferred / cascade_positions) if cascade_positions else 0.0 + ) _cas_cycles = max(1, verify_calls) print( "[cascade-accept] NOT distribution-exact; " f"rule={_cascade_rule} " f"threshold={_cascade_alpha:.4g} alpha={_cascade_alpha:.4g} " f"positions={cascade_positions} accepted={cascade_accepted} " - f"resamples={cascade_resamples} accept_rate={_cas_rate:.4f} " + f"resamples={cascade_resamples} deferred={cascade_deferred} " + f"accept_rate={_cas_rate:.4f} defer_rate={_cas_defer_rate:.4f} " f"mean_divergence={stats.cascade_mean_divergence:.4f} " f"tokens_per_cycle={len(tokens) / _cas_cycles:.3f} " f"accepted_by_depth={accepted_by_depth} " diff --git a/mtplx/server/openai.py b/mtplx/server/openai.py index 15351125c..9873abadc 100644 --- a/mtplx/server/openai.py +++ b/mtplx/server/openai.py @@ -16538,6 +16538,11 @@ def _cascade_acceptance_health_payload() -> dict[str, Any]: "rule_name": rule_name, "rule": rule_text, "divergence": "D_TV(p,q) = sum_v max(0, p(v)-q(v)) over scored top-k", + "rates": { + "defer_rate": "cascade_deferred / cascade_positions -- the paper's deferral rate r (fraction of cascade-decided positions the rule deferred to the target); per-request value in VerifyStats.cascade_defer_rate and the [cascade-accept] verdict line", + "accept_rate": "cascade_accepted / cascade_positions -- kept-draft rate over cascade-decided positions (no-defer accepts + coin-accepted deferred tokens); NOT the deferral rate", + "resample_rate": "cascade_resamples / cascade_positions -- deferred tokens that lost the exact coin; accept_rate + resample_rate == 1 over cascade-decided positions", + }, "distribution_exact": not enabled, "mutually_exclusive_with_typical": True, "conflict": bool(enabled and typical_on), diff --git a/tests/test_cascade_acceptance.py b/tests/test_cascade_acceptance.py index 99dda4cab..52881ffee 100644 --- a/tests/test_cascade_acceptance.py +++ b/tests/test_cascade_acceptance.py @@ -472,3 +472,106 @@ def test_rule_selector_reads_env_at_use(monkeypatch): monkeypatch.setenv("MTPLX_FABLE_CASCADE_RULE", "bogus") with _pytest.raises(ValueError): _cascade_accept_rule() + + +# =========================================================================== +# defer_rate / accept_rate / resample_rate counters (David 2026-09-09: "accept +# rate is never defined for speculative cascade"). cascade_deferred counts +# positions the rule deferred (the paper's r); accept_rate is the kept-draft +# rate. Invariants: accepted + resamples == positions; resamples <= deferred +# <= positions; defer_rate == deferred / positions. +# =========================================================================== +class _DivergentMTPModel(_VerdictAcceptingMTPModel): + """Draft favours token 3, trunk target favours token 1, so under TokenV3 the + drafted token has small p(v) and the rule DEFERS (exercises cascade_deferred).""" + + def mtp_forward(self, hidden_states, next_token_ids, *, mtp_cache=None, + concat_order=None, return_hidden=False, + mtp_hidden_variant=None, position_offset=None): + length = int(next_token_ids.shape[1]) + hidden = _mx.zeros((1, length, 2), dtype=_mx.float32) + logits = _mx.zeros((1, length, 4), dtype=_mx.float32) + _mx.array( + [0.0, 0.0, 0.0, 6.0], dtype=_mx.float32) # draft -> token 3 + return (logits, hidden) if return_hidden else logits + + def __call__(self, input_ids, *, cache=None, return_hidden=False, + hidden_variant=None, emit_logits=True, logits_keep=None): + self.calls.append(int(input_ids.shape[1])) + length = int(input_ids.shape[1]) + hidden = _mx.zeros((1, length, 2), dtype=_mx.float32) + if not emit_logits: + return (None, hidden) if return_hidden else None + keep = length if logits_keep is None else min(length, max(1, int(logits_keep))) + logits = _mx.zeros((1, keep, 4), dtype=_mx.float32) + _mx.array( + [0.0, 6.0, 0.0, 0.0], dtype=_mx.float32) # target -> token 1 + return (logits, hidden) if return_hidden else logits + + +def _run_cascade(model, rule, alpha, monkeypatch, seed=0): + monkeypatch.setenv("MTPLX_FABLE_CASCADE_THRESHOLD", str(alpha)) + monkeypatch.setenv("MTPLX_FABLE_CASCADE_RULE", rule) + monkeypatch.setenv("MTPLX_BATCH_TARGET_ARRAYS", "1") + monkeypatch.delenv("MTPLX_FABLE_TYPICAL_THRESHOLD", raising=False) + previous = _mx.default_device() + _mx.set_default_device(_mx.cpu) + try: + return _generate_mtpk( + _verdict_runtime(model), [0], max_tokens=6, + sampler=_SamplerConfig(temperature=0.6, top_p=1.0, top_k=1), + speculative_depth=3, mtp_history_policy="committed", + verify_strategy="batched", stop_token_ids=set(), seed=seed, + ) + finally: + _mx.set_default_device(previous) + + +def test_defer_accept_resample_counters_are_consistent(monkeypatch): + out = _run_cascade(_DivergentMTPModel(), "tokenv3", 0.5, monkeypatch, seed=7) + st = out.stats + assert st.cascade_positions > 0 + # kept-draft identity: accepted + resamples == positions + assert st.cascade_accepted + st.cascade_resamples == st.cascade_positions + # deferred = coin-accepted-after-defer + resamples, so resamples <= deferred <= positions + assert st.cascade_resamples <= st.cascade_deferred <= st.cascade_positions + # no-defer accepts = positions - deferred, all kept, so accepted >= positions - deferred + assert st.cascade_accepted >= st.cascade_positions - st.cascade_deferred + # defer_rate is deferred / positions + assert st.cascade_defer_rate == _pytest.approx( + st.cascade_deferred / st.cascade_positions) + # the divergent draft is outside Top_alpha under TokenV3, so some positions defer + assert st.cascade_deferred > 0 + + +def test_served_order_emits_defer_rate(capsys, monkeypatch): + out = _run_cascade(_DivergentMTPModel(), "tokenv3", 0.5, monkeypatch, seed=7) + err = capsys.readouterr().err + lines = [ln for ln in err.splitlines() if "[cascade-accept]" in ln] + assert len(lines) == 1, lines + m_def = _re.search(r"defer_rate=([0-9.]+)", lines[0]) + m_pos = _re.search(r"positions=(\d+)", lines[0]) + m_dfd = _re.search(r"deferred=(\d+)", lines[0]) + assert m_def and m_pos and m_dfd, lines[0] + assert int(m_dfd.group(1)) == out.stats.cascade_deferred + assert float(m_def.group(1)) == _pytest.approx( + out.stats.cascade_deferred / out.stats.cascade_positions, abs=1e-4) + + +def test_exact_off_leaves_cascade_counters_zero(monkeypatch): + monkeypatch.delenv("MTPLX_FABLE_CASCADE_THRESHOLD", raising=False) + monkeypatch.delenv("MTPLX_FABLE_TYPICAL_THRESHOLD", raising=False) + monkeypatch.setenv("MTPLX_BATCH_TARGET_ARRAYS", "1") + previous = _mx.default_device() + _mx.set_default_device(_mx.cpu) + try: + out = _generate_mtpk( + _verdict_runtime(_VerdictAcceptingMTPModel()), [0], max_tokens=5, + sampler=_SamplerConfig(temperature=0.6, top_p=1.0, top_k=1), + speculative_depth=3, mtp_history_policy="committed", + verify_strategy="batched", stop_token_ids=set(), + ) + finally: + _mx.set_default_device(previous) + assert out.stats.cascade_accept_enabled is False + assert out.stats.cascade_positions == 0 + assert out.stats.cascade_deferred == 0 + assert out.stats.cascade_defer_rate == 0.0 From 2b41ad6bc3f4480e6910f0e9a6b26acc0842c2f1 Mon Sep 17 00:00:00 2001 From: davidtai Date: Wed, 9 Sep 2026 16:25:57 -0500 Subject: [PATCH 25/30] docs(cascade): restore the 16,384 point on the TokenV3 by-context line The TokenV3 series in cascade_decode_by_context.svg had a hole exactly at 16,384. The by-context renderer merged the 16K rung from battery475/g only, but the two rules were swept separately: the arm-T (TokenV3) 16,384 windows live in battery475/t, and the context ladder appended only the non-16K rungs for those arms. With no row to merge, the point was silently dropped. render_cascade_ladder.py now stages t/manifest.tsv alongside g/, reads past that manifest's leading comment line, and admits rc=65 rows, which are the arm-T gate-regex false negative on valid measured windows. The rc gate itself stays: ladder-cascade carries an rc=1 duplicate row whose receipt_dir points at the 65,536 runroot under a 131,072 ctx, and dropping the gate entirely would let that mis-aliased row through the ctx cross-check. The TokenV3 line now carries six points, 1K through 128K, with 16,384 at 107.10 tok/s (96.70-107.10, n=3), the same fastest-of-seeds and min-max rule as every other point. No other series or value changes. --- .../charts/cascade_decode_by_context.svg | 31 +++++++------------ .../charts/manifest.json | 6 ++-- 2 files changed, 15 insertions(+), 22 deletions(-) diff --git a/docs/perf/qwen38-cascade-acceptance/charts/cascade_decode_by_context.svg b/docs/perf/qwen38-cascade-acceptance/charts/cascade_decode_by_context.svg index 51c7bba2b..f8b483b2b 100644 --- a/docs/perf/qwen38-cascade-acceptance/charts/cascade_decode_by_context.svg +++ b/docs/perf/qwen38-cascade-acceptance/charts/cascade_decode_by_context.svg @@ -33,10 +33,10 @@ z - - + @@ -194,22 +194,10 @@ z - - - - - - - - - - + @@ -1385,7 +1376,8 @@ z @@ -1405,6 +1397,7 @@ z + diff --git a/docs/perf/qwen38-cascade-acceptance/charts/manifest.json b/docs/perf/qwen38-cascade-acceptance/charts/manifest.json index 57a518411..d4c50e0a3 100644 --- a/docs/perf/qwen38-cascade-acceptance/charts/manifest.json +++ b/docs/perf/qwen38-cascade-acceptance/charts/manifest.json @@ -19,12 +19,12 @@ }, { "alt": "Decode tok/s by context size for the exact pairing, cascade OPT alphas 0.0-2.0 and TokenV3 alpha 0.95", - "bytes": 88972, - "caption": "_Decode tok/s by context size, one line per arm: the exact pairing, the OPT rule at alpha 0.0/0.25/0.5/0.75/1.0/2.0, and the TokenV3 rule at alpha 0.95. Each point is the fastest seed with a min-max band; 16,384 is the arm-G/arm-T sweep merged in; 261,120 is absent because every arm exceeds the memory knob on the #475 base without #482._", + "bytes": 88853, + "caption": "_Decode tok/s by context size, one line per arm: the exact pairing, the OPT rule at alpha 0.0/0.25/0.5/0.75/1.0/2.0, and the TokenV3 rule at alpha 0.95. Each point is the fastest seed with a min-max band; the 16,384 rung is merged in from the 16K sweeps, the OPT arms from battery475/g and the TokenV3 arm from battery475/t; 261,120 is absent because every arm exceeds the memory knob on the #475 base without #482._", "file": "cascade_decode_by_context.svg", "key": "cascade_decode_by_context", "label": "cascade decode by context", - "sha256": "2391b851be197e57ab1b3f9746c61624340fc65d10b8742ca75463df02404fa0" + "sha256": "aaa4aeff255fc437d48b6f897b86252f5790676e7ead2662fbca3a0232d3da15" }, { "alt": "Grouped bars of decode tok/s at 16,384 tokens for release 2.11.2, #475, #478 and four #485 cascade operating points, each annotated with its HumanEval strict pass@1", From 1c0fb51d967c23b7571dc857edd310e3460cae0a Mon Sep 17 00:00:00 2001 From: davidtai Date: Wed, 9 Sep 2026 16:39:55 -0500 Subject: [PATCH 26/30] docs(cascade): acceptance chart as two panels, no twin axis cascade_accept_vs_alpha.svg plotted the kept-draft rate on a left y axis and tokens per cycle on a right y axis, and its legend named the series '(left)' and '(right)' without saying what those axes were; the right axis itself was clipped out of the figure. It is now two stacked panels on one alpha axis: top = cascade kept-draft rate (Definitions (c), cascade_accepted / cascade_positions), bottom = tokens per verify cycle, one line per rule (OPT Eq 10, TokenV3 Eq 15). Same data, same sidecar. --- .../charts/cascade_accept_vs_alpha.svg | 3689 +++++++++-------- 1 file changed, 2057 insertions(+), 1632 deletions(-) diff --git a/docs/perf/qwen38-cascade-acceptance/charts/cascade_accept_vs_alpha.svg b/docs/perf/qwen38-cascade-acceptance/charts/cascade_accept_vs_alpha.svg index c499c4c2b..78f1abe46 100644 --- a/docs/perf/qwen38-cascade-acceptance/charts/cascade_accept_vs_alpha.svg +++ b/docs/perf/qwen38-cascade-acceptance/charts/cascade_accept_vs_alpha.svg @@ -1,16 +1,16 @@ - + - 2026-09-09T09:55:36.710149 + 2026-09-09T16:39:08.969981 image/svg+xml - Matplotlib v3.11.0, https://matplotlib.org/ + Matplotlib v3.11.1, https://matplotlib.org/ @@ -21,8 +21,8 @@ - - - + - - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - - + + + + + + + - - - + + + - + - + - - + + - - - + - - - + + + - + - + - - + + + + + - + + - - - + + + - + - + - - - - - + + - + - - - + + + - + - + - - + + + - - - + + + - + - + - + @@ -289,20 +412,20 @@ L 312.090545 31.104 - - - + + + - + - + - - + + + + + - - - - - - - - - - - - - - - - - - - + + + - + - + + + + - + + - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - - - - - - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - - + - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - + - - - - - - - - + + + + + + + - + - +" style="stroke: #d55e00; stroke-linejoin: miter"/> - - - - - - + + + + + + - - - - - + - + - - - + + + + - + + @@ -1331,286 +1386,368 @@ z - - - - - - - - - - - - - - - - - - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - - - - - - - + + + + + + + + + + + + + + + + + - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + + + + + + + + + + + + + + - - - - - - - + + + + + + + + + + + + + + - - - + + + - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - - - + - - - - + + + + - + - + - + @@ -1619,14 +1756,19 @@ L 3.5 0 - + + + + - + - + - + @@ -1635,14 +1777,19 @@ L 3.5 0 - + + + + - + - + - + @@ -1651,51 +1798,40 @@ L 3.5 0 - + + + + - + - + - + - - - - - - - - - - - - - - - - - + + + + + + + + + + + + + + + + + @@ -1704,14 +1840,19 @@ z - + + + + - + - + - + @@ -1720,14 +1861,19 @@ z - + + + + - + - + - + @@ -1736,14 +1882,19 @@ z - + + + + - + - + - + @@ -1751,9 +1902,21 @@ z - - - + + + + + + @@ -1765,92 +1928,216 @@ z - - - - - + + + + + + + + + + + + - - - - - - - - - - - - - - - - - - + + + + + + + + + + + + - - - - + + + + + + + + + + + + - - - - - - - +" style="fill: #ffffff; opacity: 0.95; stroke: #cccccc; stroke-linejoin: miter"/> + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - - + + + + + - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - - + + + - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - - + + + + + From a9e9708d36463f97ded16f823f00fecb572d4a2a Mon Sep 17 00:00:00 2001 From: davidtai Date: Wed, 9 Sep 2026 16:44:53 -0500 Subject: [PATCH 27/30] docs(cascade): split the acceptance chart into two charts cascade_accept_vs_alpha.svg is replaced by cascade_kept_draft_vs_alpha.svg (kept-draft rate, Definitions (c)) and cascade_tokens_per_cycle_vs_alpha.svg (tokens per verify cycle), one quantity per chart, one line per rule. Same data, same sidecar. --- .../charts/cascade_accept_vs_alpha.svg | 2842 ----------------- .../charts/cascade_kept_draft_vs_alpha.svg | 2312 ++++++++++++++ .../cascade_tokens_per_cycle_vs_alpha.svg | 2130 ++++++++++++ .../charts/manifest.json | 23 +- 4 files changed, 4458 insertions(+), 2849 deletions(-) delete mode 100644 docs/perf/qwen38-cascade-acceptance/charts/cascade_accept_vs_alpha.svg create mode 100644 docs/perf/qwen38-cascade-acceptance/charts/cascade_kept_draft_vs_alpha.svg create mode 100644 docs/perf/qwen38-cascade-acceptance/charts/cascade_tokens_per_cycle_vs_alpha.svg diff --git a/docs/perf/qwen38-cascade-acceptance/charts/cascade_accept_vs_alpha.svg b/docs/perf/qwen38-cascade-acceptance/charts/cascade_accept_vs_alpha.svg deleted file mode 100644 index 78f1abe46..000000000 --- a/docs/perf/qwen38-cascade-acceptance/charts/cascade_accept_vs_alpha.svg +++ /dev/null @@ -1,2842 +0,0 @@ - - - - - - - - 2026-09-09T16:39:08.969981 - image/svg+xml - - - Matplotlib v3.11.1, https://matplotlib.org/ - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - diff --git a/docs/perf/qwen38-cascade-acceptance/charts/cascade_kept_draft_vs_alpha.svg b/docs/perf/qwen38-cascade-acceptance/charts/cascade_kept_draft_vs_alpha.svg new file mode 100644 index 000000000..d57ff5503 --- /dev/null +++ b/docs/perf/qwen38-cascade-acceptance/charts/cascade_kept_draft_vs_alpha.svg @@ -0,0 +1,2312 @@ + + + + + + + + 2026-09-09T16:44:22.380018 + image/svg+xml + + + Matplotlib v3.11.1, https://matplotlib.org/ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/docs/perf/qwen38-cascade-acceptance/charts/cascade_tokens_per_cycle_vs_alpha.svg b/docs/perf/qwen38-cascade-acceptance/charts/cascade_tokens_per_cycle_vs_alpha.svg new file mode 100644 index 000000000..d408900d8 --- /dev/null +++ b/docs/perf/qwen38-cascade-acceptance/charts/cascade_tokens_per_cycle_vs_alpha.svg @@ -0,0 +1,2130 @@ + + + + + + + + 2026-09-09T16:44:22.404861 + image/svg+xml + + + Matplotlib v3.11.1, https://matplotlib.org/ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/docs/perf/qwen38-cascade-acceptance/charts/manifest.json b/docs/perf/qwen38-cascade-acceptance/charts/manifest.json index d4c50e0a3..402c3f76d 100644 --- a/docs/perf/qwen38-cascade-acceptance/charts/manifest.json +++ b/docs/perf/qwen38-cascade-acceptance/charts/manifest.json @@ -9,13 +9,22 @@ "sha256": "f191004d8dc77f46079ac5ba7123ccde2377a743be69cbdcdba649344fdb1f5f" }, { - "alt": "Cascade accept rate and tokens per cycle by alpha on two y-axes", - "bytes": 86003, - "caption": "_Cascade accept rate (left) and tokens per cycle (right) by alpha for BOTH rules, OPT (Equation 10) and TokenV3 (Equation 15), from each run's cascade-accept metrics rather than the receipts. Equal alphas are NOT comparable across the two rules._", - "file": "cascade_accept_vs_alpha.svg", - "key": "cascade_accept_vs_alpha", - "label": "cascade acceptance and tokens/cycle by alpha", - "sha256": "47531a838091993d3c928b4d9cf4f3b292fbfc1fed0081edeef86bcb84eba860" + "alt": "Cascade kept-draft rate by alpha for the OPT and TokenV3 rules", + "bytes": 84026, + "caption": "_Cascade kept-draft rate, Definitions (c), by alpha for BOTH rules, OPT (Equation 10) and TokenV3 (Equation 15), from each run's `[cascade-accept]` verdict line rather than the receipts. Equal alphas are NOT comparable across the two rules._", + "file": "cascade_kept_draft_vs_alpha.svg", + "key": "cascade_kept_draft_vs_alpha", + "label": "cascade kept-draft rate by alpha", + "sha256": "cf90fc510da90ab94cceb5a4e511fc18601dd8ce4af047680903e606463ffa89" + }, + { + "alt": "Tokens per verify cycle by alpha for the OPT and TokenV3 rules", + "bytes": 74232, + "caption": "_Tokens per verify cycle by alpha for BOTH rules, OPT (Equation 10) and TokenV3 (Equation 15), from each run's `[cascade-accept]` verdict line rather than the receipts. Equal alphas are NOT comparable across the two rules._", + "file": "cascade_tokens_per_cycle_vs_alpha.svg", + "key": "cascade_tokens_per_cycle_vs_alpha", + "label": "tokens per verify cycle by alpha", + "sha256": "9a36a010ec9ee5d7bf86c3efb5ff15fceb157658bb38bd7f3ea2ff8655ecb2fe" }, { "alt": "Decode tok/s by context size for the exact pairing, cascade OPT alphas 0.0-2.0 and TokenV3 alpha 0.95", From d41a68e0d28b1113a9e9541aec4e50012ee41b5e Mon Sep 17 00:00:00 2001 From: davidtai Date: Wed, 9 Sep 2026 17:06:02 -0500 Subject: [PATCH 28/30] docs(cascade): drop "exact" arm labels; use #478 typical 0.2 as the reference MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit David's rulings applied to the #485 chart set: RULING 1 ("#475 is NOT exact"): no chart, legend, reference line, caption or manifest labels #475 or release 2.11.2 as "exact". The acceptance-off state is now named "acceptance mode off" / "cascade off"; the #475 arm is "#475 (base)". * cascade_decode_vs_alpha: reference lines relabeled "cascade off, #475 base (82.80)" and (see below) "typical 0.2 (99.16)". * cascade_decode_by_context: arm-C-caspair legend "exact (no cascade)" -> "cascade off (#475 base)". * cascade_vs_475_478_16k / _ladder: "#475 (base)", "acceptance mode off" / "cascade off"; the exact-acceptance disclaimer now names the ordinary speculative-decoding acceptance law, not any arm. RULING 2 (typical 0.2, not 0.09, is the #478 reference): every #485-vs-#478 comparison now references #478 at typical threshold 0.2 (pooled 16,384 window, 99.16 tok/s fastest of n=9, HumanEval strict pass@1 0.9695), replacing typical 0.09. Applied to cascade_decode_vs_alpha (dotted reference line), the 16K bars and the context ladder. The §2.4 ABAB (typical 0.09 vs TokenV3 0.95) is a separate measurement and is unchanged. Charts re-rendered from receipts (fastest-of-seeds, min-max band); manifest bytes/sha256 refreshed. Docs only; no receipts touched. --- .../charts/cascade_decode_by_context.svg | 280 ++-- .../charts/cascade_decode_vs_alpha.svg | 379 +++-- .../charts/cascade_vs_475_478_16k.svg | 1332 +++++++++------- .../charts/cascade_vs_475_478_ladder.svg | 1334 ++++++++++------- .../charts/manifest.json | 30 +- 5 files changed, 1989 insertions(+), 1366 deletions(-) diff --git a/docs/perf/qwen38-cascade-acceptance/charts/cascade_decode_by_context.svg b/docs/perf/qwen38-cascade-acceptance/charts/cascade_decode_by_context.svg index f8b483b2b..963486359 100644 --- a/docs/perf/qwen38-cascade-acceptance/charts/cascade_decode_by_context.svg +++ b/docs/perf/qwen38-cascade-acceptance/charts/cascade_decode_by_context.svg @@ -1509,21 +1509,21 @@ z - - + - - + - - + + + + + + - - - - - - - - - - - - - - - - - - + + + + + + + + + + + + + + + + + + + + + + - - + - + - - + - - - - + @@ -1723,17 +1807,17 @@ z - - + - + @@ -1754,29 +1838,17 @@ L 475.36125 83.86125 - - + - - - - + @@ -1798,17 +1870,17 @@ z - - + - + @@ -1829,17 +1901,17 @@ L 475.36125 107.8625 - - + - + @@ -1860,17 +1932,17 @@ L 475.36125 119.863125 - - + - + - 2026-09-09T09:55:36.647146 + 2024-09-09T00:00:00+00:00 image/svg+xml - Matplotlib v3.11.0, https://matplotlib.org/ + Matplotlib v3.11.1, https://matplotlib.org/ @@ -42,16 +42,16 @@ z +" clip-path="url(#p5dae6b0956)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - - + @@ -88,11 +88,11 @@ z +" clip-path="url(#p5dae6b0956)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - + @@ -167,11 +167,11 @@ z +" clip-path="url(#p5dae6b0956)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - + @@ -187,11 +187,11 @@ L 211.915636 31.104 +" clip-path="url(#p5dae6b0956)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - + @@ -220,11 +220,11 @@ z +" clip-path="url(#p5dae6b0956)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - + @@ -272,11 +272,11 @@ z +" clip-path="url(#p5dae6b0956)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - + @@ -293,11 +293,11 @@ L 327.760364 31.104 +" clip-path="url(#p5dae6b0956)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - + @@ -327,11 +327,11 @@ z +" clip-path="url(#p5dae6b0956)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - + @@ -784,16 +784,16 @@ z +" clip-path="url(#p5dae6b0956)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - - + @@ -849,11 +849,11 @@ z +" clip-path="url(#p5dae6b0956)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - + @@ -868,11 +868,11 @@ L 623.808 237.273389 +" clip-path="url(#p5dae6b0956)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - + @@ -888,11 +888,11 @@ L 623.808 197.825067 +" clip-path="url(#p5dae6b0956)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - + @@ -908,11 +908,11 @@ L 623.808 158.376744 +" clip-path="url(#p5dae6b0956)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - + @@ -928,11 +928,11 @@ L 623.808 118.928422 +" clip-path="url(#p5dae6b0956)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - + @@ -982,11 +982,11 @@ z +" clip-path="url(#p5dae6b0956)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - + @@ -1146,12 +1146,12 @@ z +" clip-path="url(#p5dae6b0956)" style="fill: none; stroke-dasharray: 4.44,1.92; stroke-dashoffset: 0; stroke: #888888; stroke-width: 1.2"/> - + +" clip-path="url(#p5dae6b0956)" style="fill: none; stroke: #009e73; stroke-width: 1.4"/> +" clip-path="url(#p5dae6b0956)" style="fill: none; stroke: #009e73; stroke-width: 1.4"/> +" clip-path="url(#p5dae6b0956)" style="fill: none; stroke: #009e73; stroke-width: 1.4"/> +" clip-path="url(#p5dae6b0956)" style="fill: none; stroke: #009e73; stroke-width: 1.4"/> +" clip-path="url(#p5dae6b0956)" style="fill: none; stroke: #009e73; stroke-width: 1.4"/> +" clip-path="url(#p5dae6b0956)" style="fill: none; stroke: #009e73; stroke-width: 1.4"/> - - - - - - - - + + + + + + + - - - - - - - + + + + + + + +" clip-path="url(#p5dae6b0956)" style="fill: none; stroke: #0072b2; stroke-width: 1.4"/> +" clip-path="url(#p5dae6b0956)" style="fill: none; stroke: #0072b2; stroke-width: 1.4"/> +" clip-path="url(#p5dae6b0956)" style="fill: none; stroke: #0072b2; stroke-width: 1.4"/> +" clip-path="url(#p5dae6b0956)" style="fill: none; stroke: #0072b2; stroke-width: 1.4"/> +" clip-path="url(#p5dae6b0956)" style="fill: none; stroke: #0072b2; stroke-width: 1.4"/> - - - - - - - + + + + + + - - - - - - + + + + + + @@ -1552,9 +1552,9 @@ L 211.915636 148.874247 L 276.273818 116.560374 L 340.632 43.580827 L 598.064727 42.768 -" clip-path="url(#p5b6668e1f0)" style="fill: none; stroke: #009e73; stroke-width: 1.8; stroke-linecap: square"/> +" clip-path="url(#p5dae6b0956)" style="fill: none; stroke: #009e73; stroke-width: 1.8; stroke-linecap: square"/> - - - - - - - - + + + + + + + @@ -1581,21 +1581,21 @@ L 211.915636 242.401804 L 276.273818 216.385409 L 314.888727 211.53676 L 327.760364 169.82751 -" clip-path="url(#p5b6668e1f0)" style="fill: none; stroke-dasharray: 6.66,2.88; stroke-dashoffset: 0; stroke: #0072b2; stroke-width: 1.8"/> +" clip-path="url(#p5dae6b0956)" style="fill: none; stroke-dasharray: 6.66,2.88; stroke-dashoffset: 0; stroke: #0072b2; stroke-width: 1.8"/> - - - - - - - + + + + + + @@ -1618,7 +1618,7 @@ L 644.67152 41.58275 L 652.67152 41.58275 " style="fill: none; stroke: #009e73; stroke-width: 1.5; stroke-linecap: square"/> - - + @@ -1743,7 +1743,7 @@ L 644.67152 53.583375 L 652.67152 53.583375 " style="fill: none; stroke-dasharray: 5.55,2.4; stroke-dashoffset: 0; stroke: #0072b2; stroke-width: 1.5"/> - - + @@ -1816,25 +1816,111 @@ L 652.67152 65.584 " style="fill: none; stroke-dasharray: 5.55,2.4; stroke-dashoffset: 0; stroke: #888888; stroke-width: 1.5"/> - + - - - - - - - - - - - - - - - - - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + @@ -1844,7 +1930,7 @@ L 652.67152 77.584625 " style="fill: none; stroke-dasharray: 1.5,2.475; stroke-dashoffset: 0; stroke: #d55e00; stroke-width: 1.5"/> - + @@ -1856,17 +1942,15 @@ L 652.67152 77.584625 - - - - - - - - - - - + + + + + + + + + @@ -2436,43 +2520,6 @@ L 1259 0 L 628 0 L 628 4666 z -" transform="scale(0.015625)"/> - @@ -2521,7 +2568,7 @@ z - + diff --git a/docs/perf/qwen38-cascade-acceptance/charts/cascade_vs_475_478_16k.svg b/docs/perf/qwen38-cascade-acceptance/charts/cascade_vs_475_478_16k.svg index 775faa77d..fc38b857a 100644 --- a/docs/perf/qwen38-cascade-acceptance/charts/cascade_vs_475_478_16k.svg +++ b/docs/perf/qwen38-cascade-acceptance/charts/cascade_vs_475_478_16k.svg @@ -6,11 +6,11 @@ - 2026-09-09T16:02:57.596299 + 2024-09-09T00:00:00+00:00 image/svg+xml - Matplotlib v3.11.0, https://matplotlib.org/ + Matplotlib v3.11.1, https://matplotlib.org/ @@ -41,12 +41,12 @@ z - - + @@ -228,24 +228,9 @@ z - - + + - + - + + + + - - - - - - - - - + + + + + + + + + + + + + + + + + + - + - - + + + + + + + + + + + + - - - - - - - - - - - + + + + + + + + + + + + + + + + + + + + - + @@ -493,8 +696,8 @@ z - - + + - - @@ -614,15 +761,14 @@ z - - + - + @@ -731,34 +877,13 @@ z - + - - - + @@ -887,6 +993,38 @@ z + + + @@ -906,16 +1044,16 @@ z +" clip-path="url(#p1f4906ac96)" style="fill: none; stroke: #dddddd; stroke-width: 0.7; stroke-linecap: square"/> - - + @@ -929,11 +1067,11 @@ L -3.5 0 +" clip-path="url(#p1f4906ac96)" style="fill: none; stroke: #dddddd; stroke-width: 0.7; stroke-linecap: square"/> - + @@ -948,11 +1086,11 @@ L 851.04 322.384947 +" clip-path="url(#p1f4906ac96)" style="fill: none; stroke: #dddddd; stroke-width: 0.7; stroke-linecap: square"/> - + @@ -967,11 +1105,11 @@ L 851.04 277.137894 +" clip-path="url(#p1f4906ac96)" style="fill: none; stroke: #dddddd; stroke-width: 0.7; stroke-linecap: square"/> - + @@ -1018,11 +1156,11 @@ z +" clip-path="url(#p1f4906ac96)" style="fill: none; stroke: #dddddd; stroke-width: 0.7; stroke-linecap: square"/> - + @@ -1037,11 +1175,11 @@ L 851.04 186.643787 +" clip-path="url(#p1f4906ac96)" style="fill: none; stroke: #dddddd; stroke-width: 0.7; stroke-linecap: square"/> - + @@ -1057,11 +1195,11 @@ L 851.04 141.396734 +" clip-path="url(#p1f4906ac96)" style="fill: none; stroke: #dddddd; stroke-width: 0.7; stroke-linecap: square"/> - + @@ -1077,51 +1215,12 @@ L 851.04 96.149681 - - + - - @@ -1282,7 +1369,7 @@ L 175.700731 367.632 L 175.700731 204.806128 L 96.414545 204.806128 z -" clip-path="url(#pf55da03b24)" style="fill: #666666; stroke: #222222; stroke-width: 0.6; stroke-linejoin: miter"/> +" clip-path="url(#p1f4906ac96)" style="fill: #666666; stroke: #222222; stroke-width: 0.6; stroke-linejoin: miter"/> +" clip-path="url(#p1f4906ac96)" style="fill: #0072b2; stroke: #222222; stroke-width: 0.6; stroke-linejoin: miter"/> +" clip-path="url(#p1f4906ac96)" style="fill: #e69f00; stroke: #222222; stroke-width: 0.6; stroke-linejoin: miter"/> +" clip-path="url(#p1f4906ac96)" style="fill: #d55e00; stroke: #222222; stroke-width: 0.6; stroke-linejoin: miter"/> +" clip-path="url(#p1f4906ac96)" style="fill: #cc79a7; stroke: #222222; stroke-width: 0.6; stroke-linejoin: miter"/> +" clip-path="url(#p1f4906ac96)" style="fill: #009e73; stroke: #222222; stroke-width: 0.6; stroke-linejoin: miter"/> +" clip-path="url(#p1f4906ac96)" style="fill: none; stroke: #222222; stroke-width: 1.5"/> +" clip-path="url(#p1f4906ac96)" style="fill: none; stroke: #222222; stroke-width: 1.5"/> +" clip-path="url(#p1f4906ac96)" style="fill: none; stroke: #222222; stroke-width: 1.5"/> +" clip-path="url(#p1f4906ac96)" style="fill: none; stroke: #222222; stroke-width: 1.5"/> +" clip-path="url(#p1f4906ac96)" style="fill: none; stroke: #222222; stroke-width: 1.5"/> +" clip-path="url(#p1f4906ac96)" style="fill: none; stroke: #222222; stroke-width: 1.5"/> - + - - + + +" clip-path="url(#p1f4906ac96)" style="fill: none; stroke: #222222; stroke-width: 1.5"/> +" clip-path="url(#p1f4906ac96)" style="fill: none; stroke: #222222; stroke-width: 1.5"/> +" clip-path="url(#p1f4906ac96)" style="fill: none; stroke: #222222; stroke-width: 1.5"/> +" clip-path="url(#p1f4906ac96)" style="fill: none; stroke: #222222; stroke-width: 1.5"/> +" clip-path="url(#p1f4906ac96)" style="fill: none; stroke: #222222; stroke-width: 1.5"/> +" clip-path="url(#p1f4906ac96)" style="fill: none; stroke: #222222; stroke-width: 1.5"/> +" clip-path="url(#p1f4906ac96)" style="fill: none; stroke: #222222; stroke-width: 1.5"/> +" clip-path="url(#p1f4906ac96)" style="fill: none; stroke: #222222; stroke-width: 1.5"/> +" clip-path="url(#p1f4906ac96)" style="fill: none; stroke: #222222; stroke-width: 1.5"/> - - + + - - - - - - + + + + - - - + + @@ -1789,13 +1829,36 @@ z - - + + + + + @@ -1829,6 +1892,31 @@ z + +" clip-path="url(#p1f4906ac96)" style="fill: none; stroke: #222222; stroke-width: 1.5; stroke-linecap: square"/> - - - + + - - +" clip-path="url(#p1f4906ac96)" style="fill: none; stroke: #222222; stroke-width: 1.5; stroke-linecap: square"/> + + - - - + + + - - +" clip-path="url(#p1f4906ac96)" style="fill: none; stroke: #222222; stroke-width: 1.5; stroke-linecap: square"/> + + - - +" clip-path="url(#p1f4906ac96)" style="fill: none; stroke: #222222; stroke-width: 1.5; stroke-linecap: square"/> + + - - +" clip-path="url(#p1f4906ac96)" style="fill: none; stroke: #222222; stroke-width: 1.5; stroke-linecap: square"/> + + - - + + - - - - - - - - - - - - - - - - - - - - - - - - + + + + + + + + + + + + + + + + + + + + - - + + - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - + @@ -3098,34 +3178,8 @@ z - + - - + - + @@ -3581,18 +3635,212 @@ z - - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - + diff --git a/docs/perf/qwen38-cascade-acceptance/charts/cascade_vs_475_478_ladder.svg b/docs/perf/qwen38-cascade-acceptance/charts/cascade_vs_475_478_ladder.svg index fe9f80437..0b4fa0e46 100644 --- a/docs/perf/qwen38-cascade-acceptance/charts/cascade_vs_475_478_ladder.svg +++ b/docs/perf/qwen38-cascade-acceptance/charts/cascade_vs_475_478_ladder.svg @@ -6,11 +6,11 @@ - 2026-09-09T16:02:57.746972 + 2024-09-09T00:00:00+00:00 image/svg+xml - Matplotlib v3.11.0, https://matplotlib.org/ + Matplotlib v3.11.1, https://matplotlib.org/ @@ -42,16 +42,16 @@ z +" clip-path="url(#pdba595bf2d)" style="fill: none; stroke: #dddddd; stroke-width: 0.7; stroke-linecap: square"/> - - + @@ -158,11 +158,11 @@ z +" clip-path="url(#pdba595bf2d)" style="fill: none; stroke: #dddddd; stroke-width: 0.7; stroke-linecap: square"/> - + @@ -251,11 +251,11 @@ z +" clip-path="url(#pdba595bf2d)" style="fill: none; stroke: #dddddd; stroke-width: 0.7; stroke-linecap: square"/> - + @@ -338,11 +338,11 @@ z +" clip-path="url(#pdba595bf2d)" style="fill: none; stroke: #dddddd; stroke-width: 0.7; stroke-linecap: square"/> - + @@ -373,11 +373,11 @@ z +" clip-path="url(#pdba595bf2d)" style="fill: none; stroke: #dddddd; stroke-width: 0.7; stroke-linecap: square"/> - + @@ -423,11 +423,11 @@ z +" clip-path="url(#pdba595bf2d)" style="fill: none; stroke: #dddddd; stroke-width: 0.7; stroke-linecap: square"/> - + @@ -447,11 +447,11 @@ L 560.49984 72.036 +" clip-path="url(#pdba595bf2d)" style="fill: none; stroke: #dddddd; stroke-width: 0.7; stroke-linecap: square"/> - + @@ -801,16 +801,16 @@ z +" clip-path="url(#pdba595bf2d)" style="fill: none; stroke: #dddddd; stroke-width: 0.7; stroke-linecap: square"/> - - + @@ -825,11 +825,11 @@ L -3.5 0 +" clip-path="url(#pdba595bf2d)" style="fill: none; stroke: #dddddd; stroke-width: 0.7; stroke-linecap: square"/> - + @@ -844,11 +844,11 @@ L 656.208 273.960053 +" clip-path="url(#pdba595bf2d)" style="fill: none; stroke: #dddddd; stroke-width: 0.7; stroke-linecap: square"/> - + @@ -863,11 +863,11 @@ L 656.208 223.566711 +" clip-path="url(#pdba595bf2d)" style="fill: none; stroke: #dddddd; stroke-width: 0.7; stroke-linecap: square"/> - + @@ -883,11 +883,11 @@ L 656.208 173.17337 +" clip-path="url(#pdba595bf2d)" style="fill: none; stroke: #dddddd; stroke-width: 0.7; stroke-linecap: square"/> - + @@ -903,11 +903,11 @@ L 656.208 122.780029 +" clip-path="url(#pdba595bf2d)" style="fill: none; stroke: #dddddd; stroke-width: 0.7; stroke-linecap: square"/> - + @@ -1124,7 +1124,7 @@ z - - - + + - - - + + - - - + + - - - + + @@ -1221,9 +1221,9 @@ L 273.37536 259.61451 L 369.08352 266.908996 L 464.79168 251.222367 L 560.49984 330.951427 -" clip-path="url(#p0d0fb95840)" style="fill: none; stroke: #0072b2; stroke-width: 1.9; stroke-linecap: square"/> +" clip-path="url(#pdba595bf2d)" style="fill: none; stroke: #0072b2; stroke-width: 1.9; stroke-linecap: square"/> - - - - - - - - + + + + + + + - + - - - - - - - - + + + + + + + @@ -1276,21 +1276,21 @@ L 273.37536 149.099981 L 369.08352 131.094078 L 464.79168 176.980453 L 560.49984 199.90543 -" clip-path="url(#p0d0fb95840)" style="fill: none; stroke: #d55e00; stroke-width: 1.9; stroke-linecap: square"/> +" clip-path="url(#pdba595bf2d)" style="fill: none; stroke: #d55e00; stroke-width: 1.9; stroke-linecap: square"/> - - - - - - - - + + + + + + + @@ -1300,22 +1300,22 @@ L 273.37536 137.407833 L 369.08352 161.497412 L 464.79168 194.337214 L 560.49984 244.347364 -" clip-path="url(#p0d0fb95840)" style="fill: none; stroke: #009e73; stroke-width: 1.9; stroke-linecap: square"/> +" clip-path="url(#pdba595bf2d)" style="fill: none; stroke: #009e73; stroke-width: 1.9; stroke-linecap: square"/> - - - - - - - - + + + + + + + @@ -2183,11 +2183,11 @@ L 675.91888 83.562625 L 684.71888 83.562625 " style="fill: none; stroke: #0072b2; stroke-width: 1.9; stroke-linecap: square"/> - + - + - - @@ -2255,17 +2267,25 @@ z - - - - - - - - - - - + + + + + + + + + + + + + + + + + + + @@ -2274,13 +2294,20 @@ L 675.91888 96.763312 L 684.71888 96.763312 " style="fill: none; stroke: #e69f00; stroke-width: 1.9; stroke-linecap: square"/> - + - + + - - + @@ -2315,7 +2341,7 @@ L 675.91888 109.964 L 684.71888 109.964 " style="fill: none; stroke: #d55e00; stroke-width: 1.9; stroke-linecap: square"/> - + @@ -2382,7 +2408,7 @@ L 675.91888 123.164687 L 684.71888 123.164687 " style="fill: none; stroke: #009e73; stroke-width: 1.9; stroke-linecap: square"/> - + @@ -2566,8 +2592,8 @@ z - - + + - @@ -2731,37 +2751,36 @@ z - - - - - - - - - - - - - - - - - - - - - - - - - - - - + + + + + + + + + + + + + + + + + + + + + + + + + + + - - + + + - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - - + + - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - + @@ -3306,10 +3348,224 @@ z + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - + diff --git a/docs/perf/qwen38-cascade-acceptance/charts/manifest.json b/docs/perf/qwen38-cascade-acceptance/charts/manifest.json index 402c3f76d..cca3c4035 100644 --- a/docs/perf/qwen38-cascade-acceptance/charts/manifest.json +++ b/docs/perf/qwen38-cascade-acceptance/charts/manifest.json @@ -1,12 +1,12 @@ [ { - "alt": "Decode tok/s by cascade alpha, fastest of three with a min-max band, and two reference lines for the exact law and typical acceptance at 0.09", - "bytes": 94167, - "caption": "_Decode by alpha for BOTH cascade rules: OPT (Equation 10) and TokenV3 (Equation 15), each fastest of three seeds with a min-max band. Reference lines are the exact law (82.80) and typical acceptance at 0.09 (104.65) from the pooled 16K windows in #478. The two rules' alphas are different quantities and are NOT comparable at equal alpha; compare at equal decode speed._", + "alt": "Decode tok/s by cascade alpha, fastest of three with a min-max band, and two reference lines for cascade off (#475 base) and typical acceptance at 0.2", + "bytes": 95540, + "caption": "_Decode by alpha for BOTH cascade rules: OPT (Equation 10) and TokenV3 (Equation 15), each fastest of three seeds with a min-max band. Reference lines are cascade off / #475 base (82.80) and typical acceptance at 0.2 (99.16) from the pooled 16K windows in #478. The two rules' alphas are different quantities and are NOT comparable at equal alpha; compare at equal decode speed._", "file": "cascade_decode_vs_alpha.svg", "key": "cascade_decode_vs_alpha", "label": "cascade decode by alpha", - "sha256": "f191004d8dc77f46079ac5ba7123ccde2377a743be69cbdcdba649344fdb1f5f" + "sha256": "2ac1a4d3ffb281a0d218f67b09e4c2af50a85023707844ae259555f2ba6a79c7" }, { "alt": "Cascade kept-draft rate by alpha for the OPT and TokenV3 rules", @@ -27,30 +27,30 @@ "sha256": "9a36a010ec9ee5d7bf86c3efb5ff15fceb157658bb38bd7f3ea2ff8655ecb2fe" }, { - "alt": "Decode tok/s by context size for the exact pairing, cascade OPT alphas 0.0-2.0 and TokenV3 alpha 0.95", - "bytes": 88853, - "caption": "_Decode tok/s by context size, one line per arm: the exact pairing, the OPT rule at alpha 0.0/0.25/0.5/0.75/1.0/2.0, and the TokenV3 rule at alpha 0.95. Each point is the fastest seed with a min-max band; the 16,384 rung is merged in from the 16K sweeps, the OPT arms from battery475/g and the TokenV3 arm from battery475/t; 261,120 is absent because every arm exceeds the memory knob on the #475 base without #482._", + "alt": "Decode tok/s by context size for cascade off (#475 base), cascade OPT alphas 0.0-2.0 and TokenV3 alpha 0.95", + "bytes": 90125, + "caption": "_Decode tok/s by context size, one line per arm: cascade off (#475 base), the OPT rule at alpha 0.0/0.25/0.5/0.75/1.0/2.0, and the TokenV3 rule at alpha 0.95. Each point is the fastest seed with a min-max band; the 16,384 rung is merged in from the 16K sweeps, the OPT arms from battery475/g and the TokenV3 arm from battery475/t; 261,120 is absent because every arm exceeds the memory knob on the #475 base without #482._", "file": "cascade_decode_by_context.svg", "key": "cascade_decode_by_context", "label": "cascade decode by context", - "sha256": "aaa4aeff255fc437d48b6f897b86252f5790676e7ead2662fbca3a0232d3da15" + "sha256": "345985e2361ead5d9f0d1edc1f307bd57c31eaaad2f8642bc04fe0c6a5ba14cf" }, { "alt": "Grouped bars of decode tok/s at 16,384 tokens for release 2.11.2, #475, #478 and four #485 cascade operating points, each annotated with its HumanEval strict pass@1", - "bytes": 141088, - "caption": "_Decode tok/s at 16,384 tokens on the Optimized-Speed pack: release 2.11.2 at the exact speculative law (71.97), #475 at the exact law (82.84), #478 at typical acceptance 0.09 (104.65), and #485 at the OPT rule (Equation 10) alpha 0.25 (104.78) and the TokenV3 rule (Equation 15) alpha 0.75 (95.30) and 0.95 (107.10). Bar height is the fastest seed, the whisker is the min-max band, and the figure above each whisker is HumanEval strict pass@1 for the same arm. Alpha is NOT comparable between the two rules; compare at equal decode speed._", + "bytes": 157421, + "caption": "_Decode tok/s at 16,384 tokens on the Optimized-Speed pack: release 2.11.2 with acceptance mode off (71.97), #475 (base) with acceptance mode off (82.84), #478 at typical acceptance 0.2 (99.16), and #485 at the OPT rule (Equation 10) alpha 0.25 (104.78) and the TokenV3 rule (Equation 15) alpha 0.75 (95.30) and 0.95 (107.10). Bar height is the fastest seed, the whisker is the min-max band, and the figure above each whisker is HumanEval strict pass@1 for the same arm. Alpha is NOT comparable between the two rules; compare at equal decode speed._", "file": "cascade_vs_475_478_16k.svg", "key": "cascade_vs_475_478_16k", "label": "cascade against #475 and #478 at 16K", - "sha256": "393363aa89e48508f653a5a76718f2bbc94611f2a8cf6e1bbd51dfa171772ce5" + "sha256": "cfbbae204ddbb54a16bcbc690f543fa644a177e5091725e3cdda9215f8ef88e9" }, { - "alt": "Decode tok/s against context size from 1,024 to 131,072 tokens for #475 exact, #478 typical 0.09, #485 OPT alpha 0.25 and #485 TokenV3 alpha 0.95", - "bytes": 127817, - "caption": "_Decode tok/s against context size on the Optimized-Speed pack for #475 at the exact speculative law, #478 at typical acceptance 0.09, #485 at the OPT rule alpha 0.25 and #485 at the TokenV3 rule alpha 0.95. Each point is the fastest seed and the band is the min-max range. No arm has a 261,120 point: on the #475 base without #482 every arm on this pack exceeds the memory knob._", + "alt": "Decode tok/s against context size from 1,024 to 131,072 tokens for #475 (base) cascade off, #478 typical 0.2, #485 OPT alpha 0.25 and #485 TokenV3 alpha 0.95", + "bytes": 143817, + "caption": "_Decode tok/s against context size on the Optimized-Speed pack for #475 (base) with cascade off, #478 at typical acceptance 0.2, #485 at the OPT rule alpha 0.25 and #485 at the TokenV3 rule alpha 0.95. Each point is the fastest seed and the band is the min-max range. No arm has a 261,120 point: on the #475 base without #482 every arm on this pack exceeds the memory knob._", "file": "cascade_vs_475_478_ladder.svg", "key": "cascade_vs_475_478_ladder", "label": "cascade against #475 and #478 by context", - "sha256": "4743c0ef31d41736e771b1d082e9641050cf6822c213fc9f3e2671f0d8f9e324" + "sha256": "c4747722fe0d98c71038b21cd52b7a20ca8ce2f7dfb8bc60087e3bcb56acecb4" } ] From 7bf945879f4f4cbd0a8d7a310f2412404ad9f4a8 Mon Sep 17 00:00:00 2001 From: davidtai Date: Wed, 9 Sep 2026 18:57:30 -0500 Subject: [PATCH 29/30] docs(cascade): HumanEval pass@1 bars per cascade operating point cascade_humaneval_by_alpha.svg: strict and completed-task pass@1 as grouped bars for cascade off (#475 base), the OPT rule at alpha 0.0/0.25/0.5/0.75 and the TokenV3 rule at 0.75/0.95, with the cascade-off strict level as a reference line, so the quality cost of the OPT rule is visible next to TokenV3. Data: evalsweep478/armG_summary.json and the #478 sweep's typical-off cell; nothing re-measured. --- .../charts/cascade_humaneval_by_alpha.svg | 3145 +++++++++++++++++ .../charts/manifest.json | 9 + 2 files changed, 3154 insertions(+) create mode 100644 docs/perf/qwen38-cascade-acceptance/charts/cascade_humaneval_by_alpha.svg diff --git a/docs/perf/qwen38-cascade-acceptance/charts/cascade_humaneval_by_alpha.svg b/docs/perf/qwen38-cascade-acceptance/charts/cascade_humaneval_by_alpha.svg new file mode 100644 index 000000000..b969bfb0d --- /dev/null +++ b/docs/perf/qwen38-cascade-acceptance/charts/cascade_humaneval_by_alpha.svg @@ -0,0 +1,3145 @@ + + + + + + + + 2026-09-09T18:57:30.620175 + image/svg+xml + + + Matplotlib v3.11.1, https://matplotlib.org/ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/docs/perf/qwen38-cascade-acceptance/charts/manifest.json b/docs/perf/qwen38-cascade-acceptance/charts/manifest.json index cca3c4035..a45a01d4f 100644 --- a/docs/perf/qwen38-cascade-acceptance/charts/manifest.json +++ b/docs/perf/qwen38-cascade-acceptance/charts/manifest.json @@ -52,5 +52,14 @@ "key": "cascade_vs_475_478_ladder", "label": "cascade against #475 and #478 by context", "sha256": "c4747722fe0d98c71038b21cd52b7a20ca8ce2f7dfb8bc60087e3bcb56acecb4" + }, + { + "alt": "HumanEval strict and completed-task pass@1 as grouped bars per cascade operating point: cascade off, OPT alpha 0.0 to 0.75, TokenV3 alpha 0.75 and 0.95", + "bytes": 123272, + "caption": "_HumanEval pass@1 per cascade operating point, one seed at the cell's sampler, re-scored offline: strict (truncated tasks count as failures) and completed-task (truncated tasks excluded). The OPT rule (Equation 10) loses accuracy at every alpha; the TokenV3 rule (Equation 15) holds the cascade-off level at 0.75 and 0.95. OPT alpha 1.0 did not finish 164 tasks in 3 h and has no bar. Equal alphas are NOT comparable across the two rules._", + "file": "cascade_humaneval_by_alpha.svg", + "key": "cascade_humaneval_by_alpha", + "label": "HumanEval by cascade operating point", + "sha256": "d7295797e5702a8fe63c260b0ed5e733183a933e7619358e33f5f14f2294fdd8" } ] From 373c06fc0f8e3c2fd413d0618c2d483c56b09689 Mon Sep 17 00:00:00 2001 From: davidtai Date: Thu, 10 Sep 2026 06:55:33 -0500 Subject: [PATCH 30/30] docs(cascade): add #478 typical 0.2 to the HumanEval-by-operating-point chart cascade_humaneval_by_alpha.svg gains a #478 typical 0.2 bar pair (strict 0.9695, completed-task 1.0000, the #478 sweep's own cell on the same base, seed and sampler) beside cascade off, so the quality-neutral reference the speed comparisons use is on the quality chart too. Both sweep cells are now read from evalsweep478/results.csv instead of being pasted constants. --- .../charts/cascade_humaneval_by_alpha.svg | 1713 +++++++++-------- .../charts/manifest.json | 8 +- 2 files changed, 939 insertions(+), 782 deletions(-) diff --git a/docs/perf/qwen38-cascade-acceptance/charts/cascade_humaneval_by_alpha.svg b/docs/perf/qwen38-cascade-acceptance/charts/cascade_humaneval_by_alpha.svg index b969bfb0d..9770842e6 100644 --- a/docs/perf/qwen38-cascade-acceptance/charts/cascade_humaneval_by_alpha.svg +++ b/docs/perf/qwen38-cascade-acceptance/charts/cascade_humaneval_by_alpha.svg @@ -6,7 +6,7 @@ - 2026-09-09T18:57:30.620175 + 2026-09-10T06:55:06.485511 image/svg+xml @@ -40,23 +40,23 @@ z - + - - + - + - + - + - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - + - + - - - - @@ -593,52 +778,26 @@ z - - - + + + - + - + - + - + - - - - + @@ -652,26 +811,26 @@ z - - - + + + - + - + - + - + - + @@ -684,26 +843,26 @@ L 430.92 32.256 - - - + + + - + - + - + - + - + @@ -717,20 +876,20 @@ L 531.367552 32.256 - - - + + + - + - + - + - + - - + @@ -831,20 +971,20 @@ z - - - + + + - + - + - + - + @@ -854,7 +994,7 @@ L 732.262657 32.256 - + - + +" clip-path="url(#p7b5192a712)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - + - - + - + @@ -1007,17 +1147,17 @@ z - + +" clip-path="url(#p7b5192a712)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - + - + - + @@ -1027,80 +1167,39 @@ L 804.384 238.842947 - + - - - - - - - - - - - - - - - - - - +" clip-path="url(#p7b5192a712)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - + - - - - - + + + + + + + + + + + + + + + + + + + + @@ -1108,17 +1207,17 @@ z - + +" clip-path="url(#p7b5192a712)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - + - + - + @@ -1128,17 +1227,17 @@ L 804.384 144.939789 - + +" clip-path="url(#p7b5192a712)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - + - + - + @@ -1148,17 +1247,17 @@ L 804.384 113.638737 - + +" clip-path="url(#p7b5192a712)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - + - + - + @@ -1168,17 +1267,17 @@ L 804.384 82.337684 - + +" clip-path="url(#p7b5192a712)" style="fill: none; stroke: #b0b0b0; stroke-opacity: 0.25; stroke-width: 0.8; stroke-linecap: square"/> - + - + - + @@ -1204,7 +1303,7 @@ z - + @@ -1357,27 +1456,6 @@ Q 5953 1231 5456 756 Q 4959 281 4084 263 L 4084 744 z -" transform="scale(0.015625)"/> - - + + +" clip-path="url(#p7b5192a712)" style="fill: #7f7f7f; stroke: #000000; stroke-width: 0.6; stroke-linejoin: miter"/> - +" clip-path="url(#p7b5192a712)" style="fill: #7f7f7f; opacity: 0.45; stroke: #000000; stroke-width: 0.6; stroke-linejoin: miter"/> - + - + - + - + @@ -1635,24 +1713,89 @@ z - +" clip-path="url(#p7b5192a712)" style="fill: #e69f00; stroke: #000000; stroke-width: 0.6; stroke-linejoin: miter"/> - +" clip-path="url(#p7b5192a712)" style="fill: #e69f00; opacity: 0.45; stroke: #000000; stroke-width: 0.6; stroke-linejoin: miter"/> - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - + - + - + @@ -1696,25 +1839,25 @@ z - - + +" clip-path="url(#p7b5192a712)" style="fill: #d55e00; stroke: #000000; stroke-width: 0.6; stroke-linejoin: miter"/> - - + +" clip-path="url(#p7b5192a712)" style="fill: #d55e00; opacity: 0.45; stroke: #000000; stroke-width: 0.6; stroke-linejoin: miter"/> - + - + - - + - + @@ -1810,25 +1928,25 @@ z - - + +" clip-path="url(#p7b5192a712)" style="fill: #d55e00; stroke: #000000; stroke-width: 0.6; stroke-linejoin: miter"/> - - + +" clip-path="url(#p7b5192a712)" style="fill: #d55e00; opacity: 0.45; stroke: #000000; stroke-width: 0.6; stroke-linejoin: miter"/> - + - + @@ -1837,9 +1955,9 @@ z - + - + @@ -1848,25 +1966,25 @@ z - - + +" clip-path="url(#p7b5192a712)" style="fill: #d55e00; stroke: #000000; stroke-width: 0.6; stroke-linejoin: miter"/> - - + +" clip-path="url(#p7b5192a712)" style="fill: #d55e00; opacity: 0.45; stroke: #000000; stroke-width: 0.6; stroke-linejoin: miter"/> - + - + @@ -1875,9 +1993,9 @@ z - + - + @@ -1886,25 +2004,25 @@ z - - + +" clip-path="url(#p7b5192a712)" style="fill: #0072b2; stroke: #000000; stroke-width: 0.6; stroke-linejoin: miter"/> - - + +" clip-path="url(#p7b5192a712)" style="fill: #0072b2; opacity: 0.45; stroke: #000000; stroke-width: 0.6; stroke-linejoin: miter"/> - + - + @@ -1913,9 +2031,9 @@ z - + - + @@ -1924,25 +2042,25 @@ z - - + +" clip-path="url(#p7b5192a712)" style="fill: #0072b2; stroke: #000000; stroke-width: 0.6; stroke-linejoin: miter"/> - - + +" clip-path="url(#p7b5192a712)" style="fill: #0072b2; opacity: 0.45; stroke: #000000; stroke-width: 0.6; stroke-linejoin: miter"/> - + - + @@ -1951,9 +2069,9 @@ z - + - + @@ -1962,9 +2080,9 @@ z - + - + - @@ -2024,27 +2129,10 @@ z - - - + + + - + @@ -2106,91 +2209,98 @@ z - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - + - + - + @@ -2274,7 +2384,7 @@ z - + - + @@ -2327,7 +2437,7 @@ z - + - + @@ -2374,7 +2484,7 @@ z - + - + @@ -2422,7 +2532,7 @@ z - + - + @@ -2491,9 +2601,63 @@ z + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - + @@ -2742,7 +2906,7 @@ z - + - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - + - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + - + diff --git a/docs/perf/qwen38-cascade-acceptance/charts/manifest.json b/docs/perf/qwen38-cascade-acceptance/charts/manifest.json index a45a01d4f..8366eb879 100644 --- a/docs/perf/qwen38-cascade-acceptance/charts/manifest.json +++ b/docs/perf/qwen38-cascade-acceptance/charts/manifest.json @@ -54,12 +54,12 @@ "sha256": "c4747722fe0d98c71038b21cd52b7a20ca8ce2f7dfb8bc60087e3bcb56acecb4" }, { - "alt": "HumanEval strict and completed-task pass@1 as grouped bars per cascade operating point: cascade off, OPT alpha 0.0 to 0.75, TokenV3 alpha 0.75 and 0.95", - "bytes": 123272, - "caption": "_HumanEval pass@1 per cascade operating point, one seed at the cell's sampler, re-scored offline: strict (truncated tasks count as failures) and completed-task (truncated tasks excluded). The OPT rule (Equation 10) loses accuracy at every alpha; the TokenV3 rule (Equation 15) holds the cascade-off level at 0.75 and 0.95. OPT alpha 1.0 did not finish 164 tasks in 3 h and has no bar. Equal alphas are NOT comparable across the two rules._", + "alt": "HumanEval strict and completed-task pass@1 as grouped bars per operating point: cascade off, #478 typical 0.2, OPT alpha 0.0 to 0.75, TokenV3 alpha 0.75 and 0.95", + "bytes": 132371, + "caption": "_HumanEval pass@1 per cascade operating point, one seed at the cell's sampler, re-scored offline, with #478 typical 0.2 beside cascade off as the second reference: strict (truncated tasks count as failures) and completed-task (truncated tasks excluded). The OPT rule (Equation 10) loses accuracy at every alpha; the TokenV3 rule (Equation 15) holds the cascade-off level at 0.75 and 0.95. OPT alpha 1.0 did not finish 164 tasks in 3 h and has no bar. Equal alphas are NOT comparable across the two rules._", "file": "cascade_humaneval_by_alpha.svg", "key": "cascade_humaneval_by_alpha", "label": "HumanEval by cascade operating point", - "sha256": "d7295797e5702a8fe63c260b0ed5e733183a933e7619358e33f5f14f2294fdd8" + "sha256": "7008eed4da1c881dee646ee9563264eca8708e8950ed0d3594e5caadb5721bc8" } ]