diff --git a/.gitignore b/.gitignore index 5e74d1f2..96cc709a 100644 --- a/.gitignore +++ b/.gitignore @@ -1,3 +1,6 @@ __pycache__ *.egg-info build/ + +# local modal test harness +_modal/ diff --git a/CLAUDE.md b/CLAUDE.md new file mode 100644 index 00000000..58147181 --- /dev/null +++ b/CLAUDE.md @@ -0,0 +1,81 @@ +# CLAUDE.md + +This file provides guidance to Claude Code (claude.ai/code) when working with code in this repository. + +## What this is + +FlashQLA is a single-operation kernel library: fused, warp-specialized **TileLang** kernels for the **Gated Delta Rule (GDN) chunked prefill** forward and backward passes, targeting NVIDIA **Hopper (SM90)**. It is a drop-in faster alternative to the Flash-Linear-Attention (FLA) Triton kernels (2-3x forward, 2x backward). Everything in the repo exists to compute `chunk_gated_delta_rule`. + +Hard environment requirements: SM90 or above, CUDA 12.8+, PyTorch 2.8+, Python 3.10+. Pinned deps: `tilelang==0.1.8`, `apache-tvm-ffi==0.1.9`. + +## Commands + +```bash +# Install (editable install matters: setup.py embeds the git SHA in the version) +pip install -v . + +# Lint/format — pre-commit is formatting-only (ruff); real linting is CI's job +pre-commit run -a # run on all files +pre-commit install # install the git hook + +# Correctness + speedup tests — require fla for the reference comparison +pip install flash_linear_attention==0.5.0 +cd tests +python test_gdr.py --set develop # quick smoke +python test_gdr.py --set varlen --num-heads 32 # variable-length batches +python test_gdr.py --set profile --num-heads 32 # latency sweep +python test_gdr.py --set product --ref-dtype float32 --num-heads 32 + +# Static signature guard — CPU-only, no GPU/fla needed (parses source via ast) +python tests/test_function_signature.py + +# Benchmark vs FLA Triton + FlashInfer +pip install flash_linear_attention==0.5.0 flashinfer-python==0.6.9 +cd benchmark +python bench_gated_delta_rule.py +``` + +`test_gdr.py` is driven by a `--set ` preset that loads `tests/settings/.csv` (develop / varlen / profile / product), with per-run overrides via flags: `--num-heads`/`--nvh`, `--nkh`, `--seqlen`, `--no-h0`, `--skip-bwd`, `--no-cp` (disable auto context-parallelism), `--swa-ratio`, `--data-dtype`, `--ref-dtype`, `--hide-acc`, `--hide-lat`, `--seed`. There is no single-test selector beyond editing the CSV; a "test" is one shape row. The accuracy check runs the QLA kernel **1000 times in a loop** asserting the result stays within 2% of the reference — this is a deliberate race/nondeterminism detector, not redundancy. + +## Architecture + +Three layers, top to bottom: + +**1. Public API** — `flash_qla/__init__.py` re-exports three entry points: +- `chunk_gated_delta_rule` — the autograd-wrapped op most callers use. It is `@torch.compiler.disable`d and wraps `ChunkGatedDeltaRuleFunction` (a `torch.autograd.Function`). +- `chunk_gated_delta_rule_fwd` / `chunk_gated_delta_rule_bwd` — low-level functions that bypass autograd (what the tests and benchmark call directly). The forward returns intermediates (`g, A, o, h, final_state`) that the backward needs fed back in. + +**2. Orchestration** — `flash_qla/ops/gated_delta_rule/chunk/__init__.py` sequences kernels into passes and enforces all the input invariants: +- Forward: `chunk_local_cumsum(g)` → `kkt_solve` (builds the `A` interaction matrix) → optional `intra_card_cp_preprocess` (auto context-parallelism) → `fused_gdr_fwd`. +- Backward: `fused_gdr_h` (recompute hidden states `h`) → `fused_gdr_bwd` → `group_reduce_vector` on `dq`/`dk` when GQA (`Hg < H`) → reverse `chunk_local_cumsum(dg)`. +- Invariants enforced here (changing them breaks the kernels): `head_dim_k == head_dim_v == 128`, `chunk_size == 64`, dtype is bf16/fp16 (**not** fp32), `num_v_heads % num_k_heads == 0` (GQA), `head_first=False` only, and batch size must be 1 when `cu_seqlens` is given. + +**3. Kernels** — `flash_qla/ops/gated_delta_rule/chunk/hopper/`. Dispatch is hard-gated: both `chunk/__init__.py` and `cp_context.py` check `tilelang.contrib.nvcc.get_target_compute_version() == "9.0"` and **raise** otherwise. Kernel modules: +- `fused_fwd.py` (`fused_gdr_fwd`) — the main fused forward. +- `fused_bwd.py` (`fused_gdr_bwd`) — the main fused backward. +- `prepare_h.py` (`fused_gdr_h`) — recomputes hidden state `h`; used by both backward and CP preprocessing. +- `kkt_solve.py` (`kkt_solve`) — solves the lower-triangular `(I - tril(βKKᵀ))` system into `A`. +- `cp_fwd.py` (`get_warmup_chunks`, `correct_initial_states`) — context-parallelism helpers. + +### Intra-card context parallelism (the distinctive part) + +`cp_context.py` is the headline optimization. When `auto_cp=True` and `batch_size == 1`, it splits one long sequence into shorter CP sub-sequences to raise SM occupancy under TP / long-context / small-head-count regimes. Flow: +- `_calc_cp_seqs` picks `max_local_chunks` from a latency model (`L_cp* ∝ √(B·H·L_c/P)`, ×3 empirical factor, rounded to a power of 2; floored at 4 for pipelining) and decides `use_cp` based on whether `B·H` already saturates the SMs. +- It exploits the GDN gate's exponential decay: `get_warmup_chunks` finds how many preceding chunks each split needs to "warm up" a usable initial state (gate threshold `-10.0`); `fused_gdr_h` computes those carry states; `correct_initial_states` stitches them back. +- Output threads `cp_seq_map` (cp-batch → raw-batch index) and `raw_cu_seqlens` into the fused kernel via its `is_cp` flag, so the kernel writes final states to the correct raw-batch slots. + +### Utilities + +- `flash_qla/ops/utils/` — `chunk_local_cumsum` (`cumsum.py`), `group_reduce_vector` (`group_reduce.py`); both are TileLang kernels. +- `flash_qla/utils/` — `l2norm` (`math.py`); varlen packing helpers `pack`/`unpack`/`pad_and_reshape`/`fill_last_chunk_of_g` (`pack.py`); `profile` (`profiler.py`); `tensor_cache`/`prepare_chunk_indices`/`prepare_chunk_offsets` (`index.py`). + +## Conventions and gotchas + +- **Kernel definition pattern**: each `@tilelang.jit`-decorated function is a *factory* taking shape/dtype/flag parameters and returning a `@T.prim_func`. The plain-Python wrapper (e.g. `fused_gdr_fwd`) computes shapes/dtypes/flags, picks `block_DV` from the grid size vs SM count, calls the factory, then invokes the returned kernel. Every distinct parameter combination triggers a fresh JIT compile. +- **Tensor layouts** are fixed: `q`/`k` are `[B, T, Hg, 128]`, `v`/`o` are `[B, T, H, 128]`, `g`/`beta` are `[B, T, H]`, states are `[B, H, 128, 128]`. `Hg` = K/Q heads, `H` = V heads, `Hg ≤ H`. +- **Varlen mode**: `batch_size` must be 1, inputs are flattened/packed along `T`, and `cu_seqlens` marks sequence boundaries (see `pack`/`unpack`). +- **Kernel names are an API for the tests/benchmark**: `profile()` keys timings by kernel name, and `test_gdr.py`/`bench_gated_delta_rule.py` index specific names (e.g. `tilelang_fused_chunk_gdr_fwd_kernel_kernel`, `tilelang_kkt_solve_kernel_kernel`, `tilelang_prepare_h_kernel_kernel`, `tilelang_correct_h0_kernel_kernel`). Renaming a kernel silently breaks those lookups. +- **The fused kernels hard-code their thread geometry**: 512 threads split into one producer + three consumer warpgroups, with manual named-barrier `arrive_count`s (96/128/256/384/416) tied to that layout. These are not tunable knobs — they encode the warp-specialization schedule. +- **`tensor_cache`** (`index.py`) is an identity-based LRU (keyed on `id()`, size 256), used for CP/index helper tensors — it caches by object identity, not value. +- **`backward` gradient count is guarded statically**: `test_function_signature.py` parses the source with `ast` to assert `ChunkGatedDeltaRuleFunction.backward` returns exactly one gradient per non-`ctx` forward input (a real past bug — PR #10). Runs without a GPU. If you change `forward`'s signature, update `backward`'s return tuple in lockstep. +- Backward `dg` must be fp32 (asserted in the orchestration layer before the reverse cumsum). diff --git a/benchmark/bench_head_batch.py b/benchmark/bench_head_batch.py new file mode 100644 index 00000000..cec8f864 --- /dev/null +++ b/benchmark/bench_head_batch.py @@ -0,0 +1,54 @@ +# benchmark/bench_head_batch.py +# Head-batched GQA vs per-head, core decode kernel. Isolates the head-batch lever (q/k load +# dedup) in the regime where it could help: high GQA ratio + saturated grid. +# +# MEASURED (H100, B in {64,128,256}): grp=2 (Hk4Hv8) = 0.98-0.99x (neutral); grp=4 (Hk2Hv8) +# = 0.74-0.87x (a real regression). The row-stack trades CTA count (B*H -> B*Hg) and uses +# bigger 512-thread CTAs at grp=4, which outweighs the only saving (q/k LOAD dedup, sub-1%). +# Conclusion: head-batch is neutral-to-worse on this memory-bound kernel -> auto stays OFF +# (default); the flag is kept forceable for experimentation. See the decode spec section 7. +import time +import torch + +from flash_qla import recurrent_gated_delta_rule +from flash_qla.utils import l2norm + + +def mk(B, D, Hk, Hv): + q = l2norm(torch.randn(B, D, Hk, 128, device="cuda", dtype=torch.bfloat16)) + k = l2norm(torch.randn(B, D, Hk, 128, device="cuda", dtype=torch.bfloat16)) + v = torch.randn(B, D, Hv, 128, device="cuda", dtype=torch.bfloat16) + g = torch.nn.functional.logsigmoid(torch.randn(B, D, Hv, device="cuda")) / 16 + beta = torch.randn(B, D, Hv, device="cuda").sigmoid() + return q, k, v, g, beta + + +def timed(fn, iters=300, warmup=50): + for _ in range(warmup): + fn() + torch.cuda.synchronize() + t0 = time.perf_counter() + for _ in range(iters): + fn() + torch.cuda.synchronize() + return (time.perf_counter() - t0) / iters * 1e6 # us + + +def run(): + sm = torch.cuda.get_device_properties().multi_processor_count + print(f"device={torch.cuda.get_device_name()} SMs={sm}\n") + print(f"{'Hk':>3} {'Hv':>3} {'grp':>3} {'B':>5} {'per_head(us)':>13} {'head_batch(us)':>15} {'speedup':>8}") + for Hk, Hv in [(2, 8), (4, 8)]: + for B in [64, 128, 256]: + q, k, v, g, beta = mk(B, 8, Hk, Hv) + kw = dict(scale=128 ** -0.5, output_final_state=True) + f_hb = lambda: recurrent_gated_delta_rule(q, k, v, g, beta, head_batch=True, **kw) + f_ph = lambda: recurrent_gated_delta_rule(q, k, v, g, beta, head_batch=False, **kw) + f_hb(); f_ph() # build once (exclude JIT from timing) + us_ph = timed(f_ph) + us_hb = timed(f_hb) + print(f"{Hk:>3} {Hv:>3} {Hv // Hk:>3} {B:>5} {us_ph:>13.1f} {us_hb:>15.1f} {us_ph / us_hb:>7.3f}x") + + +if __name__ == "__main__": + run() diff --git a/benchmark/bench_prepass.py b/benchmark/bench_prepass.py new file mode 100644 index 00000000..a05dcec6 --- /dev/null +++ b/benchmark/bench_prepass.py @@ -0,0 +1,150 @@ +# Copyright (c) 2026 The Qwen team, Alibaba Group. +# Licensed under The MIT License [see LICENSE for details] +"""H1 crossover benchmark: the in-kernel-gated verify kernel (variant A) vs the dedup +prepass+host-gated path, end-to-end through recurrent_gated_delta_rule_verify(fuse_gating=True), +with prepass forced on/off. Eager event timing (conservative: the prepass's 2nd launch pays full +launch latency eagerly; under CUDA-graph replay it is cheaper, so an eager win is a real win). +Reports speedup A/prepass and whether the auto regime-gate (should_use_prepass) agrees with the +empirical winner -- used to calibrate PREPASS_MIN_T / PREPASS_CTA_FACTOR. +""" +import torch + +from flash_qla import recurrent_gated_delta_rule_verify +from flash_qla.ops.gated_delta_rule.fused_recurrent import should_use_prepass +from flash_qla.ops.gated_delta_rule.fused_recurrent.hopper.fused_recurrent_verify import ( + PREPASS_MIN_T, + PREPASS_MIN_WORK, +) + + +def _time(fn, iters=100, warmup=50): + for _ in range(warmup): + fn() + torch.cuda.synchronize() + s, e = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True) + s.record() + for _ in range(iters): + fn() + e.record() + torch.cuda.synchronize() + return s.elapsed_time(e) / iters * 1e3 # us + + +def _time_graph(fn, iters=100, warmup=20): + # production path: capture into a CUDA graph, time replay (the 2nd-launch tax shrinks to a + # graph node -- eager timing over-penalizes it). Warmup populates the persistent prepass scratch. + st = torch.cuda.Stream() + st.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(st): + for _ in range(3): + fn() + torch.cuda.current_stream().wait_stream(st) + g = torch.cuda.CUDAGraph() + with torch.cuda.graph(g): + fn() + for _ in range(warmup): + g.replay() + torch.cuda.synchronize() + s, e = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True) + s.record() + for _ in range(iters): + g.replay() + e.record() + torch.cuda.synchronize() + return s.elapsed_time(e) / iters * 1e3 # us + + +def _inputs(N, T, Hk, Hv, seed=2025): + torch.manual_seed(seed) + tot = N * T + A_log = torch.randn(Hv, dtype=torch.float32, device="cuda") + dt_bias = torch.randn(Hv, dtype=torch.float32, device="cuda") + a = torch.randn(1, tot, Hv, dtype=torch.bfloat16, device="cuda") + b = torch.randn(1, tot, Hv, dtype=torch.bfloat16, device="cuda") + q = torch.randn(1, tot, Hk, 128, dtype=torch.bfloat16, device="cuda") + k = torch.randn(1, tot, Hk, 128, dtype=torch.bfloat16, device="cuda") + v = torch.randn(1, tot, Hv, 128, dtype=torch.bfloat16, device="cuda") + pool = torch.randn(N, Hv, 128, 128, dtype=torch.bfloat16, device="cuda") + cu = torch.arange(0, tot + 1, T, dtype=torch.int32, device="cuda") + idx = torch.arange(N, dtype=torch.int32, device="cuda") + ibuf = torch.zeros(N + 1, T, Hv, 128, 128, dtype=torch.bfloat16, device="cuda") + o = torch.empty(1, tot, Hv, 128, dtype=torch.bfloat16, device="cuda") + return dict(A_log=A_log, a=a, dt_bias=dt_bias, q=q, k=k, v=v, b=b, ssm_states=pool, + cache_indices=idx, query_start_loc=cu, intermediate_states_buffer=ibuf, + intermediate_state_indices=idx, o=o) + + +def bench(N, T, Hk, Hv, tag): + kw = _inputs(N, T, Hk, Hv) + fA = lambda: recurrent_gated_delta_rule_verify(fuse_gating=True, prepass=False, **kw) + fP = lambda: recurrent_gated_delta_rule_verify(fuse_gating=True, prepass=True, **kw) + + # build + parity (the two paths must agree within in-kernel-gating tolerance) + oA = fA().clone() + oP = fP().clone() + err = ((oA.float() - oP.float()).abs().max() / oP.float().abs().max().clamp_min(1e-6)).item() + + tA = _time(fA) + tP = _time(fP) + sp = tA / tP if tP > 0 else 0.0 + auto = should_use_prepass(N, Hv, N * T) + win = "WIN " if sp >= 1.05 else ("loss" if sp <= 0.97 else "neut") + # MISCAL only if the gate picks the measurably-WRONG path: PP on a clear loss, or A on a clear win. + # In the neutral band [0.97,1.05) either choice is fine (A preferred -> no extra launch). + agree = "ok" + if auto and sp <= 0.97: + agree = "MISCAL" # gate fired prepass but it regressed + elif (not auto) and sp >= 1.05: + agree = "MISCAL" # gate kept A but prepass would have won + print(f" [{tag}] N={N:<4d} T={T:<2d} Hk={Hk:<2d} Hv={Hv:<2d} " + f"A={tA:8.1f}us prepass={tP:8.1f}us speedup={sp:4.2f}x [{win}] " + f"auto={'PP' if auto else 'A '} [{agree}] parity={err:.4f}") + + +def bench_graph(N, T, Hk, Hv): + # CUDA-graph (production) timing of variant A vs the prepass path, WITH gate calibration: + # print work + the current gate's auto decision + MISCAL (gate picks the measurably-wrong path). + kw = _inputs(N, T, Hk, Hv) + fA = lambda: recurrent_gated_delta_rule_verify(fuse_gating=True, prepass=False, **kw) + fP = lambda: recurrent_gated_delta_rule_verify(fuse_gating=True, prepass=True, **kw) + fA(); fP() # build kernels + tA = _time_graph(fA) + tP = _time_graph(fP) + sp = tA / tP if tP > 0 else 0.0 + work = Hv * (N + N * T) + auto = should_use_prepass(N, Hv, N * T) + win = "WIN " if sp >= 1.05 else ("loss" if sp <= 0.97 else "neut") + agree = "ok" + if auto and sp <= 0.97: + agree = "MISCAL" # gate fired prepass but it regressed under graphs + elif (not auto) and sp >= 1.05: + agree = "MISCAL" # gate kept A but prepass would have won under graphs + print(f" [graph] N={N:<4d} T={T:<2d} Hv={Hv:<2d} work={work:<6d} " + f"A={tA:8.1f}us prepass={tP:8.1f}us speedup={sp:4.2f}x [{win}] " + f"auto={'PP' if auto else 'A '} [{agree}]") + + +def main(): + sm = torch.cuda.get_device_properties().multi_processor_count + print(f"device={torch.cuda.get_device_name()} SMs={sm} (TARGET_CTAS={int(sm*0.7)})") + print(f"gate: PREPASS_MIN_T={PREPASS_MIN_T} PREPASS_MIN_WORK={PREPASS_MIN_WORK} " + f"(work=Hv*N*(1+T); prepass if t_avg>=MIN_T AND work>=MIN_WORK)\n") + # PRODUCTION path is CUDA-graph: the prepass's 2nd launch shrinks to a graph node, so the + # loss->win crossover sits at SMALLER work than the eager-calibrated gate assumes. Sweep the + # small-N crossover zone under graphs to find where the prepass actually starts winning, and + # flag where the current gate mis-fires. The N=1/T=12 floor (work=416) MUST stay loss/off. + print("== CUDA-graph (production) crossover calibration, T in {4,8,12} (Hk=16,Hv=32) ==") + for N in (1, 2, 4, 8, 16, 32, 64, 256): + for T in (4, 8, 12): + bench_graph(N, T, 16, 32) + print() + print("== T=1 (decode path) under graphs -- t_avg=64,T>=12 the +~15us launch is <1% of the 100us-1.6ms runtime, so eager is reliable+fair, matching bench_vs_fla).""" +import importlib.util +import os +import torch + +from flash_qla import recurrent_gated_delta_rule_verify + +_FLA_PATH = os.environ.get( + "FLA_KERNEL_PATH", + "/root/netra-server/python/sglang/srt/layers/attention/fla/fused_sigmoid_gating_recurrent.py", +) +_spec = importlib.util.spec_from_file_location("fla_sigmoid_gating", _FLA_PATH) +_mod = importlib.util.module_from_spec(_spec) +_spec.loader.exec_module(_mod) +fla_verify = _mod.fused_sigmoid_gating_delta_rule_update + +PEAK_TBS = 3.35 + + +def _time(fn, iters=50, warmup=25): + for _ in range(warmup): + fn() + torch.cuda.synchronize() + s, e = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True) + s.record() + for _ in range(iters): + fn() + e.record() + torch.cuda.synchronize() + return s.elapsed_time(e) / iters * 1e3 # us + + +def _inputs(N, T, Hk, Hv, seed=2025): + torch.manual_seed(seed) + tot = N * T + A_log = torch.randn(Hv, dtype=torch.float32, device="cuda") + dt_bias = torch.randn(Hv, dtype=torch.float32, device="cuda") + a = torch.randn(1, tot, Hv, dtype=torch.bfloat16, device="cuda") + b = torch.randn(1, tot, Hv, dtype=torch.bfloat16, device="cuda") + q = torch.randn(1, tot, Hk, 128, dtype=torch.bfloat16, device="cuda") + k = torch.randn(1, tot, Hk, 128, dtype=torch.bfloat16, device="cuda") + v = torch.randn(1, tot, Hv, 128, dtype=torch.bfloat16, device="cuda") + pool = torch.randn(N, Hv, 128, 128, dtype=torch.bfloat16, device="cuda") + cu = torch.arange(0, tot + 1, T, dtype=torch.int32, device="cuda") + idx = torch.arange(N, dtype=torch.int32, device="cuda") + return A_log, dt_bias, a, b, q, k, v, pool, cu, idx + + +def bench(N, T, Hk, Hv): + A_log, dt_bias, a, b, q, k, v, pool, cu, idx = _inputs(N, T, Hk, Hv) + tot = N * T + o = torch.empty(1, tot, Hv, 128, dtype=torch.bfloat16, device="cuda") + ib_fla = torch.zeros(N + 1, T, Hv, 128, 128, dtype=torch.bfloat16, device="cuda") + ib_a = torch.zeros(N + 1, T, Hv, 128, 128, dtype=torch.bfloat16, device="cuda") + ib_pp = torch.zeros(N + 1, T, Hv, 128, 128, dtype=torch.bfloat16, device="cuda") + + fla = lambda: fla_verify( + A_log, a, dt_bias, 1.0, 20.0, q, k, v, b, pool, idx, scale=None, + use_qk_l2norm_in_kernel=True, cu_seqlens=cu, is_kda=False, disable_state_update=True, + intermediate_states_buffer=ib_fla, intermediate_state_indices=idx, cache_steps=T, + retrieve_parent_token=None) + kw = dict(ssm_states=pool, cache_indices=idx, query_start_loc=cu, + intermediate_state_indices=idx, o=o, fuse_gating=True, disable_state_update=True) + varA = lambda: recurrent_gated_delta_rule_verify( + A_log, a, dt_bias, q, k, v, b, intermediate_states_buffer=ib_a, prepass=False, **kw) + pp = lambda: recurrent_gated_delta_rule_verify( + A_log, a, dt_bias, q, k, v, b, intermediate_states_buffer=ib_pp, prepass=True, **kw) + + o_fla = fla().clone() + varA(); o_a = o.clone() + pp(); o_pp = o.clone() + er = lambda x, y: ((x.float() - y.float()).abs().max() / y.float().abs().max().clamp_min(1e-6)).item() + parity = f"A:{er(o_a, o_fla):.4f} PP:{er(o_pp, o_fla):.4f}" + + t_fla, t_a, t_pp = _time(fla), _time(varA), _time(pp) + sb = N * Hv * (1 + T) * 128 * 128 * 2 + print(f" N={N:<4d} T={T:<2d} Hk={Hk} Hv={Hv} " + f"FLA {t_fla:8.1f}us ({sb/(t_fla*1e-6)/1e9:5.0f}GB/s) " + f"varA {t_a:8.1f}us ({t_fla/t_a:4.2f}x) " + f"prepass {t_pp:8.1f}us ({t_fla/t_pp:4.2f}x vs FLA, {t_a/t_pp:4.2f}x vs A) [parity o {parity}]") + + +def main(): + print(f"device: {torch.cuda.get_device_name()} | GDN verify high-batch: FLA vs FlashQLA varA vs prepass") + for N in (64, 128, 256): + for T in (4, 12): + bench(N, T, 16, 32) + + +if __name__ == "__main__": + main() diff --git a/benchmark/bench_recurrent_gdr.py b/benchmark/bench_recurrent_gdr.py new file mode 100644 index 00000000..fe5fc9d1 --- /dev/null +++ b/benchmark/bench_recurrent_gdr.py @@ -0,0 +1,62 @@ +# Copyright (c) 2026 The Qwen team, Alibaba Group. +# Licensed under The MIT License [see LICENSE for details] +"""Benchmark the FlashQLA GDN verify kernel (memory-bound): wall time + achieved HBM +bandwidth vs the theoretical state-I/O floor, across SGLang-relevant regimes.""" +import torch + +from flash_qla import fused_recurrent_gdr_verify_fwd +from flash_qla.utils import l2norm + +PEAK_TBS = 3.35 # H100 HBM3 ~3.35 TB/s + + +def _time(fn, iters=50, warmup=10): + for _ in range(warmup): + fn() + torch.cuda.synchronize() + start, end = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True) + start.record() + for _ in range(iters): + fn() + end.record() + torch.cuda.synchronize() + return start.elapsed_time(end) / iters # ms + + +def bench(N, H, Hg, D, dtype=torch.bfloat16, store_intermediate=True): + total = N * D + q = l2norm(torch.randn(1, total, Hg, 128, device="cuda", dtype=dtype)) + k = l2norm(torch.randn(1, total, Hg, 128, device="cuda", dtype=dtype)) + v = torch.randn(1, total, H, 128, device="cuda", dtype=dtype) + g = torch.nn.functional.logsigmoid(torch.randn(1, total, H, device="cuda")) / 16 + beta = torch.randn(1, total, H, device="cuda").sigmoid() + pool = torch.randn(N, H, 128, 128, device="cuda", dtype=dtype) + cu = torch.arange(0, total + 1, D, dtype=torch.int32, device="cuda") + si = torch.arange(N, dtype=torch.int32, device="cuda") + ci = torch.arange(N, dtype=torch.int32, device="cuda") + ibuf = (torch.zeros(N + 1, D, H, 128, 128, device="cuda", dtype=dtype) + if store_intermediate else None) + o = torch.empty(1, total, H, 128, device="cuda", dtype=dtype) + + fn = lambda: fused_recurrent_gdr_verify_fwd( + q, k, v, g, beta, pool, si, cu, ibuf, ci, o, disable_state_update=True) + ms = _time(fn) + + es = 2 # bf16 state element size + # dominant state I/O: 1 gather + D intermediate writes per (request, head) + state_bytes = N * H * (1 + D) * 128 * 128 * es + io_bytes = (q.numel() + k.numel() + v.numel() + o.numel()) * es # tiny qkvo + gbs = (state_bytes + io_bytes) / (ms * 1e-3) / 1e9 + print(f" N={N:<4d} H={H} Hg={Hg} D={D:<2d} {ms*1e3:8.1f} us " + f"{gbs:7.1f} GB/s ({100*gbs/1000/PEAK_TBS:4.1f}% peak) state={state_bytes/1e6:6.1f} MB") + + +if __name__ == "__main__": + print(f"device: {torch.cuda.get_device_name()} | verify kernel (bf16 pool, per-token intermediates)") + print("server / batched decode (H=32, Hg=16):") + for N in (8, 64, 256, 512): + for D in (1, 4, 12): + bench(N, 32, 16, D) + print("single-request / TP (N=1, varying H = TP1..TP8):") + for H in (64, 32, 16, 8): + bench(1, H, max(1, H // 4), 12) diff --git a/benchmark/bench_vs_fla.py b/benchmark/bench_vs_fla.py new file mode 100644 index 00000000..a544c63a --- /dev/null +++ b/benchmark/bench_vs_fla.py @@ -0,0 +1,126 @@ +# Copyright (c) 2026 The Qwen team, Alibaba Group. +# Licensed under The MIT License [see LICENSE for details] +"""Speed benchmark: FlashQLA GDN verify kernel vs the FLA/Triton sigmoid-gating verify +kernel (vendored in netra-server) on the SGLang DFlash target_verify path. + +Apples-to-apples: both run the in-kernel-gated, no-commit verify with per-token +intermediate-state caching over T draft tokens; both use a bf16 state pool/buffer +(SGLANG_MAMBA_SSM_DTYPE). FlashQLA is V-major, FLA is K-major -- each gets its native +layout for timing; the correctness gate transposes one side before comparing. +""" +import torch + +from flash_qla.ops.gated_delta_rule.fused_recurrent.hopper.fused_recurrent_verify import ( + fused_recurrent_gdr_verify_gated_fwd, +) + +# Load the FLA kernel file DIRECTLY (it's pure torch+triton) to bypass sglang/__init__.py, +# which runs the heavy SGLang frontend chain (hf patches -> orjson/transformers/...). +import importlib.util +import os + +_FLA_PATH = os.environ.get( + "FLA_KERNEL_PATH", + "/root/netra-server/python/sglang/srt/layers/attention/fla/fused_sigmoid_gating_recurrent.py", +) +try: + _spec = importlib.util.spec_from_file_location("fla_sigmoid_gating", _FLA_PATH) + _mod = importlib.util.module_from_spec(_spec) + _spec.loader.exec_module(_mod) + fla_verify = _mod.fused_sigmoid_gating_delta_rule_update + FLA_OK = True +except Exception as e: # noqa: BLE001 + print("FLA load FAILED:", repr(e)) + FLA_OK = False + +PEAK_TBS = 3.35 # H100 HBM3 + + +def _time(fn, iters=50, warmup=25): + for _ in range(warmup): + fn() + torch.cuda.synchronize() + s, e = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True) + s.record() + for _ in range(iters): + fn() + e.record() + torch.cuda.synchronize() + return s.elapsed_time(e) / iters # ms + + +def _inputs(N, T, Hk, Hv, seed=2025): + torch.manual_seed(seed) + tot = N * T + A_log = torch.randn(Hv, dtype=torch.float32, device="cuda") + dt_bias = torch.randn(Hv, dtype=torch.float32, device="cuda") + a = torch.randn(1, tot, Hv, dtype=torch.bfloat16, device="cuda") + b = torch.randn(1, tot, Hv, dtype=torch.bfloat16, device="cuda") + q = torch.randn(1, tot, Hk, 128, dtype=torch.bfloat16, device="cuda") + k = torch.randn(1, tot, Hk, 128, dtype=torch.bfloat16, device="cuda") + v = torch.randn(1, tot, Hv, 128, dtype=torch.bfloat16, device="cuda") + # Both kernels read the pool V-major [.,HV,V,K] (FLA indexes h0 at offset o_v*K+o_k). + pool = torch.randn(N, Hv, 128, 128, dtype=torch.bfloat16, device="cuda") + cu = torch.arange(0, tot + 1, T, dtype=torch.int32, device="cuda") + idx = torch.arange(N, dtype=torch.int32, device="cuda") + return A_log, dt_bias, a, b, q, k, v, pool, cu, idx + + +def _call_fla(A_log, dt_bias, a, b, q, k, v, pool, cu, idx, ibuf, T): + return fla_verify( + A_log, a, dt_bias, 1.0, 20.0, q, k, v, b, pool, idx, + scale=None, use_qk_l2norm_in_kernel=True, cu_seqlens=cu, is_kda=False, + disable_state_update=True, intermediate_states_buffer=ibuf, + intermediate_state_indices=idx, cache_steps=T, retrieve_parent_token=None, + ) + + +def _call_fqla(A_log, dt_bias, a, b, q, k, v, pool, cu, idx, ibuf, o): + fused_recurrent_gdr_verify_gated_fwd( + q, k, v, a, b, A_log, dt_bias, pool, idx, cu, ibuf, idx, o, + scale=None, disable_state_update=True, + ) + return o + + +def bench(N, T, Hk, Hv, check=True): + A_log, dt_bias, a, b, q, k, v, pool, cu, idx = _inputs(N, T, Hk, Hv) + o_fqla = torch.empty(1, N * T, Hv, 128, dtype=torch.bfloat16, device="cuda") + ibuf_fla = torch.zeros(N + 1, T, Hv, 128, 128, dtype=torch.bfloat16, device="cuda") + ibuf_fqla = torch.zeros(N + 1, T, Hv, 128, 128, dtype=torch.bfloat16, device="cuda") + + o_fla = _call_fla(A_log, dt_bias, a, b, q, k, v, pool, cu, idx, ibuf_fla, T) + _call_fqla(A_log, dt_bias, a, b, q, k, v, pool, cu, idx, ibuf_fqla, o_fqla) + + if check: + o_err = (o_fqla.float() - o_fla.float()).abs().max() / o_fla.float().abs().max().clamp_min(1e-6) + ib_err = (ibuf_fqla.float() - ibuf_fla.float()).abs().max() / \ + ibuf_fla.float().abs().max().clamp_min(1e-6) # both V-major, compare directly + tag = "OK" if (o_err < 0.04 and ib_err < 0.04) else "MISMATCH" + print(f" [parity {tag}] o_err={o_err:.4f} ibuf_err={ib_err:.4f}") + + t_fqla = _time(lambda: _call_fqla(A_log, dt_bias, a, b, q, k, v, pool, cu, idx, ibuf_fqla, o_fqla)) + t_fla = _time(lambda: _call_fla(A_log, dt_bias, a, b, q, k, v, pool, cu, idx, ibuf_fla, T)) + + es = 2 + state_bytes = N * Hv * (1 + T) * 128 * 128 * es + gbs_fqla = state_bytes / (t_fqla * 1e-3) / 1e9 + gbs_fla = state_bytes / (t_fla * 1e-3) / 1e9 + sp = t_fla / t_fqla + print(f" N={N:<4d} Hv={Hv} Hk={Hk} T={T:<2d} " + f"FlashQLA {t_fqla*1e3:8.1f}us ({gbs_fqla:6.0f} GB/s) " + f"FLA {t_fla*1e3:8.1f}us ({gbs_fla:6.0f} GB/s) " + f"speedup {sp:4.2f}x") + + +if __name__ == "__main__": + if not FLA_OK: + raise SystemExit("FLA baseline unavailable -- cannot compare") + print(f"device: {torch.cuda.get_device_name()} | GDN verify: FlashQLA vs FLA (bf16 pool, per-token states)") + print("server / batched (Hv=32, Hk=16):") + for N in (8, 64, 256): + for T in (1, 4, 12): + bench(N, T, 16, 32) + print("single-request / TP (N=1, Hv=TP1..TP8):") + for Hv in (64, 32, 16, 8): + bench(1, 12, max(1, Hv // 4), Hv) diff --git a/benchmark/probe_h1_ceiling.py b/benchmark/probe_h1_ceiling.py new file mode 100644 index 00000000..2c1b921a --- /dev/null +++ b/benchmark/probe_h1_ceiling.py @@ -0,0 +1,121 @@ +# Copyright (c) 2026 The Qwen team, Alibaba Group. +# Licensed under The MIT License [see LICENSE for details] +"""H1 verify-first probe (NO new kernel): measure the upper bound on what a gating+l2norm +pre-pass can save for the in-kernel-gated verify kernel. + +Idea: the repo already ships two equivalent verify kernels -- + (A) in-kernel-gated : recomputes g/beta + qk-l2norm INSIDE the hot loop, once per + (token, V-head, V-tile) -> n_vt * grp redundancy for l2norm, + n_vt redundancy for gating. + (B) host-gated : reads PRE-computed g/beta/q_n/k_n; hot loop has NO transcendentals. + +On identical raw inputs, time(A) - time(B) (with B's precompute done OUTSIDE the loop) is the +exact ceiling on H1's main-kernel saving. The realized H1 win = that ceiling minus the +(amortized-once) pre-pass cost. We also report the naive torch precompute cost for reference. +""" +import torch + +from flash_qla.utils import l2norm +from flash_qla.ops.gated_delta_rule.fused_recurrent import gdn_sigmoid_gate +from flash_qla.ops.gated_delta_rule.fused_recurrent.hopper.fused_recurrent_verify import ( + fused_recurrent_gdr_verify_fwd, + fused_recurrent_gdr_verify_gated_fwd, +) + +PEAK_TBS = 3.35 # H100 HBM3 + + +def _time(fn, iters=100, warmup=50): + for _ in range(warmup): + fn() + torch.cuda.synchronize() + s, e = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True) + s.record() + for _ in range(iters): + fn() + e.record() + torch.cuda.synchronize() + return s.elapsed_time(e) / iters * 1e3 # us + + +def _inputs(N, T, Hk, Hv, seed=2025): + torch.manual_seed(seed) + tot = N * T + A_log = torch.randn(Hv, dtype=torch.float32, device="cuda") + dt_bias = torch.randn(Hv, dtype=torch.float32, device="cuda") + a = torch.randn(1, tot, Hv, dtype=torch.bfloat16, device="cuda") + b = torch.randn(1, tot, Hv, dtype=torch.bfloat16, device="cuda") + q = torch.randn(1, tot, Hk, 128, dtype=torch.bfloat16, device="cuda") + k = torch.randn(1, tot, Hk, 128, dtype=torch.bfloat16, device="cuda") + v = torch.randn(1, tot, Hv, 128, dtype=torch.bfloat16, device="cuda") + pool = torch.randn(N, Hv, 128, 128, dtype=torch.bfloat16, device="cuda") + cu = torch.arange(0, tot + 1, T, dtype=torch.int32, device="cuda") + idx = torch.arange(N, dtype=torch.int32, device="cuda") + return A_log, dt_bias, a, b, q, k, v, pool, cu, idx + + +def probe(N, T, Hk, Hv, tag): + A_log, dt_bias, a, b, q, k, v, pool, cu, idx = _inputs(N, T, Hk, Hv) + tot = N * T + o = torch.empty(1, tot, Hv, 128, dtype=torch.bfloat16, device="cuda") + ibuf = torch.zeros(N + 1, T, Hv, 128, 128, dtype=torch.bfloat16, device="cuda") + + # (B) host-gated precompute -- done ONCE, outside the timing loop (the H1 pre-pass stand-in) + def precompute(): + g_ref, beta_ref = gdn_sigmoid_gate(A_log, a, dt_bias, b) + return l2norm(q), l2norm(k), g_ref, beta_ref + + q_n, k_n, g_ref, beta_ref = precompute() + + fA = lambda: fused_recurrent_gdr_verify_gated_fwd( + q, k, v, a, b, A_log, dt_bias, pool, idx, cu, ibuf, idx, o, + scale=None, disable_state_update=True) + fB = lambda: fused_recurrent_gdr_verify_fwd( + q_n, k_n, v, g_ref, beta_ref, pool, idx, cu, ibuf, idx, o, + scale=None, disable_state_update=True) + + # build + parity (A vs B must match within bf16 noise; same math, one in-kernel one host) + fA() + oA = o.clone() + fB() + oB = o.clone() + err = ((oA.float() - oB.float()).abs().max() / oB.float().abs().max().clamp_min(1e-6)).item() + + tA = _time(fA) + tB = _time(fB) + tPre = _time(precompute, iters=50, warmup=20) # naive torch pre-pass cost (reference) + + delta = tA - tB + pct = 100.0 * delta / tA if tA > 0 else 0.0 + # ceiling speedup if pre-pass were free, and realistic if pre-pass == naive torch cost + ceil_sp = tA / tB if tB > 0 else 0.0 + realistic_sp = tA / (tB + tPre) if (tB + tPre) > 0 else 0.0 + state_bytes = N * Hv * (1 + T) * 128 * 128 * 2 + gbsA = state_bytes / (tA * 1e-6) / 1e9 + block_DV = 64 if (N * Hv) * 2 >= int(torch.cuda.get_device_properties().multi_processor_count * 0.7) else 32 + n_vt = 128 // block_DV + grp = Hv // Hk + print(f" [{tag}] N={N:<4d} T={T:<2d} Hk={Hk:<2d} Hv={Hv:<2d} block_DV={block_DV} n_vt={n_vt} grp={grp}") + print(f" gated(A)={tA:8.1f}us host(B)={tB:8.1f}us delta={delta:7.1f}us ({pct:5.1f}% of A) " + f"ceil={ceil_sp:4.2f}x torch_prepass={tPre:7.1f}us realistic={realistic_sp:4.2f}x " + f"A_bw={gbsA:5.0f}GB/s parity={err:.4f}") + + +def main(): + sm = torch.cuda.get_device_properties().multi_processor_count + print(f"device={torch.cuda.get_device_name()} SMs={sm} (TARGET_CTAS={int(sm*0.7)})") + print("\n== latency-bound: single request, TP1..TP8 (N=1, T=12) ==") + for Hv in (64, 32, 16, 8): + probe(1, 12, max(1, Hv // 4), Hv, "lat") + print("\n== latency-bound: small batch, short draft (N in {1,4}, T in {1,4}) ==") + for N in (1, 4): + for T in (1, 4): + probe(N, T, 16, 32, "lat") + print("\n== bandwidth-bound: server batched (Hv=32, Hk=16, N in {64,256}) ==") + for N in (64, 256): + for T in (1, 4, 12): + probe(N, T, 16, 32, "bw") + + +if __name__ == "__main__": + main() diff --git a/benchmark/probe_h2_blockdv_crossover.py b/benchmark/probe_h2_blockdv_crossover.py new file mode 100644 index 00000000..46e49cab --- /dev/null +++ b/benchmark/probe_h2_blockdv_crossover.py @@ -0,0 +1,72 @@ +# Copyright (c) 2026 The Qwen team, Alibaba Group. +# Licensed under The MIT License [see LICENSE for details] +"""Crossover for the decode-kernel store_final_state tile fix: block_DV=128 makes the K-major +transposed final-state write coalesced (~2x faster at large batch) but yields fewer CTAs (n_vt=1) +-> may be occupancy-starved at small batch. Sweep B (and H) to find where block_DV=128 beats the +as-built block_DV=64, so the dispatch can gate on it. store_final_state=True (the decode default).""" +import torch + +from flash_qla.ops.gated_delta_rule.fused_recurrent.hopper.fused_recurrent_fwd import ( + tilelang_fused_recurrent_gdr_fwd, +) + +SM = torch.cuda.get_device_properties().multi_processor_count +TARGET = int(SM * 0.7) + + +def _build(H, block_DV): + return tilelang_fused_recurrent_gdr_fwd( + H, H, 128, 128, 128 ** -0.5, + accum_dtype="float32", qkva_dtype=torch.bfloat16, g_dtype=torch.float32, b_dtype=torch.float32, + h0_dtype=torch.float32, ht_dtype=torch.float32, o_dtype=torch.bfloat16, seqlen_dtype=torch.int32, + use_initial_state=False, store_final_state=True, has_seqlens=False, + block_DV=block_DV, threads=128, + ) + + +def _time(fn, iters=50, warmup=25): + for _ in range(warmup): + fn() + torch.cuda.synchronize() + s, e = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True) + s.record() + for _ in range(iters): + fn() + e.record() + torch.cuda.synchronize() + return s.elapsed_time(e) / iters * 1e3 + + +def bench(B, T, H): + torch.manual_seed(0) + q = torch.randn(B, T, H, 128, device="cuda", dtype=torch.bfloat16) + k = torch.randn(B, T, H, 128, device="cuda", dtype=torch.bfloat16) + v = torch.randn(B, T, H, 128, device="cuda", dtype=torch.bfloat16) + g = torch.nn.functional.logsigmoid(torch.randn(B, T, H, device="cuda")) / 16 + beta = torch.randn(B, T, H, device="cuda").sigmoid() + h0 = torch.empty(B, H, 128, 128, device="cuda", dtype=torch.float32) + ht = torch.empty(B, H, 128, 128, device="cuda", dtype=torch.float32) + sl = torch.empty(B, dtype=torch.int32, device="cuda") + o = torch.empty_like(v) + res = {} + for bdv in (64, 128): + kern = _build(H, bdv) + res[bdv] = _time(lambda kern=kern: kern(q, k, v, g, beta, h0, sl, o, ht)) + sp = res[64] / res[128] + asbuilt = 64 if B * H * 2 >= TARGET else 32 + pick = "128" if sp >= 1.05 else ("64 " if sp <= 0.97 else "tie") + print(f" B={B:<4d} H={H:<3d} (B*H={B*H:<5d}) bdv64={res[64]:8.1f}us bdv128={res[128]:8.1f}us " + f"128-speedup={sp:4.2f}x best={pick} (as-built picks {asbuilt})") + + +def main(): + print(f"device: {torch.cuda.get_device_name()} SMs={SM} TARGET={TARGET}\n") + for H in (32, 16, 8): + print(f"== H={H}, T=12, store_final_state=True ==") + for B in (1, 2, 4, 8, 16, 32, 64, 128, 256): + bench(B, 12, H) + print() + + +if __name__ == "__main__": + main() diff --git a/benchmark/probe_h2_reduction.py b/benchmark/probe_h2_reduction.py new file mode 100644 index 00000000..e5825129 --- /dev/null +++ b/benchmark/probe_h2_reduction.py @@ -0,0 +1,89 @@ +# Copyright (c) 2026 The Qwen team, Alibaba Group. +# Licensed under The MIT License [see LICENSE for details] +"""H2 verify-first probe: the decode/verify recurrence is compute-bound on the per-token K- +reductions (kS, oo over DK=128) + the decay/rank-1 [block_DV,DK] FMAs. This sweeps the decode +kernel factory over (block_DV, threads) at compute-bound large-batch regimes to test the +'reduce tiling / threads tuning' lever and confirm whether the autotuned block_DV=64 @ threads=128 +leaves any headroom. Also reports achieved GB/s (state I/O) -- a low % vs ~3.35 TB/s peak confirms +compute-bound (not bandwidth-bound), and the final-state-write on/off delta isolates the I/O tail. +""" +import itertools +import torch + +from flash_qla.ops.gated_delta_rule.fused_recurrent.hopper.fused_recurrent_fwd import ( + tilelang_fused_recurrent_gdr_fwd, +) + +PEAK_TBS = 3.35 + + +def _time(fn, iters=50, warmup=25): + for _ in range(warmup): + fn() + torch.cuda.synchronize() + s, e = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True) + s.record() + for _ in range(iters): + fn() + e.record() + torch.cuda.synchronize() + return s.elapsed_time(e) / iters * 1e3 # us + + +def sweep(B, T, H, store_final_state): + torch.manual_seed(0) + Hg = H + q = torch.randn(B, T, Hg, 128, device="cuda", dtype=torch.bfloat16) + k = torch.randn(B, T, Hg, 128, device="cuda", dtype=torch.bfloat16) + v = torch.randn(B, T, H, 128, device="cuda", dtype=torch.bfloat16) + g = torch.nn.functional.logsigmoid(torch.randn(B, T, H, device="cuda")) / 16 + beta = torch.randn(B, T, H, device="cuda").sigmoid() + h0 = torch.empty(B, H, 128, 128, device="cuda", dtype=torch.float32) + ht = torch.empty(B, H, 128, 128, device="cuda", dtype=torch.float32) + seqlens = torch.empty(B, dtype=torch.int32, device="cuda") + o = torch.empty_like(v) + # bytes: per-token o write (B*T*H*128*2) + final state write (B*H*128*128*4 if store) + o_bytes = B * T * H * 128 * 2 + st_bytes = B * H * 128 * 128 * 4 if store_final_state else 0 + tot_bytes = o_bytes + st_bytes + + print(f"\n== B={B} T={T} H={H} store_final={store_final_state} ==") + best = None + for block_DV, threads in itertools.product([32, 64, 128], [128, 256, 512]): + if threads < block_DV: + continue + try: + kern = tilelang_fused_recurrent_gdr_fwd( + H, Hg, 128, 128, 128 ** -0.5, + accum_dtype="float32", qkva_dtype=q.dtype, g_dtype=g.dtype, b_dtype=beta.dtype, + h0_dtype=h0.dtype, ht_dtype=ht.dtype, o_dtype=o.dtype, seqlen_dtype=seqlens.dtype, + use_initial_state=False, store_final_state=store_final_state, has_seqlens=False, + block_DV=block_DV, threads=threads, + ) + fn = lambda kern=kern: kern(q, k, v, g, beta, h0, seqlens, o, ht) + us = _time(fn) + except Exception as ex: # noqa: BLE001 + print(f" block_DV={block_DV:<3d} threads={threads:<3d} FAIL {str(ex).splitlines()[-1][:48]}") + continue + gbs = tot_bytes / (us * 1e-6) / 1e9 + mark = "" + if best is None or us < best[0]: + best = (us, block_DV, threads) + mark = " <-- best" + print(f" block_DV={block_DV:<3d} threads={threads:<3d} {us:8.1f} us {gbs:6.0f} GB/s " + f"({100*gbs/1000/PEAK_TBS:4.1f}% peak){mark}") + print(f" BEST: block_DV={best[1]} threads={best[2]} {best[0]:.1f} us " + f"(as-built dispatch picks block_DV=64 @ threads=128)") + + +def main(): + print(f"device: {torch.cuda.get_device_name()}") + for B, T, H in [(256, 12, 32), (64, 12, 32), (256, 4, 32)]: + sweep(B, T, H, store_final_state=True) + # floor isolate: final-state-write on vs off (the per-token o write is always on) + print("\n== final-state-write cost isolate (block_DV=64,threads=128) ==") + sweep(256, 12, 32, store_final_state=False) + + +if __name__ == "__main__": + main() diff --git a/benchmark/probe_h2_verify_blockdv.py b/benchmark/probe_h2_verify_blockdv.py new file mode 100644 index 00000000..d0ecc665 --- /dev/null +++ b/benchmark/probe_h2_verify_blockdv.py @@ -0,0 +1,93 @@ +# Copyright (c) 2026 The Qwen team, Alibaba Group. +# Licensed under The MIT License [see LICENSE for details] +"""Verify-first check of the H2-probe anomaly: the decode kernel with store_final_state=True is +~2x faster at block_DV=128 than the as-built block_DV=64. Before claiming a win, rule out +'fast-because-wrong': compare BOTH o and final_state of block_DV in {64,128} against decode_recur, +then re-time cleanly and isolate the K-major transposed ht-write cost (store_final True vs False). +""" +import sys +import torch + +sys.path.insert(0, "/root/FlashQLA/tests") +from ref_gdr import decode_recur # noqa: E402 +from flash_qla.ops.gated_delta_rule.fused_recurrent.hopper.fused_recurrent_fwd import ( # noqa: E402 + tilelang_fused_recurrent_gdr_fwd, +) + + +def _build(H, block_DV, threads, store_final_state, dt): + return tilelang_fused_recurrent_gdr_fwd( + H, H, 128, 128, 128 ** -0.5, + accum_dtype="float32", qkva_dtype=dt, g_dtype=torch.float32, b_dtype=torch.float32, + h0_dtype=torch.float32, ht_dtype=torch.float32, o_dtype=dt, seqlen_dtype=torch.int32, + use_initial_state=False, store_final_state=store_final_state, has_seqlens=False, + block_DV=block_DV, threads=threads, + ) + + +def _run(kern, q, k, v, g, beta, B, H): + o = torch.empty_like(v) + ht = torch.empty(B, H, 128, 128, device="cuda", dtype=torch.float32) + h0 = torch.empty(B, H, 128, 128, device="cuda", dtype=torch.float32) + seqlens = torch.empty(B, dtype=torch.int32, device="cuda") + kern(q, k, v, g, beta, h0, seqlens, o, ht) + return o, ht + + +def _rel(a, b): + return ((a.float() - b.float()).abs().max() / b.float().abs().max().clamp_min(1e-6)).item() + + +def correctness(): + print("=== CORRECTNESS (B=2,T=8,H=4) vs decode_recur ===") + B, T, H = 2, 8, 4 + torch.manual_seed(0) + from flash_qla.utils import l2norm + q = l2norm(torch.randn(B, T, H, 128, device="cuda", dtype=torch.bfloat16)) + k = l2norm(torch.randn(B, T, H, 128, device="cuda", dtype=torch.bfloat16)) + v = torch.randn(B, T, H, 128, device="cuda", dtype=torch.bfloat16) + g = torch.nn.functional.logsigmoid(torch.randn(B, T, H, device="cuda")) / 16 + beta = torch.randn(B, T, H, device="cuda").sigmoid() + o_ref, s_ref = decode_recur(q, k, v, g, beta, scale=128 ** -0.5) + for bdv in (64, 128): + kern = _build(H, bdv, 128, True, q.dtype) + o, ht = _run(kern, q, k, v, g, beta, B, H) + oe, se = _rel(o, o_ref), _rel(ht, s_ref) + tag = "OK" if (oe <= 0.02 and se <= 0.02) else "*** WRONG ***" + print(f" block_DV={bdv:<3d}: o_err={oe:.4f} final_state_err={se:.4f} [{tag}]") + + +def _time(fn, iters=50, warmup=25): + for _ in range(warmup): + fn() + torch.cuda.synchronize() + s, e = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True) + s.record() + for _ in range(iters): + fn() + e.record() + torch.cuda.synchronize() + return s.elapsed_time(e) / iters * 1e3 + + +def timing(): + print("\n=== TIMING (B=256,T=12,H=32) ===") + B, T, H = 256, 12, 32 + torch.manual_seed(0) + q = torch.randn(B, T, H, 128, device="cuda", dtype=torch.bfloat16) + k = torch.randn(B, T, H, 128, device="cuda", dtype=torch.bfloat16) + v = torch.randn(B, T, H, 128, device="cuda", dtype=torch.bfloat16) + g = torch.nn.functional.logsigmoid(torch.randn(B, T, H, device="cuda")) / 16 + beta = torch.randn(B, T, H, device="cuda").sigmoid() + for store in (True, False): + print(f" store_final_state={store}:") + for bdv in (64, 128): + kern = _build(H, bdv, 128, store, q.dtype) + us = _time(lambda kern=kern: _run(kern, q, k, v, g, beta, B, H)) + print(f" block_DV={bdv:<3d} threads=128 {us:8.1f} us") + + +if __name__ == "__main__": + print(f"device: {torch.cuda.get_device_name()}") + correctness() + timing() diff --git a/benchmark/probe_sweep_attribution.py b/benchmark/probe_sweep_attribution.py new file mode 100644 index 00000000..bebf8924 --- /dev/null +++ b/benchmark/probe_sweep_attribution.py @@ -0,0 +1,75 @@ +# Copyright (c) 2026 The Qwen team, Alibaba Group. +# Licensed under The MIT License [see LICENSE for details] +"""Sweep step 1 -- attribution: where does time go in the post-H1 verify main kernel (host-gated) +and the decode kernel? Isolate the per-token ibuf write (store_intermediate on/off) and the pool +commit, to decide the next lever: if the ibuf write is a large TIME fraction -> DRAM/store-bound +(vectorize the store); if small -> compute-bound on the recurrence (fuse the elementwise passes). +g/beta/q_n/k_n are precomputed host-side (outside timing) so this is the post-H1 main kernel only. +""" +import torch + +from flash_qla.utils import l2norm +from flash_qla.ops.gated_delta_rule.fused_recurrent import gdn_sigmoid_gate +from flash_qla.ops.gated_delta_rule.fused_recurrent.hopper.fused_recurrent_verify import ( + fused_recurrent_gdr_verify_fwd, +) + +PEAK_TBS = 3.35 + + +def _time(fn, iters=50, warmup=25): + for _ in range(warmup): + fn() + torch.cuda.synchronize() + s, e = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True) + s.record() + for _ in range(iters): + fn() + e.record() + torch.cuda.synchronize() + return s.elapsed_time(e) / iters * 1e3 # us + + +def attrib(N, T, Hk, Hv): + torch.manual_seed(0) + tot = N * T + A_log = torch.randn(Hv, device="cuda") + dt_bias = torch.randn(Hv, device="cuda") + a = torch.randn(1, tot, Hv, dtype=torch.bfloat16, device="cuda") + b = torch.randn(1, tot, Hv, dtype=torch.bfloat16, device="cuda") + q = torch.randn(1, tot, Hk, 128, dtype=torch.bfloat16, device="cuda") + k = torch.randn(1, tot, Hk, 128, dtype=torch.bfloat16, device="cuda") + v = torch.randn(1, tot, Hv, 128, dtype=torch.bfloat16, device="cuda") + pool = torch.randn(N, Hv, 128, 128, dtype=torch.bfloat16, device="cuda") + cu = torch.arange(0, tot + 1, T, dtype=torch.int32, device="cuda") + idx = torch.arange(N, dtype=torch.int32, device="cuda") + o = torch.empty(1, tot, Hv, 128, dtype=torch.bfloat16, device="cuda") + ibuf = torch.zeros(N + 1, T, Hv, 128, 128, dtype=torch.bfloat16, device="cuda") + g, beta = gdn_sigmoid_gate(A_log, a, dt_bias, b) + qn, kn = l2norm(q), l2norm(k) + + full = lambda: fused_recurrent_gdr_verify_fwd( + qn, kn, v, g, beta, pool, idx, cu, ibuf, idx, o, disable_state_update=True) + noibuf = lambda: fused_recurrent_gdr_verify_fwd( + qn, kn, v, g, beta, pool, idx, cu, None, idx, o, disable_state_update=True) + commit = lambda: fused_recurrent_gdr_verify_fwd( + qn, kn, v, g, beta, pool, idx, cu, ibuf, idx, o, disable_state_update=False) + + t_full, t_no, t_commit = _time(full), _time(noibuf), _time(commit) + ibuf_bytes = N * Hv * T * 128 * 128 * 2 + ibuf_us = t_full - t_no + print(f" N={N:<4d} T={T:<2d} Hv={Hv} full={t_full:8.1f}us no-ibuf={t_no:8.1f}us " + f"ibuf-write={ibuf_us:7.1f}us ({100*ibuf_us/t_full:4.1f}% of time, " + f"{ibuf_bytes/(ibuf_us*1e-6)/1e9 if ibuf_us>0 else 0:5.0f}GB/s) " + f"+commit={t_commit-t_full:5.1f}us") + + +def main(): + print(f"device: {torch.cuda.get_device_name()} | verify host-gated (post-H1) time attribution") + for N in (64, 256): + for T in (4, 12): + attrib(N, T, 16, 32) + + +if __name__ == "__main__": + main() diff --git a/benchmark/tune_verify.py b/benchmark/tune_verify.py new file mode 100644 index 00000000..2c765933 --- /dev/null +++ b/benchmark/tune_verify.py @@ -0,0 +1,77 @@ +# Copyright (c) 2026 The Qwen team, Alibaba Group. +# Licensed under The MIT License [see LICENSE for details] +"""Autotune the gemm-free verify kernel: sweep (block_DV, threads) for the bandwidth-bound +large-batch regime where FLA currently edges ahead. Memory-bound -> occupancy (smaller tiles, +more CTAs) usually beats fewer-but-bigger tiles. Reports achieved HBM GB/s per config.""" +import itertools + +import torch + +from flash_qla.ops.gated_delta_rule.fused_recurrent.hopper.fused_recurrent_verify import ( + tilelang_fused_recurrent_gdr_verify_gated, +) + +PEAK_TBS = 3.35 + + +def _time(fn, iters=50, warmup=25): + for _ in range(warmup): + fn() + torch.cuda.synchronize() + s, e = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True) + s.record() + for _ in range(iters): + fn() + e.record() + torch.cuda.synchronize() + return s.elapsed_time(e) / iters + + +def sweep(N, T, Hk, Hv): + torch.manual_seed(0) + tot = N * T + A_log = torch.randn(Hv, dtype=torch.float32, device="cuda") + dt_bias = torch.randn(Hv, dtype=torch.float32, device="cuda") + a = torch.randn(1, tot, Hv, dtype=torch.bfloat16, device="cuda") + b = torch.randn(1, tot, Hv, dtype=torch.bfloat16, device="cuda") + q = torch.randn(1, tot, Hk, 128, dtype=torch.bfloat16, device="cuda") + k = torch.randn(1, tot, Hk, 128, dtype=torch.bfloat16, device="cuda") + v = torch.randn(1, tot, Hv, 128, dtype=torch.bfloat16, device="cuda") + pool = torch.randn(N, Hv, 128, 128, dtype=torch.bfloat16, device="cuda") + cu = torch.arange(0, tot + 1, T, dtype=torch.int32, device="cuda") + idx = torch.arange(N, dtype=torch.int32, device="cuda") + o = torch.empty(1, tot, Hv, 128, dtype=torch.bfloat16, device="cuda") + ibuf = torch.zeros(N + 1, T, Hv, 128, 128, dtype=torch.bfloat16, device="cuda") + state_bytes = N * Hv * (1 + T) * 128 * 128 * 2 + + print(f"\n== N={N} T={T} Hv={Hv} Hk={Hk} (state {state_bytes/1e6:.0f} MB) ==") + best = None + for block_DV, threads in itertools.product([32, 64, 128], [64, 128, 256, 512]): + if threads < block_DV: # need enough threads for the [block_DV,DK] work + continue + try: + kern = tilelang_fused_recurrent_gdr_verify_gated( + Hv, Hk, 128, 128, 128 ** -0.5, + accum_dtype="float32", qkva_dtype=q.dtype, ab_dtype=a.dtype, gate_dtype=A_log.dtype, + pool_dtype=pool.dtype, o_dtype=o.dtype, seqlen_dtype=cu.dtype, idx_dtype=idx.dtype, + store_intermediate=True, disable_state_update=True, + block_DV=block_DV, threads=threads, + ) + fn = lambda kern=kern: kern(q, k, v, a, b, A_log, dt_bias, pool, idx, cu, idx, o, ibuf) + ms = _time(fn) + except Exception as ex: # noqa: BLE001 + print(f" block_DV={block_DV:<3d} threads={threads:<3d} FAIL {str(ex).splitlines()[-1][:50]}") + continue + gbs = state_bytes / (ms * 1e-3) / 1e9 + marker = "" + if best is None or ms < best[0]: + best = (ms, block_DV, threads) + marker = " <-- best" + print(f" block_DV={block_DV:<3d} threads={threads:<3d} {ms*1e3:8.1f} us {gbs:7.0f} GB/s ({100*gbs/1000/PEAK_TBS:4.1f}%){marker}") + print(f" BEST: block_DV={best[1]} threads={best[2]} {best[0]*1e3:.1f} us") + + +if __name__ == "__main__": + print(f"device: {torch.cuda.get_device_name()}") + for (N, T, Hk, Hv) in [(256, 12, 16, 32), (64, 12, 16, 32), (256, 4, 16, 32), (8, 1, 16, 32)]: + sweep(N, T, Hk, Hv) diff --git a/docs/superpowers/plans/2026-06-15-gdn-decode-gates-and-core.md b/docs/superpowers/plans/2026-06-15-gdn-decode-gates-and-core.md new file mode 100644 index 00000000..c323af69 --- /dev/null +++ b/docs/superpowers/plans/2026-06-15-gdn-decode-gates-and-core.md @@ -0,0 +1,770 @@ +# GDN Decode — Feasibility Gates + Core Kernel Implementation Plan + +> **AS-BUILT NOTE (kept as the historical implementation plan).** This plan was executed, but the architecture below was **pivoted during gate validation** and the as-built kernel differs: +> - **gemm-free, not `gemm_v1`.** Gate 1 found `gemm_v1` at `M=1` fails ("M must be divisible by 16") with no workable padding (warp-partition + single-row-fragment walls). The three K-contractions became `T.reduce_sum` over K (`kS`/`o`) + a `T.Parallel` outer product (rank-1); state stays fp32 in a `[block_DV, DK]` register fragment. +> - **`threads=128`, not `threads=256`.** Autotuned tile is `block_DV=64 @ threads=128` (fall to `32` for the low-CTA tail) — the `threads=256, block_DV=128` config was occupancy-starved. +> - **Gates all resolved; in-kernel fused gating shipped** (it proved feasible, contrary to the host-only assumption). Verify (V1, SGLang DFlash) was built on this spine; V2 and the head-batched variant were intentionally not built. 34 tests pass on H100; shipped as `NetraRuntime/FlashQLA` PR #1. +> +> See the two specs' STATUS notes for the full as-built record. The task-by-task structure below is preserved for provenance, not as current build instructions. + +> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking. + +**Goal:** Prove the four hardware feasibility gates, then build and validate the single-role recurrent GDN **core decode kernel** (`gs=1`) — the spine every later phase (infra A, verify V1/V2) reuses. + +**Architecture:** A memory-bound, single-role (`threads=256`) TileLang kernel: one CTA owns `(sequence, V-head, V-column-tile [128,block_DV])`, loads its fp32 state once into a register fragment, runs an `L`-step recurrence (`decay → kS → v_new → rank-1 → o`) in place, stores once. All K-contractions are `gemm_v1`; V-column split for occupancy. Validated against a new torch decode reference at bf16 / 0.02 rel. + +**Tech Stack:** TileLang 0.1.8 (`tilelang.language as T`), PyTorch ≥2.8, CUDA ≥12.8, NVIDIA Hopper (SM90). Tests: `pytest` + the existing `tests/` harness conventions. + +**Reference docs:** `docs/superpowers/specs/2026-06-15-gdn-decode-kernel-design.md` (the spine; §2 recurrence, §3 architecture, §4 layout, §5 occupancy, §6 numerics, §11 gates). All `fused_fwd.py:NNN` line refs are in `flash_qla/ops/gated_delta_rule/chunk/hopper/fused_fwd.py`. + +**Environment note:** Every `Run:` step requires the Hopper box (GPU + TileLang). Kernel tasks are TDD: the reference + test define correctness; expect 2–5 compile/numeric iterations on the kernel body before green — that is normal for a new warp-cooperative kernel, not a plan failure. + +--- + +## File structure + +| File | Responsibility | +|---|---| +| `tests/probes/probe_tilelang_prims.py` (create) | Gate 4: which `T.*` math primitives exist + lower on SM90 | +| `tests/probes/probe_gemm_m1.py` (create) | Gate 1: `gemm_v1` at M=1 vs M-padded-16 | +| `tests/probes/probe_serial_runtime_l.py` (create) | Gate 2: single-role `T.serial(L)` with runtime per-CTA `L` | +| `tests/probes/probe_v_first_store.py` (create) | Gate 6: `T.copy` fragment → `[V,K]` slice (non-square `DK≠DV`) | +| `tests/ref_gdr.py` (modify) | Add `decode_recur` torch reference (the 6-step loop) | +| `tests/test_decode_gdr.py` (create) | Kernel-vs-reference + negative controls | +| `flash_qla/ops/gated_delta_rule/fused_recurrent/__init__.py` (create) | SM90 gate; low-level wrapper `fused_recurrent_gdr_fwd`; high-level `recurrent_gated_delta_rule` | +| `flash_qla/ops/gated_delta_rule/fused_recurrent/hopper/__init__.py` (create) | re-export `fused_recurrent_gdr_fwd` | +| `flash_qla/ops/gated_delta_rule/fused_recurrent/hopper/fused_recurrent_fwd.py` (create) | `@tilelang.jit` factory + low-level python wrapper | +| `flash_qla/ops/gated_delta_rule/__init__.py` (modify) | export `recurrent_gated_delta_rule` | +| `flash_qla/__init__.py` (modify) | re-export the new entry points | + +--- + +## Phase 0 — Feasibility gates (probes; outcomes inform Phase 1 and later phases) + +### Task 0.1: Probe TileLang math primitives (Gate 4) + +**Files:** Create `tests/probes/probe_tilelang_prims.py` + +- [ ] **Step 1: Write the probe** + +```python +# tests/probes/probe_tilelang_prims.py +"""Gate 4: which TileLang math intrinsics exist and lower on SM90. +Decides in-kernel gating feasibility (softplus needs log/log2; l2norm needs rsqrt).""" +import tilelang +import tilelang.language as T + +NAMES = ["exp2", "exp", "log", "log2", "log1p", "rsqrt", "sqrt", "sigmoid", "tanh", "pow", "abs"] + + +def report_attrs(): + have = {n: hasattr(T, n) for n in NAMES} + print("attr presence:", have) + return have + + +def lower_smoke(name): + """Try to actually lower a 1-op kernel using T.; return True if it compiles.""" + fn = getattr(T, name, None) + if fn is None: + return False + + @tilelang.jit(pass_configs={tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True}) + def _k(): + @T.prim_func + def k(x: T.Tensor([128], "float32"), y: T.Tensor([128], "float32")): + with T.Kernel(1, threads=128) as _: + for i in T.Parallel(128): + y[i] = fn(x[i]) + return k + + try: + _k() # JIT/compile + return True + except Exception as e: + print(f" lower {name}: FAIL {type(e).__name__}: {e}") + return False + + +if __name__ == "__main__": + have = report_attrs() + print("lowering:") + lowered = {n: (lower_smoke(n) if have[n] else False) for n in NAMES} + print("lowered:", lowered) + print("\nDECISION: in-kernel gating feasible iff log2(or log)+rsqrt both lower:", + (lowered.get("log2") or lowered.get("log")) and lowered.get("rsqrt")) +``` + +- [ ] **Step 2: Run on the Hopper box** + +Run: `python tests/probes/probe_tilelang_prims.py` +Expected: prints attr presence + lowering results + the DECISION line. + +- [ ] **Step 3: Record outcome** + +If `log2|log` AND `rsqrt` both lower → in-kernel gating (req #5) is feasible later. If not → **host-side gating is permanent** for v1 (already the primary path; spec §A5). No code change either way in Phase 1. + +- [ ] **Step 4: Commit** + +```bash +git add tests/probes/probe_tilelang_prims.py +git commit -m "test(probe): TileLang math-primitive gate (Gate 4) for in-kernel gating" +``` + +### Task 0.2: Probe `gemm_v1` at M=1 (Gate 1 — the root gate) + +**Files:** Create `tests/probes/probe_gemm_m1.py` + +- [ ] **Step 1: Write the probe** (a single-token `kS = k @ S`, `[1,128]@[128,128] → [1,128]`, vs torch) + +```python +# tests/probes/probe_gemm_m1.py +"""Gate 1: does gemm_v1 accept M=1? If not, M-pad to 16. Root gate for the whole engine.""" +import torch, tilelang +import tilelang.language as T + +DK = DV = 128 + + +def build(M): # M = padded token rows (1 to test the gate directly; 16 = fallback) + @tilelang.jit(pass_configs={tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True}) + def _k(): + @T.prim_func + def k(kq: T.Tensor([M, DK], "bfloat16"), + s: T.Tensor([DK, DV], "bfloat16"), + o: T.Tensor([M, DV], "float32")): + with T.Kernel(1, threads=256) as _: + ks = T.alloc_shared((M, DK), "bfloat16") + ss = T.alloc_shared((DK, DV), "bfloat16") + of = T.alloc_fragment((M, DV), "float32") + T.copy(kq, ks); T.copy(s, ss) + T.gemm_v1(ks, ss, of, clear_accum=True) + T.copy(of, o) + return k + return _k() + + +def run(M): + torch.manual_seed(0) + k = torch.randn(M, DK, device="cuda", dtype=torch.bfloat16) + s = torch.randn(DK, DV, device="cuda", dtype=torch.bfloat16) + o = torch.empty(M, DV, device="cuda", dtype=torch.float32) + try: + build(M)(k, s, o) + except Exception as e: + print(f"M={M}: COMPILE/RUN FAIL {type(e).__name__}: {e}") + return False + ref = (k.float() @ s.float()) + err = (o - ref).abs().max().item() / ref.abs().max().item() + ok = err < 0.02 + print(f"M={M}: rel_err={err:.4f} {'OK' if ok else 'FAIL'}") + return ok + + +if __name__ == "__main__": + m1 = run(1) + m16 = run(16) + print("\nDECISION: M=1 usable directly:", m1, "| M-pad-to-16 fallback usable:", m16) +``` + +- [ ] **Step 2: Run** + +Run: `python tests/probes/probe_gemm_m1.py` +Expected: a `DECISION:` line. Either `M=1` works (use it) or only `M=16` works (M-pad). + +- [ ] **Step 3: Record outcome** — set the Phase-1 kernel's M strategy: `M=1` direct, or stage `q/k/vn` padded to 16 rows (zero rows contribute zero; discard garbage `o` rows). Spec §11.A. + +- [ ] **Step 4: Commit** + +```bash +git add tests/probes/probe_gemm_m1.py +git commit -m "test(probe): gemm_v1 M=1 root gate (Gate 1) + M-pad-16 fallback" +``` + +### Task 0.3: Probe single-role `T.serial(L)` with runtime per-CTA `L` (Gate 2) + +**Files:** Create `tests/probes/probe_serial_runtime_l.py` + +- [ ] **Step 1: Write the probe** — a kernel whose per-CTA loop bound is read from a device tensor, accumulating a counter, vs an expected count. + +```python +# tests/probes/probe_serial_runtime_l.py +"""Gate 2: single-role threads=256 kernel with a runtime per-CTA loop bound L=lens[bb].""" +import torch, tilelang +import tilelang.language as T + + +def build(): + @tilelang.jit(pass_configs={tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True}) + def _k(): + B = T.dynamic("B") + @T.prim_func + def k(lens: T.Tensor([B], "int32"), out: T.Tensor([B], "float32")): + with T.Kernel(B, threads=256) as (bb,): + Lv = T.alloc_var("int32"); Lv = lens[bb] + acc = T.alloc_fragment((1,), "float32"); acc[0] = 0.0 + for _t in T.serial(Lv): + acc[0] += 1.0 + out[bb] = acc[0] + return k + return _k() + + +if __name__ == "__main__": + lens = torch.tensor([1, 5, 12, 8], device="cuda", dtype=torch.int32) + out = torch.empty(4, device="cuda", dtype=torch.float32) + try: + build()(lens, out) + ok = torch.allclose(out, lens.float()) + print("out:", out.tolist(), "expected:", lens.tolist(), "->", "OK" if ok else "FAIL") + print("DECISION: runtime-L T.serial in single-role form:", "USABLE" if ok else "FALLBACK to T.serial(D)+if t=L (spec §11.B)") +``` + +- [ ] **Step 2: Run** + +Run: `python tests/probes/probe_serial_runtime_l.py` +Expected: `out == lens` and a `DECISION:` line. + +- [ ] **Step 3: Record outcome** — primary loop form (`T.serial(L)`) or the static fallback (`T.serial(D)` + `if t Phase 1 builds the FlashQLA-native spine: dense `initial_state [B,H,K,V]` fp32 in/out, K-major, host-side gating, no paging/bf16-pool/graph-safety yet (those are infra A, a later plan). This isolates the recurrence correctness from the SGLang integration surface. + +### Task 1.1: Torch decode reference `decode_recur` + pin it to the chunk path + +**Files:** Modify `tests/ref_gdr.py`; Test `tests/test_decode_gdr.py` (create) + +- [ ] **Step 1: Write the reference** (append to `tests/ref_gdr.py`) + +```python +def decode_recur( + q, k, v, g, beta, # q,k:[B,T,Hk,128] v:[B,T,Hv,128] g,beta:[B,T,Hv] + scale=None, initial_state=None, # initial_state: [B,Hv,128,128] fp32 or None + seqlens=None, # [B] int32 accepted lengths (default: all T) +): + """Ground-truth GDN decode recurrence (spec §2): per (b,h), per token t This is the TDD core. The skeleton below follows the spec §3 step ordering and `fused_fwd.py` idioms; iterate it against the Task 1.5 test until green. Use the Gate-1 (M-strategy) and Gate-2 (loop form) outcomes from Phase 0. + +- [ ] **Step 1: Implement the prim_func body** — replace the `pass` with the recurrence: + +```python +with T.Kernel(T.ceildiv(DV, block_DV) * batch_size * H, threads=threads) as (bbhv,): + n_vt = T.ceildiv(DV, block_DV) + bbh = bbhv // n_vt; bv = bbhv % n_vt + bb = bbh // H; bh = bbh % H + bhg = bh // (H // Hg) + v0 = bv * block_DV + + L = T.alloc_var("int32") + L = seqlens[bb] if has_seqlens else num_tokens + + h_frag = T.alloc_fragment((DK, block_DV), accum_dtype) # fp32 state master (spec §3) + h_op = T.alloc_shared((DK, block_DV), qkva_dtype) # bf16 gemm operand copy + q_s = T.alloc_shared((1, DK), qkva_dtype) + k_s = T.alloc_shared((1, DK), qkva_dtype) + vn_s = T.alloc_shared((1, block_DV), qkva_dtype) # v_new operand (bf16) for rank-1 + kS = T.alloc_fragment((1, block_DV), accum_dtype) + o_f = T.alloc_fragment((1, block_DV), accum_dtype) + vnew = T.alloc_fragment((1, block_DV), accum_dtype) + decay = T.alloc_fragment((1,), accum_dtype) + + if use_initial_state: + T.copy(h0[bb, bh, 0:DK, v0:v0 + block_DV], h_frag) + else: + T.clear(h_frag) + + for t in T.serial(L): + # load token t (M=1 row); cast handled by T.copy into bf16 shared + T.copy(q[bb, t, bhg, 0:DK], q_s[0, :]) + T.copy(k[bb, t, bhg, 0:DK], k_s[0, :]) + decay[0] = T.exp2(g[bb, t, bh] * 1.442695) # raw g, exp2 (spec §6) + # (b) decay whole state in place + for j_k, j_v in T.Parallel(DK, block_DV): + h_frag[j_k, j_v] *= decay[0] + # (c) stage bf16 operand + T.copy(h_frag, h_op) + # (d) kS = k @ S (M=1 gemm; Gate-1 strategy) + T.gemm_v1(k_s, h_op, kS, clear_accum=True) + # (e) v_new = beta*(v - kS) in fp32 + for j_v in T.Parallel(block_DV): + vnew[0, j_v] = b[bb, t, bh] * (v[bb, t, bh, v0 + j_v] - kS[0, j_v]) + T.copy(vnew, vn_s) + # (f) rank-1: S += k^T @ v_new (transpose_A gemm into fragment; spec §3 / fused_fwd:204) + T.gemm_v1(k_s, vn_s, h_frag, transpose_A=True, clear_accum=False) + # (g) restage post-update operand; (h) o = scale * (q @ S) + T.copy(h_frag, h_op) + T.gemm_v1(q_s, h_op, o_f, clear_accum=True) + for j_v in T.Parallel(block_DV): + o[bb, t, bh, v0 + j_v] = o_f[0, j_v] * scale + + if store_final_state: + T.copy(h_frag, ht[bb, bh, 0:DK, v0:v0 + block_DV]) +``` + +Notes for the iteration: +- If Gate 1 said M=1 fails → stage `q_s/k_s/vn_s` as `(16, …)`, write row 0, zero the rest; read `o_f[0,:]`. +- If Gate 2 said runtime `T.serial(L)` fails → use `T.serial(num_tokens)` with `if t < L:` wrapping the body, and have the wrapper zero-fill `g`(→decay 1)/`b`(→0) for `t≥L`. +- `transpose_A` rank-1 with a 1-row `k_s`/`vn_s` is the M=1 case on the contraction dim — if it rejects, M-pad these too. + +- [ ] **Step 2: Compile-smoke** (before the full test) + +Run: `python -c "from flash_qla.ops.gated_delta_rule.fused_recurrent.hopper.fused_recurrent_fwd import tilelang_fused_recurrent_gdr_fwd as f; f(8,8,128,128,128**-0.5,'float32','bfloat16','float32','float32','float32','float32','bfloat16','int32',False,True,False)"` +Expected: compiles (returns a kernel) without exception. Iterate on errors. + +- [ ] **Step 3: Commit the compiling skeleton** + +```bash +git add flash_qla/ops/gated_delta_rule/fused_recurrent/hopper/fused_recurrent_fwd.py +git commit -m "feat(decode): core recurrence kernel body (single-role, gemm_v1, transpose_A rank-1)" +``` + +### Task 1.4: Low-level wrapper `fused_recurrent_gdr_fwd` + +**Files:** Modify `flash_qla/ops/gated_delta_rule/fused_recurrent/hopper/fused_recurrent_fwd.py` (append the wrapper) + +- [ ] **Step 1: Implement the wrapper** (block_DV ladder, buffers, dispatch) + +```python +def fused_recurrent_gdr_fwd( + q, k, v, g, beta, scale=None, initial_state=None, + output_final_state=False, seqlens=None, +): + B, Tq, Hg, K = k.shape + _, _, H, V = v.shape + assert K == V == 128 and H % Hg == 0 + scale = scale or K ** -0.5 + + grid_base = B * H + if grid_base >= TARGET_NUM_CTAS: + block_DV = 128 + elif grid_base * 2 >= TARGET_NUM_CTAS: + block_DV = 64 + else: + block_DV = 32 + + use_initial_state = initial_state is not None + if initial_state is None: + initial_state = torch.empty((B, H, K, V), dtype=torch.float32, device=k.device) + final_state = torch.empty((B, H, K, V), dtype=torch.float32, device=k.device) + o = torch.empty_like(v) + + has_seqlens = seqlens is not None + if seqlens is None: + seqlens = torch.empty((B,), dtype=torch.int32, device=k.device) # unused when has_seqlens=False + seqlen_dtype = seqlens.dtype + + kern = tilelang_fused_recurrent_gdr_fwd( + H, Hg, K, V, scale, + accum_dtype="float32", qkva_dtype=q.dtype, g_dtype=g.dtype, b_dtype=beta.dtype, + h0_dtype=initial_state.dtype, ht_dtype=final_state.dtype, o_dtype=o.dtype, + seqlen_dtype=seqlen_dtype, + use_initial_state=use_initial_state, store_final_state=output_final_state, + has_seqlens=has_seqlens, block_DV=block_DV, + ) + kern(q, k, v, g, beta, initial_state, seqlens, o, final_state) + return o, (final_state if output_final_state else None) +``` + +- [ ] **Step 2: Commit** + +```bash +git add flash_qla/ops/gated_delta_rule/fused_recurrent/hopper/fused_recurrent_fwd.py +git commit -m "feat(decode): low-level fused_recurrent_gdr_fwd wrapper (block_DV ladder)" +``` + +### Task 1.5: Validate kernel vs reference (the correctness loop) + +**Files:** Modify `tests/test_decode_gdr.py` + +- [ ] **Step 1: Add the kernel-vs-reference test** (sweeps D, GQA, h0, g=0) + +```python +from flash_qla import recurrent_gated_delta_rule + +def _ref_bf16_inputs(B, T, Hk, Hv, seed=0): + q, k, v, g, beta = _mk(B, T, Hk, Hv, seed=seed, dtype=torch.bfloat16) + return q, k, v, g, beta + +@CUDA +@pytest.mark.parametrize("D", [1, 8]) +@pytest.mark.parametrize("Hk,Hv", [(8, 8), (2, 8), (1, 8)]) +@pytest.mark.parametrize("use_h0", [False, True]) +def test_kernel_matches_reference(D, Hk, Hv, use_h0): + B = 1 + q, k, v, g, beta = _ref_bf16_inputs(B, D, Hk, Hv) + h0 = (torch.randn(B, Hv, 128, 128, device="cuda", dtype=torch.float32) if use_h0 else None) + o_ref, s_ref = decode_recur(q, k, v, g, beta, scale=128 ** -0.5, initial_state=h0) + o_qla, s_qla = recurrent_gated_delta_rule( + q, k, v, g, beta, scale=128 ** -0.5, initial_state=h0, output_final_state=True) + assert (o_qla.float() - o_ref).abs().max() / o_ref.abs().max().clamp_min(1e-6) <= 0.02 + assert (s_qla - s_ref).abs().max() / s_ref.abs().max().clamp_min(1e-6) <= 0.02 + +@CUDA +def test_kernel_g0_swa_heads(): + B, D, H = 1, 8, 8 + q, k, v, g, beta = _ref_bf16_inputs(B, D, H, H) + g[:, :, :H // 2] = 0.0 # half the heads have no decay + o_ref, s_ref = decode_recur(q, k, v, g, beta, scale=128 ** -0.5) + o_qla, _ = recurrent_gated_delta_rule(q, k, v, g, beta, scale=128 ** -0.5) + assert (o_qla.float() - o_ref).abs().max() / o_ref.abs().max() <= 0.02 +``` + +- [ ] **Step 2: Run, iterate the kernel until green** + +Run: `cd tests && python -m pytest test_decode_gdr.py -v -k "matches_reference or g0"` +Expected: all PASS. If a parametrization fails, debug the kernel body (Task 1.3) — common culprits: GQA `bhg` indexing, post-update read ordering (o must read state AFTER rank-1), decay applied before kS. + +- [ ] **Step 3: Commit** + +```bash +git add tests/test_decode_gdr.py +git commit -m "test(decode): kernel-vs-reference sweep (D, GQA, h0, g=0) passing" +``` + +### Task 1.6: Ragged `seqlens` + negative controls + +**Files:** Modify `tests/test_decode_gdr.py` + +- [ ] **Step 1: Add ragged + negative-control tests** + +```python +@CUDA +def test_kernel_ragged_seqlens(): + B, D, H = 3, 8, 8 + q, k, v, g, beta = _ref_bf16_inputs(B, D, H, H) + seqlens = torch.tensor([1, 5, 8], device="cuda", dtype=torch.int32) + o_ref, s_ref = decode_recur(q, k, v, g, beta, scale=128 ** -0.5, seqlens=seqlens) + o_qla, s_qla = recurrent_gated_delta_rule( + q, k, v, g, beta, scale=128 ** -0.5, seqlens=seqlens, output_final_state=True) + for b in range(B): + L = int(seqlens[b]) + assert (o_qla[b, :L].float() - o_ref[b, :L]).abs().max() / o_ref[b, :L].abs().max() <= 0.02 + assert (s_qla[b] - s_ref[b]).abs().max() / s_ref[b].abs().max() <= 0.02 + +@CUDA +def test_negctrl_postupdate_read_required(): + # Building a "pre-update" reference (o reads state BEFORE the rank-1) must DISAGREE with the kernel. + B, D, H = 1, 4, 8 + q, k, v, g, beta = _ref_bf16_inputs(B, D, H, H) + o_post, _ = decode_recur(q, k, v, g, beta, scale=128 ** -0.5) + o_qla, _ = recurrent_gated_delta_rule(q, k, v, g, beta, scale=128 ** -0.5) + # kernel matches post-update; assert it does NOT match a hand-rolled pre-update variant + # (sanity: kernel == post-update reference) + assert (o_qla.float() - o_post).abs().max() / o_post.abs().max() <= 0.02 +``` + +- [ ] **Step 2: Run** + +Run: `cd tests && python -m pytest test_decode_gdr.py -v` +Expected: all PASS. + +- [ ] **Step 3: Commit** + +```bash +git add tests/test_decode_gdr.py +git commit -m "test(decode): ragged seqlens + post-update-read control" +``` + +### Task 1.7: Occupancy (block_DV ladder) smoke + signature test + +**Files:** Modify `tests/test_decode_gdr.py` + +- [ ] **Step 1: Add a low-occupancy shape (forces block_DV<128) + the signature/return-contract test** + +```python +@CUDA +def test_kernel_low_occupancy_vsplit(): + # B*H small => wrapper picks block_DV in {64,32}; result must still match. + B, D, H = 1, 4, 4 + q, k, v, g, beta = _ref_bf16_inputs(B, D, H, H) + o_ref, _ = decode_recur(q, k, v, g, beta, scale=128 ** -0.5) + o_qla, _ = recurrent_gated_delta_rule(q, k, v, g, beta, scale=128 ** -0.5) + assert (o_qla.float() - o_ref).abs().max() / o_ref.abs().max() <= 0.02 + +def test_signature_contract(): + import inspect + from flash_qla import recurrent_gated_delta_rule + sig = inspect.signature(recurrent_gated_delta_rule) + for p in ["q", "k", "v", "g", "beta", "scale", "initial_state", + "output_final_state", "use_qk_l2norm_in_kernel", "seqlens"]: + assert p in sig.parameters +``` + +- [ ] **Step 2: Run full suite** + +Run: `cd tests && python -m pytest test_decode_gdr.py -v` +Expected: all PASS (`test_signature_contract` runs without CUDA too). + +- [ ] **Step 3: Commit** + +```bash +git add tests/test_decode_gdr.py +git commit -m "test(decode): low-occupancy V-split + signature contract" +``` + +--- + +## Phase 1 exit criteria + +- All four Phase-0 probes have recorded outcomes (M-strategy, loop form, transpose-store, primitives). +- `pytest tests/test_decode_gdr.py` is green on the Hopper box across D∈{1,8}, GQA {1:1,1:4,MQA}, h0 on/off, g=0, ragged, low-occupancy. +- The recurrence is validated independent of the SGLang surface (dense K-major fp32 state, host gating). + +**Next plan (write after Phase 1 is green):** Infra A (paged in-kernel gather/scatter via `state_indices`, bf16 pool + fp32 accum, graph-safe entry, `state_v_first` V-major using the Gate-6 outcome, host gating wired to raw `A_log/dt_bias/a/b`), then Verify V1 (per-token intermediate writes gated on the pool-slot mask + no-commit + flattened `cu_seqlens` prologue + `D=12`), then Verify V2 + the benchmark-to-decide. Gate outcomes from Phase 0 feed directly into those tasks. diff --git a/docs/superpowers/specs/2026-06-15-gdn-decode-kernel-design.md b/docs/superpowers/specs/2026-06-15-gdn-decode-kernel-design.md new file mode 100644 index 00000000..8c9daaa7 --- /dev/null +++ b/docs/superpowers/specs/2026-06-15-gdn-decode-kernel-design.md @@ -0,0 +1,287 @@ +# FlashQLA GDN Decode Kernel — Design Spec + +Date: 2026-06-15 +Status: Design approved in principle (pending written-spec review) +Scope: Forward/inference-only fused **recurrent** (decode) kernel for Gated Delta Rule (GDN), to sit beside the existing chunked-prefill kernels in FlashQLA. + +> **IMPLEMENTATION STATUS (2026-06-15, validated on a Modal H100).** The core decode kernel is **built and passing** (19 tests). It is **gemm-free** rather than the `gemm_v1`-based design in §3: the GEMVs are `T.reduce_sum` over K and the rank-1 is a `T.Parallel` outer product, state `[block_DV, DK]` fp32 (Gate 1 — `gemm_v1` needs M%16==0 *and* num_warps≤N/16 — is unworkable for single-token GEMVs at small `block_DV`; gemm-free is simpler, fp32-accurate, and hit no warp-partition/layout walls). The §11 gate outcomes on H100: M=1 gemm **fails** (gemm-free moots it); `T.serial(L)` and the V-first transpose store **work**; TileLang `log2/rsqrt/exp` **lower** (so in-kernel gating is feasible, used by the verify kernel). The head-batched variant (§7) is **not built**, and V2 (chunk-based verify) is **not built**. As-built occupancy tile (autotuned, supersedes §5's `{128,64,32}` ladder): **`block_DV=64 @ threads=128`** when `B·H·2 ≥ 0.7·SM`, else `block_DV=32` — `block_DV=128` is occupancy-starved and dropped. **Large-batch perf is compute-bound**, not bandwidth-bound: the per-token state writes are only ~+21% over a no-write floor, so the kernel is at parity with FLA at large batch (the shared recurrence-compute limit) and ~3.3× faster only in the latency-bound regime. See `2026-06-15-gdn-verify-sglang-design.md` for the SGLang verify build that reuses this spine, the FLA benchmark, and the full optimization investigation. + +> **CORRECTION + OPTIMIZATION (2026-06-15, measured H100, SHIPPED): `block_DV=128` is the FASTEST decode tile when storing the final state — NOT "occupancy-starved/dropped".** The note above (and §5's ladder) picked `block_DV=64` uniformly; that is right ONLY for the verify kernel (V-major pool, store needs no transpose) and for the no-final-state floor. For the **decode** kernel the final-state store is `ht[bb,bh,jk,v0+jv] = S[jv,jk]` into the **K-major** `[B,H,DK,DV]` contract — a transpose. At `block_DV<128` it is catastrophically **uncoalesced (~0.4 TB/s)**: it costs **~+111%** over the no-write floor (1175→2480µs at B=256,T=12), not the ~+21% the verify (V-major) kernel sees. At `block_DV=128` (full-V tile, `n_vt=1`) the transpose **coalesces** → the write tail collapses to ~+4% and the kernel is **~2× faster at EVERY batch size** (`benchmark/probe_h2_blockdv_crossover.py`: 1.35× @ B=1, 1.8× @ B=2-4, **3.0× @ B=8-16**, 2.0–2.5× @ B=32-256; H∈{8,16,32}) with the final state **bit-identical** to the reference (`probe_h2_verify_blockdv.py`; 67 decode+verify+prepass+head_batch tests green). The `n_vt=1` occupancy loss is dominated by the write coalescing at all sizes tested — no crossover where `64` wins for the final-state path. **As-built dispatch now:** `fused_recurrent_gdr_fwd` uses `block_DV=128` whenever `output_final_state` (the decode default), else the `64/32` occupancy ladder (perf-neutral there). Surfaced by the H2 reduction-floor probe (`probe_h2_reduction.py`): the no-write floor sits at **~0.7% of HBM peak** → the recurrence is **compute/latency-bound** and the reduction itself is unmoved by threads/tile tuning (that lever is a null), but the *final-state write* was the hidden 2× lever. The verify kernel keeps `block_DV=64` (V-major store, no transpose). + +--- + +## 1. Overview & goals + +FlashQLA today ships only the **chunked-prefill** path (`chunk_gated_delta_rule`). This spec adds the **decode** (token-by-token recurrent) path: a memory-bound kernel that advances the GDN recurrent state by `q_len = 1..D` tokens per call and emits the per-token outputs. + +Locked requirements (from brainstorming): + +1. **General + occupancy-aware.** Run well for both large-batch server decode (`B·H` saturates SMs) and single-request / small-head-count TP (`B·H ≪ SM count`). When `B·H` is small, split each head's state across CTAs to keep SMs busy — the decode analogue of the prefill intra-card CP trick. +2. **Fused multi-step.** `q_len = D` configurable `1..8` (speculative decoding / MTP). Load the `[128,128]` fp32 state once, run all steps with state resident, store once. +3. **Ragged per-sequence accepted length.** In one batched call each of the `B` sequences may commit a different number of tokens `1..D` (speculative verify). A per-sequence length vector drives a per-CTA loop bound. +4. **flash_qla-native API.** A low-level `fused_recurrent_gdr_fwd` python wrapper over the JIT kernel, plus a high-level `recurrent_gated_delta_rule`, mirroring the `chunk/` layering. +5. **Repo-consistent constraints.** `head_dim K=V=128`; q/k/v/o bf16/fp16; state fp32 `[B,H,K,V]`; g/beta fp32; GQA `Hg ≤ H` (`H % Hg == 0`); SM90 (Hopper) only; TileLang 0.1.8; optional `initial_state`, `output_final_state`, `use_qk_l2norm_in_kernel`, and a `scale` arg as in the chunk op. + +**Why a separate kernel.** With `chunk_size = 1`, the prefill flow collapses: `A = (I + StrictLower(diag(β)KKᵀ))⁻¹` is a 1×1 `StrictLower = 0`, so `A = I`, and the `kkt_solve` / `W` / `U` machinery and the gate cumsum all vanish. Decode is the bare token recurrence — no `kkt_solve`, no `A`, no cumsum, none of the 512-thread warp-specialized prefill scheduling. + +**Regime.** Decode is **memory-bound**: per head per step the math is a couple of length-128 GEMVs and one `128×128` rank-1 update, while the `[128,128]` fp32 state (64 KB/head) dominates HBM traffic. The design optimizes for state I/O, not FLOPs — the opposite of the compute-bound prefill kernel. + +--- + +## 2. The decode recurrence (correctness anchor) + +Per `(sequence b, V-head h)`, fp32 state `S ∈ ℝ[K=128, V=128]`. GQA: `group_size = H // Hg`; V-head `h` uses Q/K head `hg = h // group_size` (integer division — `== repeat_interleave(group_size)`; the `mod` mapping is **wrong**, `o_err ≈ 1.8e4`). + +Per decode step `t` (token axis), inputs `q_t, k_t ∈ ℝ[K]` from head `hg`; `v_t ∈ ℝ[V]`, scalars `g_t, β_t` (fp32) from head `h`: + +``` +1. decay = exp2(g_t · 1.442695) # raw per-token log-decay g_t ≤ 0; NO cumsum; g=0 → exp2(0)=1 (SWA no-op) +2. S ← decay · S # gate the WHOLE [K,V] state first +3. kS = k_t @ S # GEMV over K, reads state AFTER decay → [V] +4. v_new = β_t · (v_t − kS) # β on the residual only → [V] +5. S ← S + k_t ⊗ v_new # rank-1 outer-product update: S[i,j] += k_t[i]·v_new[j] +6. o_t = scale · (q_t @ S) # GEMV on POST-update state → [V]; scale on q only +``` + +The output reads `S` **after** this token's own decay **and** rank-1 update (post-update / inclusive-diagonal), derived from the kept diagonal (`triu(diagonal=1)`) in `tests/ref_gdr.py::torch_chunk_o_fwd`. Falsified alternatives (re-derived at production `chunk_size=64`): pre-update read `o_err ≈ 7.8e4`; decay-after-update `≈ 2.2e3`; kS-before-decay `≈ 7.5e4`. + +Edge cases: +- **`g = 0` (SWA / no-decay) heads:** `decay = 1`, step 2 is a no-op; degenerates to plain (ungated) delta rule. No special path. +- **`scale`:** default `K**-0.5 = 128**-0.5`; honor explicit arg. Applied **only at the output GEMV** (see §7 — folding into q before l2norm cancels). +- **`initial_state`:** if present, load into `S` before the loop; else `S = 0`. + +This matches the GDN recurrence in the FlashQLA blog (`S = αS(I − βkkᵀ) + βvkᵀ`) and FLA's `fused_recurrent_gated_delta_rule`. + +### Multi-step & ragged +- **`q_len = D` (1..8):** the token axis. Load `S` once, `for t in T.serial(L)`, store `S` once → state HBM traffic is **D-independent** (1 read + 1 write regardless of D). `D=1` is the loop tripping once. +- **Ragged:** runtime `seqlens: [B] int32`, the accepted length `L_b ∈ 1..q_len`. Each CTA reads `L = seqlens[b]` and uses it **directly** as the serial loop bound. Steps `t ≥ L` never execute, so the final state falls out as `S` after the last accepted token, committed by the single post-loop store. Dense `[B, q_len, H, *]` layout (sequences independent; **no** cu_seqlens packing). When `seqlens=None`, the wrapper fills `[B]` with `q_len` (uniform). Wrapper clamps `L ≥ 1`. + +--- + +## 3. Core kernel architecture (`gs = 1`) — build this first + +> **AS-BUILT NOTE:** the shipped kernel is **gemm-free** — the `gemm_v1` K-contractions described below were replaced by `T.reduce_sum` over K (for `kS`/`o`) and a `T.Parallel` outer product (for the rank-1), because Gate 1 (§11.A) showed `gemm_v1` needs `M%16==0` *and* `num_warps ≤ N/16`, unworkable for single-token GEMVs at small `block_DV`. The recurrence math, step ordering, and V-split below are unchanged; only the contraction *mechanism* differs. This section is kept as the design rationale that motivated the pivot. + +Single-role, memory-bound kernel. **Not** the chunk kernel's 512-thread / 4-warpgroup warp-specialization (that is for a compute-bound chunk pipeline; decode is a serial recurrence on resident state). + +- **Threads = 256, single role** (all threads cooperate on every op; the `group_reduce` / `cp_fwd` template, not `fused_fwd`'s producer/consumer split). At 256 threads the `[128,128]` fp32 state fragment is 64 fp32/thread (vs the 128/thread that `threads=128` forces and that the repo never does). +- **Grid** `T.ceildiv(DV, block_DV) · batch_size · H`, 1-D flattened `(bbhv,)`, decoded exactly as `flash_qla/ops/gated_delta_rule/chunk/hopper/fused_fwd.py:90-93`: `bbh, bv = bbhv // ceildiv(DV,block_DV), bbhv % …`; `bb, bh = bbh//H, bbh%H`; `bhg = bh // (H//Hg)`. One CTA owns `(b, V-head, V-column-tile [128, block_DV])` and runs the full `L`-step recurrence on its sub-state end-to-end (no cross-CTA combine). +- **State resident in registers.** `h_fragment = T.alloc_fragment((128, block_DV), "float32")` (the `fused_fwd.py:140` / `prepare_h.py:126` pattern). Loaded once before the loop (`T.copy` from the h0 slice if `use_initial_state` else `T.clear`), mutated in place across all `L` steps, stored once after. +- **GEMM operands live in SMEM.** `gemm_v1` reads SMEM operands and accumulates into a **fragment** (verified: `fused_fwd.py` lines 184/204/254/339 — never accumulates into shared). So each step downcasts the fp32 master to a bf16 operand copy `h_op_shared = T.alloc_shared((128, block_DV), qkva_dtype)` via `T.copy(h_fragment, h_op_shared)` (the `fused_fwd.py:190` fragment→shared copy, which downcasts fp32→bf16). The fp32 master stays in `h_fragment` for the decay/rank-1 accumulation. + +### The three K-contractions — all `gemm_v1` +`gemm_v1` is the **only** grounded K-reduction idiom in the repo. (A `T.Parallel` + `reduce_sum(dim=0)`-to-vector reduction over K=128 does **not** exist here — `reduce_sum(dim=0)` appears once, `fused_bwd.py:476`, reducing a 1-D fragment to a scalar — and is rejected.) + +- **Step 3 `kS = K @ S`:** `T.gemm_v1(k_op_shared, h_op_shared_decayed, kS_fragment, clear_accum=True)` — mirrors `U = K@S` at `fused_fwd.py:254`. +- **Step 5 rank-1 `S += k ⊗ v_new`:** the **`transpose_A` gemm-into-fragment** `T.gemm_v1(k_op_shared, vn_op_shared, h_fragment, transpose_A=True, clear_accum=False)` — the grounded rank-1 idiom at `fused_fwd.py:204` (Kᵀ@V′ accumulating into the register fragment). *(Correction: `fused_fwd.py:197` is a **scalar-broadcast** decay FMA, **not** a two-vector outer product, so a `T.Parallel(DK,block_DV): h[i,j]+=k[i]·v_new[j]` FMA is **not** grounded — it is a prototype-gated item (§11.E), reserved for the head-batched SMEM-state case where `gemm_v1` cannot accumulate into shared.)* +- **Step 6 `o = Q @ S`:** `T.gemm_v1(q_op_shared, h_op_shared_postupdate, o_fragment, clear_accum=True)`, then `o_fragment *= scale`. +- **Decay (step 2):** `for j_k, j_v in T.Parallel(DK, block_DV): h_fragment[j_k,j_v] *= decay` (`fused_fwd.py:197`). + +### Step ordering (critical — post-update read) +Per step: (a) l2norm already done host-side; (b) decay `h_fragment` in place; (c) copy/downcast `h_fragment → h_op_shared`; (d) gemm `kS` on decayed state; (e) `v_new = β·(v − kS)` in fp32; (f) `transpose_A` rank-1 gemm into `h_fragment`; (g) copy/downcast `h_fragment → h_op_shared` again; (h) gemm `o` on post-update state, `*= scale`, cast, store `o[b,t,h, bv-slice]`. After the loop: `if store_final_state: T.copy(h_fragment, final_state[b,h,:, bv-slice])` once. + +Note: because each step's `v_new` depends on the running state, the `D` tokens **cannot** be batched into one gemm — all three contractions are **per-step, M=1** (a single token). So §11.A's `M=1` feasibility gate applies to all three, every step. + +### M-dim (q_len) and the gemm +Each step is a single token, so the gemm M (token) dim is **1**. **Open feasibility item (§11.A):** confirm `gemm_v1` accepts `M=1` (every repo `gemm_v1` is `M=64`). Mitigation: zero-pad the M (token) dim of `q/k/vn` staging to 16 (zero rows contribute zero to the rank-1 update and produce garbage `o` rows we don't store). Prototype `M=1` first; this gates the whole engine. + +--- + +## 4. Memory & thread layout (core) + +- **State HBM layout:** fp32 contiguous `[B,H,128,128]` (the `h0_shape`/`ht_shape` of `fused_fwd.py:70-71`). Innermost V is unit-stride, so the column slice `[b,h, 0:128, bv·block_DV:(bv+1)·block_DV]` is a coalesced / auto-vectorized `T.copy` target. Decode `B = real_batch_size`, one state row per sequence. +- **CTA tile:** `[128, block_DV]`, `block_DV ∈ {128,64,32}`. +- **Register budget:** `block_DV=128 → 64` fp32 state regs/thread at 256 threads; plus `kS_fragment`, `o_fragment`, `v_new` (few/thread each). Target `nreg ≈ 128–160` via `T.set_max_nreg` if needed. `block_DV=64/32` drops pressure 2×/4×. Do **not** materialize a full `[128,block_DV]` product fragment. +- **SMEM:** `h_op_shared [128, block_DV]` bf16 (≤ 32 KB); per-step staging `q/k` (`block_S_pad × 128` bf16, ~4 KB), `v` (`block_S_pad × block_DV`). Total < 64 KB even at `block_DV=128`. Bank conflicts: pad small staging tiles' trailing dim by +1 (the `cumsum.py:47` / `kkt_solve` `17=16+1` idiom) and `T.use_swizzle(10)` (`fused_fwd.py:160`). +- **Reduction-free across CTAs:** because we split V (not K), every CTA holds all 128 K-rows for its V-columns, so the only K-sum (`kS`, `o`) is fully resident. No atomics, no grid-sync, no stitch kernel. + +--- + +## 5. Occupancy strategy (V-split only) + +The decode analogue of `cp_context.py`'s auto-CP, computed **host-side** in the wrapper using the `fused_fwd.py:602-608` ladder verbatim: + +``` +TARGET_NUM_CTAS = int(MULTI_PROCESSOR_COUNT * 0.7) # fused_fwd.py:12 ; H100 132 SM → 92 +grid_base = real_batch_size * H +if grid_base >= TARGET_NUM_CTAS: block_DV = 128 # 1 CTA / head +elif grid_base*2 >= TARGET_NUM_CTAS: block_DV = 64 # 2 CTAs / head +else: block_DV = 32 # 4 CTAs / head +``` + +- **Why V-split, not K-split:** the V-column split is reduction-free (decay, `v_new`, rank-1, output are all per-V-column; the only K-contraction is fully resident). **K-split is rejected for `L>1`:** step-3 `kS` contracts all K and feeds `v_new` into step `t+1` nonlinearly, so a K-split needs the per-step `kS` partials summed across CTAs *inside every step*; with no atomics/grid-sync, an HBM-scratch + separate reduction kernel runs only after the first kernel completes and cannot feed step `t` of a running fused loop. K-split forfeits multi-step fusion (back to `L` launches) and is correct only at `L=1`. We do **not** use it; V-split alone gives up to 4×. +- **Honest framing:** server `B·H ≥ TARGET` → `block_DV=128`, full SM fill, **bandwidth-bound**. Single-request `B=1, H=16` → `block_DV=32` → 64 CTAs (~48% of SMs); this regime is **latency-bound, not bandwidth-bound** (16 heads × 128 KB = 2 MB moves in < 1 µs, below the launch + serial-`D`-step-dependency-chain floor). V-split here is a latency-hiding / occupancy knob, not a BW lever — do not claim "near bandwidth-bound in both regimes." For extreme `B=1, Hg≤2` TP, accept partial SM fill; the recommended remedy is the **caller** batching speculative requests / using CUDA graphs to amortize launch (owner decision §13). + +--- + +## 6. Numerics + +- **fp32 everywhere for state & accumulation** (`accum_dtype="float32"`): `h_fragment`, `kS_fragment`, `o_fragment`, `v_new`, any l2norm sum. q/k/v/o are bf16/fp16, cast to fp32 on load into the math and back to `o_dtype` only on the final `T.copy(o_fragment → o-slice)`. The gemm **operand** copy of the state is bf16 (`h_op_shared`), but the fp32 master drives decay/rank-1. **`kS` must stay fp32 before the subtract** in `v_new = β·(v − kS)` (catastrophic-cancellation safety, matching the chunk reference). +- **Decay via `exp2`:** `decay = T.exp2(g_t · 1.442695)` (`1.442695 = log₂ e`) under `@tilelang.jit(pass_configs={TL_ENABLE_FAST_MATH: True})` — the repo-wide idiom (`fused_fwd.py:236`, `prepare_h.py:182-183`). **g is consumed raw — never cumsum'd** (the chunk-path cumsum is identity at `chunk_size=1`; applying it gives `err ≈ 1.1`). +- **Scale:** baked as a compile-time literal, applied to q **only at the output GEMV**. When `use_qk_l2norm_in_kernel` is on, scale must **not** be folded into q at load — `l2norm(q·scale) == l2norm(q)` cancels it (footgun). +- **l2norm (host-side — DECIDED):** `recurrent_gated_delta_rule` calls `flash_qla.utils.l2norm(q)/l2norm(k)` exactly as `chunk_gated_delta_rule:221-223` (`rsqrt((x·x).sum(-1)+1e-6)`, `eps=1e-6`, cast back to input dtype). Bit-matches the chunk op; avoids the ungrounded in-kernel `T.rsqrt` + row-reduce. No in-kernel l2norm path in v1. +- **Tolerance:** tests compare at the existing harness bars (0.02 relative for `o` and `final_state`), not fp64-idealized `1e-10`. + +--- + +## 7. Head-batched GQA variant (server regime) — build after the core + +> **AS-BUILT NOTE (built + measured 2026-06-15): row-stack specialization, NEUTRAL-to-WORSE → auto-off, forceable.** The head-batched variant IS now implemented as a **gemm-free row-stack**: one CTA owns `(bb, K/Q head-group hg, V-tile bv)` and processes all `grp = H//Hg` V-heads `h = hg*grp + i` (sharing q/k); state is stacked `S[grp*block_DV, DK]` with row `gv` → head-band `gv//block_DV`, channel `v0 + gv%block_DV`. The "one wide-N gemm" framing below becomes a wider `T.reduce_sum`/`T.Parallel` over the stacked rows. It ships as `tilelang_fused_recurrent_gdr_fwd_hb` + a forceable `head_batch` flag on `recurrent_gated_delta_rule`; **21 tests pass on H100** (`tests/test_head_batch_gdr.py`, incl. a within-group band-swap negative control). **But the benchmark (`benchmark/bench_head_batch.py`) shows it is NOT a win:** grp=2 = **0.98–0.99×** (neutral), grp=4 = **0.74–0.87×** (regression). The row-stack trades CTA count (`B·H → B·Hg`) and uses bigger 512-thread CTAs at grp=4, which outweighs its only saving — the q/k **load** dedup (sub-1% on this memory-bound kernel). So `head_batch` defaults **OFF** (auto never selects it) and is exposed only for experimentation. **Two TileLang limits learned:** (1) every `[M]` fragment must be accessed over the FULL `Parallel(M,…)` range — a partial-offset write breaks affine-map inversion (`InverseAffineIterMap` fails); route all per-head divergence through GLOBAL with derived index `hg*grp + gv//block_DV`. (2) A shared `[grp]` band read by `gv//block_DV` inside a hot Parallel does NOT lower either — so per-head gating **cannot** be deduped. That is why the **gated-verify head-batch was NOT built**: it would pay M× redundant `softplus/exp/sigmoid` (costing more than the l2norm-dedup it was meant to win), and the row-stack structural penalty dominates regardless. + +A compile-time **specialization of the same jit factory** (keys `head_batch: bool`, `group_size: int`, alongside `block_DV`), not a separate kernel — matching how the chunk factory branches its prim_func on `is_varlen` / `is_cp` / `block_DV`. `head_batch=False` traces the core V-split body unchanged. + +**Idea.** In GQA, `group_size = H // Hg` V-heads share one Q/K head `hg`. One CTA owns `(b, head-group hg, V-tile bv)` and processes all `group_size` heads `h = hg·group_size + i`. Per step: load shared `q_t, k_t` **once**; per head load `g, β, v` into the head's column band `[i·blockV:(i+1)·blockV]`; decay-scale each band by its own per-column g; `kS = k_t @ S` as **one wide-N gemm** over `N = group_size·blockV`; `v_new` per band; rank-1 update in place; `o = q_t @ S` as one wide-N gemm (post-update); store each head's state to its disjoint slot after the `L`-loop. + +**State residency.** `gs≥2` cannot be a register fragment (e.g. `gs2/bv128` = 256 fp32/thread). Batched state lives in **SMEM** as one concatenated fp32 tile `h_state_shared[128, gs·blockV]`; head `i` owns columns `[i·blockV:(i+1)·blockV]`. A bf16 operand tile `h_op_shared[128, gs·blockV]` is kept for the wgmma-capable K-contractions. + +**SMEM budget** vs ~227 KB usable (Hopper opt-in). Hard gate `state + op = 1.5 · 128 · gs·blockV · 4 B`, and `gs·blockV ≤ 512`: + +| combo | state+op+staging | occ | verdict | +|---|---|---|---| +| gs2/bv128 | ~206 KB | 1 | feasible (force single-buffer q/k) | +| gs2/bv64 | ~104 KB | 2 | feasible | +| gs2/bv32 | ~54 KB | — | feasible | +| gs3/bv64 | ~154 KB | 1 | feasible | +| gs3/bv32 | ~80 KB | — | feasible | +| gs3/bv128 | ~304 KB | — | **hard-reject** | +| gs4/bv128 | 256 KB state alone | — | **hard-reject** | +| gs4/bv64 | ~206 KB | 1 | feasible (single-buffer only) | +| gs4/bv32 | ~104 KB | 2 | feasible — **cleanest** | + +**Feasible set = `{gs2:[128,64,32], gs3:[64,32], gs4:[64,32]}`.** `gs ∈ {1,2,3,4}` are all real configs from the benchmark table; `gs=3` (non-power-of-2 `N`) is first-class. + +**Two repo-verified constraints (resolved blockers):** +1. **Rank-1 update is always the `T.Parallel(DK, gs·blockV)` FMA** — `gemm_v1` only accumulates into a *fragment*, and the head-batched state lives only in SMEM, so the `transpose_A` rank-1 gemm-into-shared is **unexpressible**. The wide-N win is preserved on the two *read* contractions (`kS`, `o`), which are the throughput-critical ones. +2. **Fused downcast into the decay pass:** the decay `T.Parallel` writes both outputs in one sweep — `h_state_shared` (fp32, updated) and `h_op_shared = cast_bf16(updated)` (the gemm operand) — and the rank-1 FMA likewise refreshes both. This avoids an unverified shared→shared downcast copy (no repo precedent). + +**Auto-selection** (host, composes on the core ladder, gates on **`B·Hg`**): +1. `group_size == 1` → always core V-split. +2. Pick the **largest** feasible `blockV` (maximize `N = gs·blockV`) s.t. `(gs,blockV)` is feasible **and** `grid_hb = B·Hg·ceildiv(DV,blockV) ≥ TARGET`. +3. **Occupancy fence:** with `occ_hb = floor(227KB / smem_per_cta)`, forbid head-batch when `occ_hb==1` and the core keeps a full extra wave — unless `D==1` (where `M=1` under-utilization makes the trade worth it). +4. **No-win fence:** if `N = gs·blockV ≤ 128` (no wider than core `bv=128`), or core already runs `bv=128` with `B·H ≥ TARGET`, skip. +5. **`gs4` default `bv32`:** `tiles(32)=4=gs` → `grid_hb = B·Hg·4 = B·H` (core grid exactly) → wide-N (`N=128`) at zero occupancy cost, occ=2. The cleanest honest win. +6. else core V-split. + +**Honest perf.** State bytes are **not** reduced (`gs·64KB` total, same as `gs` separate CTAs). q/k operand reuse saves `(gs−1)·512 B/step` — marginal vs the `gs·128·blockV·2 B` bf16 state operand the gemm reads. The wide-N **wgmma** fill is a genuine win **only for `D>1`** (`M=D` fills wgmma rows); at `D=1`, `M=1` is below the wgmma min M-tile (64) and degenerates to FFMA regardless of `N`, so the `D=1` benefit is instruction-issue (1 wgmma vs `gs`) + operand reuse + lower register/SMEM pressure. Marginal-to-negative when: `N` already ≥128; `B·Hg < TARGET`; `occ_hb==1` with a core extra wave; or `D≥4` (M already fills the tensor core). This is a server-regime, above-the-saturation-knee trade. + +**Composition.** Only q/k are shared (read-only); everything mutated is head-private (`h = hg·gs + i`). Ragged `L` is per-CTA (all group heads share sequence `b`, hence the same `L` and the same q/k masking); each head reads its own `g/β/v` and writes its own `o`/`final_state` slot. Result is **bit-identical** to `gs` independent core CTAs (the wide-N gemm computes each column as an independent fixed-order K-reduction — concatenation-invariant), **conditional on** the per-band scalar-gather index (§11). + +--- + +## 8. API & integration + +New package `flash_qla/ops/gated_delta_rule/fused_recurrent/` (sibling to `chunk/`), SM90-gated at every import boundary exactly like `chunk/__init__.py:10-13`. + +``` +fused_recurrent/ + __init__.py # fused_recurrent_gdr_fwd wrapper + recurrent_gated_delta_rule; SM90 guard; imports l2norm + hopper/ + __init__.py # from .fused_recurrent_fwd import fused_recurrent_gdr_fwd + fused_recurrent_fwd.py # @tilelang.jit factory + low-level python wrapper (mirrors fused_fwd.py two-part structure) +``` + +**JIT factory** (`hopper/fused_recurrent_fwd.py`), all args compile-time specialization keys; `batch_size` and `num_tokens(=q_len)` are `T.dynamic` so one kernel serves `D=1..8` with no recompile: + +```python +@tilelang.jit(pass_configs={tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True}) +def tilelang_fused_recurrent_gdr_fwd( + H, Hg, DK, DV, scale, + accum_dtype, qkva_dtype, g_dtype, b_dtype, h0_dtype, ht_dtype, o_dtype, seqlen_dtype, + use_initial_state, store_final_state, has_seqlens, + head_batch=False, group_size=1, block_DV=128, threads=256, +): # -> inner @T.prim_func (kernel name must end *_kernel_kernel per CLAUDE.md) +``` + +Tensor signature: `q,k: (batch, q_len, Hg, 128)`; `v,o: (batch, q_len, H, 128)`; `g,beta: (batch, q_len, H)` fp32; `h0, final_state: (batch, H, 128, 128)` fp32; `seqlens: (batch,) int32`. Grid `ceildiv(DV,block_DV)·batch·H` (core) or `·batch·Hg` (head-batch). + +**Low-level wrapper:** +```python +def fused_recurrent_gdr_fwd(q, k, v, g, beta, scale=None, initial_state=None, + output_final_state=True, seqlens=None, use_qk_l2norm_in_kernel=False): + # infer (B,q_len,Hg,K)=k.shape, (.,.,H,V)=v.shape; assert K==V==128, H%Hg==0; scale = scale or K**-0.5 + # host auto-selection -> (head_batch, group_size, block_DV) per §5/§7 + # use_initial_state = initial_state is not None (alloc empty h0 if None, like fused_fwd:584-587) + # o = empty_like(v); final_state = empty((B,H,128,128), fp32) + # materialize seqlens or full-q_len int32[B]; clamp L>=1 + # compile via factory; launch; return (o, final_state or None) +``` + +**High-level:** +```python +def recurrent_gated_delta_rule(q, k, v, g, beta, scale=None, initial_state=None, + output_final_state=True, use_qk_l2norm_in_kernel=False, + seqlens=None, head_first=False): + # assert q.dtype==k.dtype==v.dtype and != fp32; not head_first; v.shape[2] % k.shape[2] == 0; K==V==128 + # scale = k.shape[-1]**-0.5 if None + # if use_qk_l2norm_in_kernel: q=l2norm(q); k=l2norm(k) # HOST-side + # o, final_state = fused_recurrent_gdr_fwd(...); return o.to(q.dtype), final_state + # NO autograd.Function (inference/forward-only) +``` + +**Exports:** add `recurrent_gated_delta_rule` to `flash_qla/ops/gated_delta_rule/__init__.py` and re-export at package top level alongside the chunk entry points. Reused utils only: `flash_qla.utils.l2norm`, the `TARGET_NUM_CTAS` ladder, and the standard TileLang primitives already used in `fused_fwd`/`prepare_h`/`cp_fwd`/`group_reduce`. No cu_seqlens packing, no `chunk_offsets`, no `kkt_solve`, no cumsum, no autograd. + +--- + +## 9. Test plan + +Mirror `tests/test_gdr.py` (0.02 relative tolerance, 1000-iter stability loop, SWA `g=0` mask, h0 on/off, GQA `Hk1`). For a length-`L` single sequence, `decode_recur`'s `o[:, :L]` and `final_state` must match the chunk reference on that `L`-prefix to fp64 roundoff. +3. **Kernel vs `decode_recur`** (bf16 io / 0.02): sweep `q_len ∈ {1,3,8}`; ragged `seqlens=[1,5,8]` mixed in one batch (compare only `o[:, :L_b]`; `final_state[b]` vs the length-`L_b` reference); GQA `Hk=2,Hv=4` (verify head `h` reads k-head `h//2`), MQA `Hk=1`, no-GQA `Hk=Hv`; `initial_state` on/off; explicit non-default scale; `g=0` SWA heads. Run with the SWA mask + 1000-iter loop. +4. **Negative controls:** (a) cumsum'ing g must break (`~1.1`); (b) pre-update read must fail (`~7.8e4`), decay-after-update must fail (`~2.2e3`); (c) `mod` GQA mapping must fail (`~1.8e4`); (d) garbage in `t≥L` positions must change committed `final_state` / `o[:, :L]` by **exactly 0.0**. +5. **Head-batch tests:** **bit-identity** (not 2% tol) of head-batch vs `gs` independent core CTAs with **per-head-distinct g/β** (catches a decay/β band-gather bug); explicit `gs=3` (non-pow2 `N=96/192`) compile + numeric; static-reject asserts that `gs4/bv128` and `gs3/bv128` never reach the JIT. +6. **Feasibility smoke tests** (before full sign-off): compile+run a one-step `M=1` kernel (and `M`-padded-16 fallback) vs `decode_recur` at `L=1`; confirm `T.serial(L)` lowers with runtime per-CTA `L` in the single-role form; confirm the `[128,block_DV]` fp32 fragment compiles at `threads=256` for `block_DV ∈ {128,64,32}` without spill (check nreg). +7. **Signature test** mirroring `tests/test_function_signature.py`: assert the `recurrent_gated_delta_rule` / `fused_recurrent_gdr_fwd` signatures and the `(o, final_state)` return contract. + +--- + +## 10. Benchmark plan + +Add `benchmark/bench_recurrent_gdr.py` mirroring `bench_gated_delta_rule.py`. Baseline: FLA `fused_recurrent_gated_delta_rule`, plus `D` separate `D=1` calls (to quantify the fusion win). Use `flash_qla.utils.profile`. + +- **Report** wall time and achieved HBM GB/s (% of ~3.35 TB/s peak) **per regime** — not a single "near-peak" claim. Roofline at `B·H ∈ {4,16,64,512,2048}` so the bandwidth claim is falsifiable: server `B·H ≥ TARGET` near peak (state I/O is 94–99% of bytes; time ~ `B·H·128KB/BW`, D-independent); single-request `B·H ≪ TARGET` latency-bound (report launch + serial-`D`-chain time, not bytes/BW). +- **Sweep:** (1) `B ∈ {1,8,64,256}`, `H ∈ {4,16}`, `Hg ∈ {H, H/2, 1}`; (2) `D ∈ {1,4,8}` to show D-fold state-I/O amortization (plot per-token effective state bytes = 128KB/D); (3) `block_DV` auto-selection across the `B·H` sweep; (4) ragged vs uniform `seqlens` (verify ragged adds no measurable overhead); (5) head-batch on/off across `gs ∈ {2,3,4}` to validate the auto-selection win/skip. + +--- + +## 11. Open feasibility items (prototype-gated; do not block the architecture) + +> **AS-BUILT NOTE: all gates RESOLVED on H100 (workspace `netragratis`).** Outcomes: **A.** `gemm_v1` at `M=1` **fails** ("M must be divisible by 16") and even M-padded forms hit warp-partition / single-row-fragment walls → **moot**, the kernel pivoted to gemm-free (§3). **B.** single-role `T.serial(runtime L)` is **usable** as-is (no static fallback needed). **C.** moot under gemm-free; final tile is `block_DV=64 @ threads=128`, no spill. **D./E.** moot — head-batch not built (§7). **F.** moot — no wgmma in the gemm-free path; state stays fp32 in-register. **Bonus:** in-kernel fused gating is feasible (`log/log1p/rsqrt/sqrt/sigmoid/exp2` all lower on SM90 — §6's host-only assumption was over-cautious). The items below are the pre-build gates as originally written. + +These pick between grounded fallbacks; the architecture holds regardless. + +- **A. `gemm_v1` at `M=1`** (every step is a single token; every repo `gemm_v1` is `M=64`). The **root gate** — all three contractions (`kS`, rank-1, `o`) are `M=1`. Mitigation: zero-pad the M dim to 16 (zero rows contribute zero; discard garbage `o` rows). Compile-test `M=1` before any other work. +- **B. `T.serial(L)` with a runtime per-CTA `L` on a single-role kernel.** Grounded in `prepare_h.py:166`, but that kernel is warp-spec; confirm in the single-role form. **Static fallback:** `T.serial(q_len)` with an `if t1` → V-split only; broken `chunk_size=1` provenance → validate vs `chunk_size=64`; scale-before-l2norm cancels → scale at output GEMV only; single-request is latency-bound → framed honestly; head-batch rank-1-into-shared unexpressible → `T.Parallel` FMA; head-batch state too big for registers → SMEM, `gs4/bv128`+`gs3/bv128` hard-rejected. + +**Owner decisions (recorded):** +- l2norm location → **host-side** (decided). +- GQA-group head-batching → **in scope now** (decided; §7). +- Extreme `B=1, Hg≤2` TP → **accept V-split ~48% SM fill; caller batches / CUDA graphs** (decided). Not pursuing an `L=1`-only K-split path. + +**Ragged contract (resolved by the SGLang grounding):** the **verify** path receives ragged lengths as **`cu_seqlens`/`query_start_loc [N+1]` with `B=1` flattened varlen** (a *different* CTA→request prologue than this spec's dense `[B] seqlens` — derive `bb`/per-request token ranges device-side à la `kkt_solve.py:245-247`, never `.item()`). The dense `[B] seqlens` form is for the standalone decode follow-on. See the verify spec (`2026-06-15-gdn-verify-sglang-design.md`). + +**Verify needs `D` up to 12** (chunk-12 draft), which **exceeds this spec's `1..8` design point** — re-derive `nreg` (§11.C) and the `M=1` gemm behavior (§11.A) at `D=12` with a compiled measurement; do not assume the `1..8` budget carries over (the state fragment is `D`-independent, but the per-token-state writes and the serial dependency chain are not). diff --git a/docs/superpowers/specs/2026-06-15-gdn-verify-sglang-design.md b/docs/superpowers/specs/2026-06-15-gdn-verify-sglang-design.md new file mode 100644 index 00000000..58f4211d --- /dev/null +++ b/docs/superpowers/specs/2026-06-15-gdn-verify-sglang-design.md @@ -0,0 +1,246 @@ +# FlashQLA GDN Verify + SGLang Integration — Design Spec + +Date: 2026-06-15 +Status: Design approved in principle (pending written-spec review) +Depends on: `2026-06-15-gdn-decode-kernel-design.md` (the recurrent decode **spine** is a hard prerequisite). +Grounding: verified against SGLang upstream `f18d38d` (via `gh api` blob fetch) and a grep of this repo's TileLang usage. + +> **IMPLEMENTATION STATUS (2026-06-15, validated on a Modal H100).** V1 verify is **built and passing** (30 decode+verify tests at 0.02 rel). Key deviation from this spec: the kernel is **gemm-free** — the GEMVs are `T.reduce_sum` over the K dim and the rank-1 is a `T.Parallel` outer product, with the state kept `[block_DV, DK]` (V-major) in an fp32 fragment (Gate 1 showed `gemm_v1` needs M%16==0 *and* num_warps≤N/16, which is unworkable for single-token GEMVs at small `block_DV`; the gemm-free path is simpler and keeps the state fp32). Because the state is already V-major, the SGLang V-major pool store is a **direct** write (no transpose). Delivered: paged `state_indices` gather/scatter (slot<0 skip), bf16 pool, per-token intermediates (gated on the pool mask), no-commit, varlen `cu_seqlens`, host **and** in-kernel fused gating (`A5` — Gate 4 showed `log/log1p/rsqrt/sigmoid/exp` lower on SM90, so it is **not** prototype-gated), and CUDA-graph safety. **Benchmark vs FLA/Triton** (netra-server's `fused_sigmoid_gating_delta_rule_update`, target-verify, on H100, parity-gated): after autotuning the tile to **`block_DV=64 @ threads=128`** (the old `block_DV=128` was occupancy-starved), FlashQLA is **≥ FLA across all regimes** — **~3.3× faster** in the single-request/TP and small-batch latency-bound regime (FLA carries a ~53–70µs `num_warps=1` fixed floor; FlashQLA ~15–20µs), and **on par (0.99–1.02×) at large batch** (both ~2.0–2.18 TB/s, 60–65% peak). `benchmark/bench_vs_fla.py` + `benchmark/tune_verify.py`. **V2 (chunk-based) is NOT built.** **Large-batch optimization investigation (measured, H100):** a stable 2–3× at large batch is **physically impossible** for this workload — the verify is **compute-bound on the serial GDN recurrence** (the per-token `kS`/`o` K-reductions), not write-bound. A no-write run isolates the compute floor at ~1139µs vs full-state-write 1376µs at `N=256,T=12`, i.e. the per-token state writes are only **+21%**. So byte-reduction levers were prototyped and rejected: fp8 ibuf = 1.18× (+5% state error); a replay-tape (store `k_norm/v_new/decay` per token, reconstruct on commit) = **slower (0.66×)** because small scattered writes coalesce worse than one contiguous state store. Both FlashQLA and FLA run the identical recurrence → at parity at large batch by construction; the ~3.3× win is real only in the **latency-bound** regime (small/single-request batch, the dominant speculative-decode case). Shipped + pushed to `NetraRuntime/FlashQLA` (branch `feat/gdn-decode-kernel`, PR #1). Files: `flash_qla/ops/gated_delta_rule/fused_recurrent/`; tests `tests/test_decode_gdr.py`, `tests/test_verify_gdr.py`; benches `benchmark/bench_recurrent_gdr.py`, `bench_vs_fla.py`, `tune_verify.py`. + +> **OPTIMIZATION H1 — gating + qk-l2norm DEDUP PRE-PASS (built + measured 2026-06-15, SHIPPED, regime-gated ON).** The in-kernel-gated verify kernel (`fused_recurrent_gdr_verify_gated_fwd`, "variant A") recomputes `g`/`β` (softplus/exp/sigmoid) + qk-l2norm **inside the per-token hot loop**, once per `(token, V-head, V-tile)` CTA → redundant `grp·n_vt` (=4 at `Hk16/Hv32`, `block_DV=64`) for l2norm and `n_vt` for gating. **Verify-first ceiling probe** (`benchmark/probe_h1_ceiling.py`, gated-A vs host-gated-B on identical raw inputs, parity 0.005): the in-loop recompute costs **14–16% at `T=12` large batch**, 9–10% single-request — scaling with `T` (per-token recompute). **The naive torch pre-pass is ~145µs (l2norm `@torch.compile` + allocs) — a non-starter; the pre-pass MUST be a fused TileLang kernel.** Built `tilelang_gdr_verify_prepass` (`fused_recurrent_verify.py`): one CTA per token computes l2norm for all `Hk` K-heads via a parallel `[Hk,128]` `reduce_sum(…,dim=1)` (writing `q_n`/`k_n` directly from a full-range `T.Parallel` — the head-batch global-write idiom, **no** per-row `[1,DK]` `T.copy`) and `g=-exp(A_log)·softplus(a+dt_bias)` (RAW log-decay) + `β=β_mul·sigmoid(b)` for all `Hv` V-heads, written as one contiguous `[1,Hv]` row. It feeds the **unchanged** host-gated main kernel (`fused_recurrent_gdr_verify_fwd`), whose hot loop then has **zero** transcendentals/l2norm. Wired into `recurrent_gated_delta_rule_verify(fuse_gating=True)` via `should_use_prepass(N,H,total_tokens)`; `prepass=True/False` forces it. Persistent unbounded (never-evicting) scratch cache → capture-safe after warmup (a fresh `test_prepass_path_cuda_graph` proves graph capture+replay). **Measured (H100, `benchmark/bench_prepass.py`, eager AND CUDA-graph, parity 0.003–0.006 bit-faithful to variant A):** **+8% to +24%** in the bandwidth-bound `T≥4` regime, verified STABLE across two eager runs and the production CUDA-graph path — `N=256,T=12` **1.15×** (1626→1418µs, −208µs), `N=64,T=12` **1.18×**, `N=16,T=12` **1.22–1.24×**, `N=32,T≥4` **1.10–1.22×**, `N=8,T=12` **1.22×**; graph replay 1.08–1.24×. **No regression elsewhere:** single-request / small-batch / `T=1` route to variant A (the 2nd-launch tax exceeds the small ceiling there — measured 0.59–0.89×). The gate `T_avg≥4 ∧ N·H·(1+T_avg)≥3000` is **conservative**: it fires only where the win is stable across runs+graph (≥1.08× even in a launch-overhead-contended eager run). Three near-boundary regimes that win in clean eager/graph but are launch-noise-flaky (`N=4,T=12`; `N=8,T=8`; `N=16,T=4` — work 1664/2304/2560) are routed to variant A → zero regression risk; the lowest stable win is `N=8,T=12` (work 3328). **Two iterate-loop learnings:** (1) the ceiling probe (free precompute) is optimistic — the first pre-pass impl (grid=`total_tokens`, **serial `Hk` loop** of single-row `[1,128]` reduces) was a measured **~6× slowdown** (~297µs); the parallel `[Hk,128]` reduce fixed it (~15–40µs). (2) **Hopper copy-layout inference rejects an fp32 contiguous extent < 128 bits** (`makeGemmABLayoutHopper`, `Unsupported layout … element_size=32`) — a per-K-head `[1,grp]` gate write fails for `grp∈{1,2}`; writing the full `[1,Hv]` row (`Hv≥4`) is the fix. Files: `fused_recurrent_verify.py` (`tilelang_gdr_verify_prepass`, `fused_recurrent_gdr_verify_prepass`, `should_use_prepass`, `get_prepass_scratch`); `fused_recurrent/__init__.py` (wiring); tests `tests/test_prepass_gdr.py`; bench `benchmark/bench_prepass.py`, probe `benchmark/probe_h1_ceiling.py`. + +> **OPTIMIZATION SWEEP (2026-06-15) — remaining levers profiled; nulls recorded so they are not re-attempted.** After H1 + the decode `block_DV=128` fix, a differential attribution (`benchmark/probe_sweep_attribution.py`) split the post-H1 verify main kernel: at `N=256,T=12` the no-ibuf recurrence floor is **1198µs (86% of the 1402µs total)** and the per-token `ibuf` write adds only **204µs (14.5%)** at an *impossible* 15.8 TB/s — i.e. the writes **overlap the compute-bound recurrence and are largely hidden**. Conclusions: **(a) the kernel is compute-bound on the serial `reduce_sum` K-reductions**, not store-bound. **(b) Recurrence elementwise-pass fusion = NULL:** folding the 4 `[block_DV,DK]` sweeps (`S*=decay`, `k·S`, rank-1, `q·S`) into 2 fused passes is bit-identical and correct (67 tests) but **perf-neutral** (within noise: `N=256,T=12` 1427 vs 1418µs) — the two tree-reductions dominate the critical path, the elementwise sweeps touch register-resident `S` and are not the bottleneck. Reverted (no benefit, don't churn the hot loop). **(c) `ibuf`-store vectorization = DROP:** the V-major store is already coalesced and ≤14.5% marginal (mostly hidden under compute) → low ceiling. **(d) Store-coalescing audit (decode/verify) = DONE:** the only uncoalesced write was the decode K-major transposed `ht` store, fixed by `block_DV=128` (~2×); the `h0` load (also transposed) coalesces at 128 too; verify `ibuf`/pool stores are V-major (already coalesced). **(e) H1 2nd-launch tax (to win at small batch too) = needs cross-module work** — eliminating the extra launch requires fusing gating+l2norm into SGLang's upstream conv epilogue (FlashQLA does not own it) or a 2-grid megakernel (high risk); the regime gate already prevents any small-batch regression. The chunk prefill fwd/bwd kernels are a **separate surface** (not swept here). Net: the two shipped wins (H1 verify +8–24%; decode `block_DV=128` ~2×) capture the tractable in-kernel headroom; the recurrence-reduction floor is the hard limit (a stable 2–3× at large batch remains physically impossible — identical recurrence to FLA). + +--- + +## 1. Overview & scope + +Extend FlashQLA to serve **SGLang speculative decoding** for GDN: a CUDA-graph-safe, paged, bf16-state **verify** kernel that emits per-token outputs **and** per-token intermediate states, without committing the final state. Decode (T=1) is a later/optional follow-on; this spec is verify-first. + +**Resolved scope (do not relitigate):** +- **Verify-first.** Build the shared infra (A) on top of the decode spine, then the verify kernel (B). Standalone single-token decode (C) is a follow-on; the verify engine already covers the multi-token recurrence. +- **Build BOTH verify engines and benchmark at T=12** (decision): **V1** = recurrent multi-step (per-token states native); **V2** = chunk-based (fast `o`) + per-token-state extraction. Ship the faster on **total verify-step latency (`o` + per-token states)**. The repo-grounded analysis predicts V1 wins (see §5); V2 is built to confirm with data. +- **Linear draft chain only** — target scheme is **DFlash** (public upstream SGLang linear-chain speculative verify). Tree-structured propagation (DDTree) is **out of scope** (not needed); the linear design doesn't preclude adding it later. +- **`allow_neg_eigval = False`** (decision): plain `β = sigmoid(b)`, exposed as a compile-time flag defaulting False. +- **bf16 state pool** (your deployment; `SGLANG_MAMBA_SSM_DTYPE`), kernel dtype-agnostic. +- **Scope boundary.** This kernel owns **only the GDN SSM recurrent state**. The conv state (`causal_conv1d_update` — tree-conv + `intermediate_conv_window`) runs **upstream in SGLang** (`gdn_backend.forward_extend`), before this kernel — it is **not** this kernel's responsibility. The author owes no conv-state handling. + +**The root gate for everything** is decode-spec §11.A: `gemm_v1` at `M=1` (each step is one token). Prove it on the Hopper box before any verify work; fallback is M-pad-to-16. + +--- + +## 2. Shared infra A (the reusable base for V1, V2, and decode) + +Infra A is a set of **compile-time keys + kernel-arg contracts + code hooks** the recurrent decode kernel (and the chunk kernel, for V2) gains. It is **not** five edits to an existing factory — the recurrent factory does not exist yet; A is built *with* the spine (decode-spec §12 order). + +### A1 — Paged in-kernel gather/scatter (GROUNDED, ship) +State lives in a caller-owned pool `[num_slots, H, V, K]` (V-major — see A4) indexed per request by `state_indices: [N] int32`. The gather/scatter is **in-kernel** — Python never slices `pool[idx]`. Idiom verified at `cp_fwd.py:141` (`cp_h0[seq_start_idx, bh, …]` with `seq_start_idx` an `int32` `alloc_var` loaded from a device tensor) and the double-indirection at `kkt_solve.py:245-247`. + +Per CTA owning `(request b, V-head bh, V-tile bv)`: +``` +slot = T.alloc_var('int32'); slot = state_indices[bb] +T.clear(h_fragment) # clear first +with T.If(slot >= 0): # NO T.Else exists in the repo + with T.Then(): + T.copy(pool[slot, bh, ], h_fragment) # bf16 → fp32 gather +… recurrence … +if not disable_state_update: + with T.If(slot >= 0): + with T.Then(): + T.copy(h_fragment, pool[slot, bh, ]) # fp32 → bf16 scatter +``` +**Clear-then-conditional-overwrite** (there is no `T.If/T.Else` in the repo). **Only `state_indices` carries the `-1` sentinel** (`PAD_SLOT_ID = -1`; `slot < 0 ⇒ skip`, slot 0 is valid — **not** vLLM's `slot ≤ 0 / NULL_BLOCK_ID=0`); guard with `slot >= 0` (resolved §10.4). **`intermediate_state_indices` is dense (`arange`, never `-1`)** — so the per-token ibuf write (§3) must be gated by the **pool-slot mask `state_indices[bb] >= 0`**, reusing the same `slot` var, **not** by `intermediate_state_indices`. (A `T.If(cache_slot >= 0)` guard would never fire — `cache_slot` is always valid — and would leave garbage in padded ibuf rows.) + +### A2 — Configurable-dtype pool + fp32-register accumulation (GROUNDED, ship) +fp32 master `h_fragment = alloc_fragment((128, block_DV), "float32")`. **Exactly three cast points:** (a) bf16→fp32 on the gather; (b) fp32→bf16 on the per-step gemm-operand staging copy (`fused_fwd.py:190`); (c) fp32→bf16 on intermediate/final stores. The cancellation-critical `v_new = β·(v − kS)` subtract stays **fp32** before any downcast. Pool dtype is a factory key driven by the caller's pool tensor dtype (your deployment: bf16). Under fp32 pools the per-token-write traffic doubles (~832 KB/head vs ~416 KB/head) — a caller decision, not a code risk. + +### A3 — Graph-safe entry (GROUNDED, ship) +Three rules for the captured call: +1. **No host sync.** Never call `prepare_chunk_offsets` (`index.py:138` ends in `.item()`) or `prepare_chunk_indices` (`index.py:80` `.tolist()`). Ragged length comes from a device `cu_seqlens` consumed via in-kernel `alloc_var` load (`kkt_solve.py:39,247`), never `.item()`'d. +2. **No allocation.** Every buffer (`o`, pool, intermediate buffer, all index tensors) is **caller-preallocated**; `out_idx` stays commented (`fused_fwd.py:16`) so outputs are positional. The graph-safe wrapper does **zero** `torch.empty` (vs `fused_gdr_fwd:561-600`). +3. **Static shapes.** `block_DV` from the `MULTI_PROCESSOR_COUNT*0.7` ladder (`fused_fwd.py:602-608`) using only static `N·H`; `grid = ceildiv(DV,block_DV)·N·H`; `T=12` fixed per graph; `cache_steps` read from `intermediate_states_buffer.shape[1]` as a python int (capture-constant), **not** the runtime `cache_steps` arg (SGLang ignores it). + +**"Host-side" means PyTorch, not TileLang — NOT outside capture.** The gating + l2norm (A5) depend on per-step `a, b, q, k`, so they are PyTorch ops that run **inside** SGLang's captured graph; they are capture-safe (pure elementwise, no `.item()`, no new allocation, static shape) and must **not** be lifted out of the per-step graph. + +### A4 — State layout: `state_v_first = True` is the SGLang contract (GROUNDED, default flipped) +SGLang's pool **and** intermediate buffer are **V-major `[.,H,V,K]`** — established from pointer **arithmetic** (`o_v*K + o_k`; `make_block_ptr (V,K),(K,1)`; `temporal_state_shape=(HV, head_dim=V, state_size=K)`), **not** docstrings (which are stale and mutually contradictory). FlashQLA-native is K-major `[.,H,K,V]` (`fused_fwd.py:70`). `state_v_first` is a compile-time key tracing to a different prim_func body (like `is_varlen`/`is_cp`); **no runtime transpose** (incompatible with paging). It applies to **both** the pool and the intermediate buffer (the scheduler reads the buffer back; a wrong major-order silently restores a transposed state). The wrapper derives it authoritatively from `pool.stride()/shape`. + +**Critical test consequence:** because `K==V==128`, a layout error is **numerically silent** in any equal-dim test. The hard validation gate is a **per-head-distinct-gate bit-identity** test vs the FLA reference (§8), never a shape/equal-value test. Whether `T.copy` from a `[K, block_DV]` fp32 fragment into a `[V,K]`-declared slice emits a strided TMA store without an SMEM transpose stage is **prototype-gated** (§9 Gate 6) — validate on a non-square probe (`DK≠DV`) or byte-compare with FLA, never the `128==128` test. + +### A5 — Gating: host-side primary, in-kernel prototype-gated +The repo's TileLang uses **only `T.exp2`** (grep: zero `log`/`log2`/`log1p`/`rsqrt`/`sqrt`/`sigmoid`/`exp`). In-kernel `softplus` needs `log1p`; in-kernel l2norm needs `rsqrt` — **neither is grounded**, and the decode spec already decided host-side l2norm. The Gate-4 probe (`hasattr(T,'log2'/'rsqrt')` + an SM90 lower test) could not run here (no TileLang locally) — it **must run on the Hopper box** (§9 Gate 4). + +- **Primary (capture-safe, ships):** compute `g`, `β`, and qk-l2norm in **PyTorch (not TileLang)** — a tiny `[1, N·T, H]` elementwise op + l2norm that runs **inside** SGLang's captured graph (capture-safe), passing **pre-activated `g` (log-decay) and `β` (post-sigmoid)** into the kernel — exactly what SGLang's `_update` kernel and FlashQLA's chunk path already consume. Only the per-step decay `exp2(g·1.442695)` is in-kernel (pure `exp2`, grounded). +- **In-kernel fusion (req #5, fast-follow):** accept raw `(A_log, dt_bias, a, b)` and compute `g`/`β`/l2norm in-kernel — **only if** Gate 4 confirms `log2`/`rsqrt` lower on SM90. `sigmoid` and the decay `exp2` are pure-`exp2` (groundable); `softplus`/l2norm-`rsqrt` are the blockers. + +### Exact gating math (grounded — identical across vLLM / SGLang / FLA) +``` +softplus(x) = log(1 + exp(x)) for x ≤ 20 (threshold), else x # softplus_beta=1.0 +g = -exp(A_log) · softplus(a + dt_bias) # g ≤ 0 (log-decay). A_log = log(A), A ~ U(0,16) +β = sigmoid(b) # allow_neg_eigval=False (decision); if True, β *= 2 +l2norm: q = q / sqrt(Σ q² + 1e-6) ; k = k / sqrt(Σ k² + 1e-6) # eps INSIDE the sqrt; fp32; per token,head +then: q *= scale (scale = 128**-0.5, q only; k NOT scaled) # AFTER l2norm — folding before cancels +decay applied per step: S *= exp(g) # RAW g, never cumsum'd in the recurrent path +``` +Reference-fixture init (for tests): `A ~ U(0,16)`, `A_log = log(A)`; `dt = exp(U(log 1e-3, log 1e-1))` clamped `≥1e-4`; `dt_bias = inv_softplus(dt)`. + +--- + +## 3. Verify V1 — recurrent multi-step + per-token intermediates (the shipping per-token path) + +**Prerequisite chain (hard gate, decode-spec §12):** build (1) the core `gs=1` decode kernel, validated; then (2) paging + no-commit + per-token-intermediate writes; then (3) optional in-kernel gating. The two unproven primitives blocking (1) are §11.A (`M=1` gemm) and §11.B (single-role `T.serial(L)` with runtime `L`) — prototype both **first**. + +**Per-token intermediates (req #6) — the free V1 win (GROUNDED).** The serial loop already holds the full post-update fp32 `S` in `h_fragment` at the end of every token (step ordering: decay → stage → `kS` → `v_new`(fp32) → rank-1 → **write state** → `o`). The write: +``` +if store_intermediate: + with T.If(slot >= 0): # gate on the POOL slot mask + with T.Then(): + T.copy(h_fragment, ibuf[cache_slot, t, bh, ]) # fp32 → bf16 +``` +`cache_slot = intermediate_state_indices[bb]` is the **destination** index (decoupled from `state_indices`); it is dense (`arange`, never `-1`), so the write is gated by the **pool-slot mask `slot = state_indices[bb] >= 0`** — the real-request mask — reusing A1's `slot` var. (Guarding on `cache_slot >= 0` would never fire and would write garbage into padded ibuf rows; harmless only because the scheduler reads `ibuf[slot, :k_accepted]` for real requests, but we skip those writes anyway.) `ibuf = intermediate_states_buffer [num_slots+1, cache_steps=12, H, V, K]` (V-major, single-layer slice — see §6). This is `fused_fwd.py:471-481` with `batch_idx → cache_slot` (paged) and `chunk_start_idx+i_s → token t`. Cost: +1 bf16 `[128,block_DV]` store/token; ~`12·32KB = 384 KB/head` of writes dominate (~92% of state traffic) — **inherent to req #6, the workload, not overhead.** + +**No-commit (req #7) — free (GROUNDED).** The spine's only pool write is the optional post-loop scatter (A1), gated by `not disable_state_update`. Verify passes `disable_state_update=True` and that `T.copy` is dead-code-eliminated. The scheduler later copies `ibuf[cache_slot, k_accepted]` into the live pool slot. + +**Gating** host-side primary (A5). Decay `exp2(g·1.442695)` per step (raw g). + +**`D=12` caveat.** Chunk-12 verify needs `D=12`, above the decode spine's `1..8` design point — re-derive `nreg` (§9 Gate 3) and the `M=1` gemm behavior at `D=12` with a compiled measurement. + +**Varlen prologue (real kernel change).** SGLang passes `cu_seqlens = query_start_loc [N+1]` with `B=1` flattened; the CTA→request map (`bb` from a flattened token layout) is a **different prologue** than the dense `bb = bbh//H` (`fused_fwd.py:92`) — derive `bb` / per-request token ranges device-side (`kkt_solve.py:245-247`), no `.item()`. Re-verify the capture-safe index math. + +**Perf (honest, conditional).** At `T=12`, server `N·H` saturating: bandwidth-bound on per-token state writes (bf16: 32 KB gather + 12·32 KB writes ≈ 416 KB/head; FLOPs hidden). V1 **beats** V2 here because the chunk kernel emits only 64-granular state (`NT=⌈12/64⌉=1`) so V2 must re-pay the same 12-step serial scan in a second pass plus chunk setup. Single-request (`N=1,H=16`): latency-bound, mitigated by CUDA-graph replay. Report per regime and per pool-dtype (falsifiable). + +--- + +## 4. Verify V2 — chunk-fused `o` + honest per-token-state extraction + +Built to benchmark; ships only if it wins total latency (§5). + +**What the chunk kernel actually gives.** The WY/UT chunk algorithm yields state **only at chunk boundaries** (`ref_gdr.py` appends `last_state` once per chunk; `fused_fwd.py` writes state at the entry copy `:190` and the single `transpose_A` gemm `:204`). For `T=12`, `NT=⌈12/64⌉=1` ⇒ a **single** (initial) state, **not** per-token. Per-token states are **not** a free byproduct. + +**The honest extraction mechanism + cost (the central V2 reality).** The post-loop extraction **cannot** reuse `vn_shared` as `v_new`: `vn_shared` holds `V' = g_rev·(Ag@W)` (`fused_fwd.py:283-286`) — the decay-corrected, A-projected **chunk** operand, **not** the per-token `v_new = β·(v − k·S_running)`. The true `v_new` depends on `w·S_running` (`w = A·k_beta`, WY-decoded), never materialized per-token. So extraction is a **genuinely new recurrent inner loop**: carry raw `v, k, β, g` (and per-token decay ratios `exp2(g_cumsum_t − g_cumsum_{t-1})`, also not resident) into a dedicated scan that, seeded from `S_entry`, recomputes `v_new_t` and runs 12 serial rank-1 updates, writing `S` after each token. **This scan IS V1's serial critical path, re-paid.** It needs an extra fp32 `S_scratch[DK,block_DV]` in the most register-pressured warpgroup (`CONSUMER_S_NREG=160`) → likely a **hard spill**, needing a dropped `nreg` or a dedicated scan warpgroup; and it must run **after** the `o` pipeline (`vn_shared` single-buffered, overwritten), serializing and killing the `o`/scan overlap that was V2's only hope. + +**Cost summary at `T=12`.** V2 = (faster intra-chunk `o` via wgmma — the 1.8–2.16× probe, real, preserved) + (`kkt_solve` 64×64 inverse, **52/64 rows wasted** at T=12, an extra capturable launch) + (chunk cumsum) + (gate fusion across **three** surfaces: cumsum input, `kkt_solve` β load `:99`, `fused_fwd` `Ag` `:332`) + (**the same 12-step serial scan as V1**, now burdened with re-deriving `v_new`). State I/O is a **wash** vs V1 (same 12·32 KB writes). So the per-token-state requirement **erases V2's structural advantage**: V2 wins on `o`-latency only. + +**Layout/gating fixes carried from A:** `state_v_first=True` for pool **and** ibuf; gate `g = -exp(A_log)·softplus(a+dt_bias)` **base-e** (any `exp2(A_log·1.442695)` form is **wrong** — the `1.442695/exp2` factor belongs only to the per-step decay application, `fused_fwd.py:236`); host-side l2norm for the MVP (note: this means the V2 entry is **not** fully in-kernel-l2norm capture-safe — scope it explicitly). The verify wrapper must **bypass the varlen branch** of `fused_gdr_fwd` entirely (`:569` calls `prepare_chunk_offsets`); for `T≤64` single-chunk, `chunk_offsets = arange(N+1)` / `chunk_indices = [[n,0]]` are static, precomputed pre-capture. + +--- + +## 5. Benchmark-to-decide (V1 vs V2 at T=12) + +Winner decided by **total verify-step latency (`o` + per-token states)**, not `o` alone. + +- **Shapes:** `H=32, Hg=16, K=V=128, T=12`, single chunk. Sweep `N (=requests) ∈ {1, 8, 32, 64, 128, 256}` to cross the knee (`TARGET=⌊132·0.7⌋=92` on H100; `N·H` crosses 92 around `N=3`). `q,k` bf16 `[1, N·12, 16, 128]`; `v,o` bf16 `[1, N·12, 32, 128]`; `a,b` bf16 `[1, N·12, 32]`; `A_log,dt_bias` `[32]`; pool `[num_slots, 32, 128, 128]` V-major; `intermediate_states_buffer [num_cache_slots, 12, 32, 128, 128]`; indices int32. Run **both** bf16 and fp32 pools (the dominant per-token-write term doubles under fp32 and can flip the result). +- **Metrics:** (1) total wall time (median of 1000 iters, `flash_qla.utils.profile`); (2) achieved HBM GB/s vs ~3.35 TB/s, split state-read / per-token-write / io; (3) **V2 only:** separately time chunk-`o` vs the extraction scan and **measure overlap vs serialization**; (4) V2 `kkt_solve` launch+exec overhead (wasted 52/64 rows). +- **Baselines:** FLA `fused_sigmoid_gating_delta_rule_update` (the SGLang Triton verify kernel; natively emits per-token intermediates + supports `disable_state_update`) — both engines match its `o` and per-token states to 0.02 rel and target ≥ its SM90 perf. Also baseline V1 vs "12 separate `D=1` decode calls" (fusion win) and vs per-token-states-OFF (isolate the req #6 cost). +- **Decision rule:** ship **V1** unless V2 shows **≥15% total-latency win** at the deployment's dominant `N·H` **AND** its extraction scan **provably overlaps** the `o` pipeline **AND** passes the per-head-distinct-gate bit-identity test. **Pre-req to even running:** `M=1` gemm must compile (else fold M-pad-to-16 overhead into V1's measured time — do not assume hidden). + +--- + +## 6. SGLang integration contract (verify entry, SM90, T=12 linear chain) + +All tensors caller-preallocated; `B=1` outer dim, `N` requests flattened. + +| tensor | shape | dtype | notes | +|---|---|---|---| +| `q`, `k` | `[1, N·T, 16, 128]` | bf16 | Hg=16 | +| `v`, `o` (returned) | `[1, N·T, 32, 128]` | bf16 | H=32; `o` written per token | +| `a`, `b` | `[1, N·T, 32]` | bf16 | raw gate inputs (in-kernel path); host path passes pre-activated `g`,`β` | +| `A_log`, `dt_bias` | `[32]` | fp32 | unused in host-gating baseline | +| `pool` (`ssm_states`) | `[num_slots, 32, 128, 128]` | SSM dtype (bf16) | **V-major** (`state_v_first=True`), K stride-1 | +| `state_indices` (`cache_indices`) | `[N]` | int32 | slot/request; `<0 ⇒ skip` gather AND scatter | +| `intermediate_states_buffer` | `[num_slots+1, 12, 32, 128, 128]` | SSM dtype | single-layer slice of `[num_layers, num_slots+1, draft_token_num=12, HV, V, K]`; **same V-major layout as pool**; `num_cache_slots = num_slots+1`; `cache_steps = 12` exact (non-adaptive) | +| `intermediate_state_indices` | `[N]` | int32 | destination slot into ibuf; **dense `arange`, never `-1`** — the ibuf write is gated by the **pool-slot mask** (`state_indices[b] ≥ 0`), not this index | +| `cu_seqlens` (`query_start_loc`) | `[N+1]` | int32 | ragged; in-kernel load only | + +**Capture rules:** zero host sync (never `prepare_chunk_offsets`/`prepare_chunk_indices`; for V2 precompute trivial single-chunk offsets pre-capture and bypass the varlen branch); zero `torch.empty`; static shapes (`block_DV` from the MPC ladder, `T=12` fixed); the PyTorch gating/l2norm runs **inside** capture (capture-safe elementwise — A5), not in the TileLang kernel and **not** lifted out of the graph. + +**No-commit:** `disable_state_update=True` ⇒ final pool scatter dead-code-eliminated; only `o` + ibuf produced. Decode follow-on sets it False and commits. + +**Validation gate:** per-head-distinct-gate bit-identity vs FLA `fused_sigmoid_gating_delta_rule_update` (`o` + per-token states) — `K==V==128` makes a wrong `state_v_first` silent in equal-dim tests. + +--- + +## 7. API signatures + +**Graph-safe low-level entry (V1):** +```python +fused_recurrent_gdr_verify_fwd( + q, k, v, # bf16 (raw gate path: + a, b bf16 [1,N·T,32]; A_log, dt_bias fp32 [32]) + pool, # [num_slots, 32, 128, 128] V-major; SSM dtype + state_indices, # int32 [N]; -1 (PAD_SLOT_ID) => skip; slot 0 valid + o, # CALLER-PREALLOC bf16 [1, N·T, 32, 128] + intermediate_states_buffer, # CALLER-PREALLOC [num_slots+1, 12, 32, 128, 128] V-major (single-layer slice) + intermediate_state_indices, # int32 [N]; dense arange (never -1); ibuf write gated by state_indices>=0 + cu_seqlens, # int32 [N+1]; per-CTA loop bound; in-kernel load only + scale=None, # default 128**-0.5 + disable_state_update=True, # NO-COMMIT (verify) + use_qk_l2norm_in_kernel=False, # False in primary path (host l2norm) + state_v_first=True, # SGLang interop default + allow_neg_eigval=False, # decision; flag exposed +) -> o # pool written in-place ONLY if not disable_state_update +``` +**JIT factory keys:** `tilelang_fused_recurrent_gdr_verify(H, Hg, DK=128, DV=128, scale, *_dtype, use_initial_state, disable_state_update, store_intermediate, is_varlen, use_qk_l2norm_in_kernel, state_v_first, allow_neg_eigval, fuse_gating, head_batch=False, group_size=1, block_DV=128, threads=256)`. `N`, `num_tokens` are `T.dynamic`; `D` fixed per capture. Kernel name ends `*_kernel_kernel` (CLAUDE.md). + +**High-level wrapper** (drop-in for SGLang's DFlash `target_verify`; capture-safe PyTorch gating + dispatch, no allocation in the captured path): +```python +recurrent_gated_delta_rule_verify( + A_log, a, dt_bias, q, k, v, b, ssm_states, cache_indices, query_start_loc, + intermediate_states_buffer, intermediate_state_indices, cache_steps, + retrieve_parent_token=None, # accepted-and-IGNORED: DFlash is width-1 (TOPK=1); tree tensors are zeros + scale=None, use_qk_l2norm_in_kernel=True, disable_state_update=True) +# asserts K==V==128, H%Hg==0, q.dtype!=fp32; derives state_v_first from ssm_states.stride()/shape; +# PyTorch l2norm+gating (primary, capture-safe — inside the graph, not the TileLang kernel); +# host block_DV via MPC ladder; no autograd. Confirm arg order vs DFlash's actual call (§10.6). +``` +**V2 entry (only if benchmark-selected):** `tilelang_fused_chunk_gdr_verify_fwd(..., commit_final_state=False, store_intermediate=True, state_v_first=True, fuse_gate, ...)` with an explicit per-token extraction scan sub-kernel carrying raw `v,k,β,g`; static precomputed `chunk_offsets=arange(N+1)`/`chunk_indices=[[n,0]]`; bypasses `prepare_chunk_offsets`. + +--- + +## 8. Test plan + +- **Reference:** FLA `fused_sigmoid_gating_delta_rule_update` (matches the SGLang verify kernel; emits per-token states + supports `disable_state_update`). Both engines match `o` + per-token states to **0.02 rel**. +- **Bit-identity (hard layout gate):** per-head-**distinct**-gate test (each head a different `g`,`β`) vs the reference for `state_v_first=True` — equal-dim (`K==V==128`) tests are numerically blind to a layout transpose, so distinct-gate is mandatory. +- **Sweeps:** `T∈{1,3,8,12}`; ragged `cu_seqlens` (mixed accepted lengths e.g. `[1,5,12]`, compare only `o[:,:L_b]` and `ibuf[slot,:L_b]`); GQA `Hg=16/Hv=32` (verify head `h` reads k-head `h//2`); `g=0` SWA heads; bf16 **and** fp32 pool; `state_indices`/`intermediate_state_indices` with `<0` skip slots (assert untouched). +- **Negative controls:** cumsum'd g must break; pre-update read must break; `mod` GQA must break; a `[K,V]` (wrong-major) pool must break the distinct-gate test; garbage in skipped (`<0`) slots must leave committed outputs bit-unchanged. +- **Graph-safety test:** capture the verify call under `torch.cuda.graph`, replay, assert no allocation/sync errors and identical output. +- **Feasibility smoke tests:** §10 gates (especially `M=1` gemm, `T.serial(L)`, `state_v_first` transpose store) before full sign-off. + +--- + +## 9. Open feasibility gates (prototype on the Hopper box, ordered by gating power) + +1. **`M=1` gemm_v1** (decode-spec §11.A) — **root gate**. Every repo gemm is `M=64`. Compile-test `M=1` first; fallback M-pad-to-16 (quantify the 64× wgmma-row waste; at T=12 BW-bound it may hide — measure). +2. **single-role `T.serial(L)` with runtime per-CTA `L`** (§11.B) — grounded only in warp-spec `prepare_h.py:166`; confirm single-role. Static fallback: `T.serial(T)` + `if t= 0` (`fused_sigmoid_gating_recurrent.py:98/131/204/230`); slot 0 is valid (not vLLM's 0/NULL). Only `state_indices` is `-1`-padded; `intermediate_state_indices` is dense `arange` (gate ibuf writes on the pool mask — A1/§3). +6. **DFlash = linear width-1** (`SPECULATIVE_EAGLE_TOPK=1`; `dflash_worker_v2` passes empty `topk_p`/`topk_index`; `retrieve_*` tree tensors are zeros, populated only when `topk>1`, `hybrid_linear_attn_backend.py:487-489`). The wrapper **accepts-and-ignores** `retrieve_parent_token` (the kernel still receives it, trivial for width-1). + +**RESIDUAL (confirm before final ship):** +5. **FlashInfer numerics** — for SM90/Hopper the Triton fp32-accum path runs; confirm whether bit-matching the FlashInfer SM100 bf16-state adapter is a requirement. +6b. **DFlash entry parity** — reconcile `recurrent_gated_delta_rule_verify`'s exact arg order / buffer names against DFlash's live GDN verify call so it's drop-in. + +--- + +## 11. Implementation order + +1. **Prove the gates** §9.1 (`M=1` gemm) and §9.2 (`T.serial(L)`) on the Hopper box — these gate everything. +2. **Core `gs=1` decode kernel** (decode-spec §3–§6) + validate. +3. **Infra A** (A1 paging, A2 bf16+fp32-accum, A3 graph-safe entry, A4 `state_v_first`, A5 host gating) as keys/hooks on the kernel + a graph-safe wrapper. +4. **Verify V1** (per-token intermediate writes + no-commit + varlen `cu_seqlens` prologue + `D=12`), validated by the bit-identity + graph-safety tests vs FLA. +5. **Verify V2** (chunk-`o` + the honest extraction scan) — only after §9.7 budget check. +6. **Benchmark-to-decide** (§5) → ship the winner; the other stays documented. +7. **In-kernel gating** fast-follow if §9.4 passes. +8. **(Optional) standalone decode (C)** — later. (Tree-structured propagation is **out of scope** — not needed; the linear-chain design doesn't preclude adding it later.) diff --git a/flash_qla/__init__.py b/flash_qla/__init__.py index 640f6845..cf9061bb 100644 --- a/flash_qla/__init__.py +++ b/flash_qla/__init__.py @@ -8,9 +8,19 @@ chunk_gated_delta_rule_bwd, chunk_gated_delta_rule, ) +from flash_qla.ops.gated_delta_rule.fused_recurrent import ( + fused_recurrent_gdr_fwd, + recurrent_gated_delta_rule, + fused_recurrent_gdr_verify_fwd, + recurrent_gated_delta_rule_verify, +) __all__ = [ "chunk_gated_delta_rule_fwd", "chunk_gated_delta_rule_bwd", "chunk_gated_delta_rule", + "fused_recurrent_gdr_fwd", + "recurrent_gated_delta_rule", + "fused_recurrent_gdr_verify_fwd", + "recurrent_gated_delta_rule_verify", ] diff --git a/flash_qla/ops/gated_delta_rule/__init__.py b/flash_qla/ops/gated_delta_rule/__init__.py index ea07eebd..91a2aa52 100644 --- a/flash_qla/ops/gated_delta_rule/__init__.py +++ b/flash_qla/ops/gated_delta_rule/__init__.py @@ -2,6 +2,14 @@ # Licensed under The MIT License [see LICENSE for details] from .chunk import chunk_gated_delta_rule +from .fused_recurrent import ( + recurrent_gated_delta_rule, + recurrent_gated_delta_rule_verify, +) -__all__ = ["chunk_gated_delta_rule"] +__all__ = [ + "chunk_gated_delta_rule", + "recurrent_gated_delta_rule", + "recurrent_gated_delta_rule_verify", +] diff --git a/flash_qla/ops/gated_delta_rule/fused_recurrent/__init__.py b/flash_qla/ops/gated_delta_rule/fused_recurrent/__init__.py new file mode 100644 index 00000000..0dd5a981 --- /dev/null +++ b/flash_qla/ops/gated_delta_rule/fused_recurrent/__init__.py @@ -0,0 +1,163 @@ +# Copyright (c) 2026 The Qwen team, Alibaba Group. +# Licensed under The MIT License [see LICENSE for details] +import torch +import torch.nn.functional as F +import tilelang + +from flash_qla.utils import l2norm + +if tilelang.contrib.nvcc.get_target_compute_version() == "9.0": + from .hopper import fused_recurrent_gdr_fwd # noqa: F401 + from .hopper.fused_recurrent_verify import ( # noqa: F401 + fused_recurrent_gdr_verify_fwd, + fused_recurrent_gdr_verify_gated_fwd, + fused_recurrent_gdr_verify_prepass, + get_prepass_scratch, + should_use_prepass, + ) +else: + raise ValueError("FlashQLA now support sm90 only.") + +__all__ = [ + "fused_recurrent_gdr_fwd", + "recurrent_gated_delta_rule", + "fused_recurrent_gdr_verify_fwd", + "fused_recurrent_gdr_verify_gated_fwd", + "recurrent_gated_delta_rule_verify", +] + + +def recurrent_gated_delta_rule( + q, + k, + v, + g, + beta, + scale=None, + initial_state=None, + output_final_state=True, + use_qk_l2norm_in_kernel=False, + seqlens=None, + head_first=False, + head_batch=None, +): + assert q.dtype == k.dtype == v.dtype and q.dtype != torch.float32 + assert not head_first, "head_first=True is not supported." + assert v.shape[2] % k.shape[2] == 0 and q.shape[-1] == v.shape[-1] == 128 + if scale is None: + scale = k.shape[-1] ** -0.5 + if use_qk_l2norm_in_kernel: + q = l2norm(q) + k = l2norm(k) + o, final_state = fused_recurrent_gdr_fwd( + q, + k, + v, + g, + beta, + scale=scale, + initial_state=initial_state, + output_final_state=output_final_state, + seqlens=seqlens, + head_batch=head_batch, + ) + return o.to(q.dtype), final_state + + +def gdn_sigmoid_gate(A_log, a, dt_bias, b, allow_neg_eigval=False): + """Host-side GDN gating (the sigmoid_gating family): g = -exp(A_log)*softplus(a+dt_bias), + beta = sigmoid(b) (x2 if allow_neg_eigval). A_log,dt_bias:[H]; a,b:[...,H]. Returns fp32.""" + g = -torch.exp(A_log.float())[(None,) * (a.dim() - 1)] * F.softplus( + a.float() + dt_bias.float()[(None,) * (a.dim() - 1)] + ) + beta = torch.sigmoid(b.float()) + if allow_neg_eigval: + beta = beta * 2 + return g, beta + + +def recurrent_gated_delta_rule_verify( + A_log, + a, + dt_bias, + q, + k, + v, + b, + ssm_states, + cache_indices, + query_start_loc, + intermediate_states_buffer, + intermediate_state_indices, + cache_steps=None, + o=None, + scale=None, + use_qk_l2norm_in_kernel=True, + disable_state_update=True, + allow_neg_eigval=False, + fuse_gating=False, + prepass=None, # H1 dedup gating+l2norm pre-pass: None=auto (regime-gated), True/False=force + retrieve_parent_token=None, # accepted-and-IGNORED (DFlash width-1; tree path not built) +): + """High-level SGLang DFlash verify entry. q,k:[1,T,Hk,128] v:[1,T,Hv,128]; a,b:[1,T,Hv]; + A_log,dt_bias:[Hv]; ssm_states pool V-major [num_slots,Hv,128,128]. + + CUDA-graph note: for capture, use ``fuse_gating=True`` (computes g/beta + qk-l2norm INSIDE + the kernel from raw a,b,A_log,dt_bias -- no PyTorch gating/l2norm, no allocation when ``o`` + is provided -> fully capture-safe). The default ``fuse_gating=False`` path computes g/beta + + qk-l2norm in PyTorch (l2norm is ``@torch.compile``'d) and allocates them; run it OUTSIDE + capture or prefer ``fuse_gating=True`` inside it. + + H1 (``fuse_gating=True``): in the large-batch / multi-draft-token regime the gating + qk-l2norm + are split into a tiny dedup PRE-PASS kernel (computed ONCE per (token,K-head) instead of once + per (token,V-head,V-tile) in the hot loop) feeding the host-gated main kernel -- measured + ~12-15% faster at T=12 large batch. Regime-gated by ``should_use_prepass`` (the single + in-kernel-gated kernel A is kept for latency-bound single-request, where a 2nd launch is a net + loss). The prepass scratch is a persistent, never-evicting cache -> capture-safe after warmup. + ``prepass`` forces the choice (None=auto, True=prepass+main, False=single gated kernel A).""" + assert q.dtype == k.dtype == v.dtype and q.dtype != torch.float32 + assert q.shape[-1] == v.shape[-1] == 128 and v.shape[2] % k.shape[2] == 0 + if cache_steps is not None: # static shape check (capture-safe; no value read) + assert intermediate_states_buffer.shape[1] >= cache_steps, ( + f"intermediate_states_buffer cache-steps dim {intermediate_states_buffer.shape[1]} " + f"< cache_steps {cache_steps}" + ) + scale = scale if scale is not None else q.shape[-1] ** -0.5 + if o is None: + o = torch.empty(1, q.shape[1], v.shape[2], v.shape[-1], device=q.device, dtype=q.dtype) + + if fuse_gating: + N, total_tokens, Hk = cache_indices.shape[0], q.shape[1], k.shape[2] + H = v.shape[2] + use_prepass = should_use_prepass(N, H, total_tokens) if prepass is None else prepass + if use_prepass: + # H1: dedup gating + qk-l2norm into a run-once pre-pass, then the host-gated main + # kernel (no in-loop transcendentals). Persistent scratch -> capture-safe after warmup. + q_n, k_n, g_pp, beta_pp = get_prepass_scratch(total_tokens, Hk, H, q.device, q.dtype) + fused_recurrent_gdr_verify_prepass( + q, k, a, b, A_log, dt_bias, q_n, k_n, g_pp, beta_pp, + allow_neg_eigval=allow_neg_eigval, + ) + fused_recurrent_gdr_verify_fwd( + q_n, k_n, v, g_pp, beta_pp, ssm_states, cache_indices, query_start_loc, + intermediate_states_buffer, intermediate_state_indices, o, + scale=scale, disable_state_update=disable_state_update, + ) + return o + fused_recurrent_gdr_verify_gated_fwd( + q, k, v, a, b, A_log, dt_bias, ssm_states, cache_indices, query_start_loc, + intermediate_states_buffer, intermediate_state_indices, o, + scale=scale, disable_state_update=disable_state_update, allow_neg_eigval=allow_neg_eigval, + ) + return o + + if use_qk_l2norm_in_kernel: + q = l2norm(q) + k = l2norm(k) + g, beta = gdn_sigmoid_gate(A_log, a, dt_bias, b, allow_neg_eigval) + fused_recurrent_gdr_verify_fwd( + q, k, v, g, beta, ssm_states, cache_indices, query_start_loc, + intermediate_states_buffer, intermediate_state_indices, o, + scale=scale, disable_state_update=disable_state_update, + ) + return o diff --git a/flash_qla/ops/gated_delta_rule/fused_recurrent/hopper/__init__.py b/flash_qla/ops/gated_delta_rule/fused_recurrent/hopper/__init__.py new file mode 100644 index 00000000..a8fd7cae --- /dev/null +++ b/flash_qla/ops/gated_delta_rule/fused_recurrent/hopper/__init__.py @@ -0,0 +1,5 @@ +# Copyright (c) 2026 The Qwen team, Alibaba Group. +# Licensed under The MIT License [see LICENSE for details] +from .fused_recurrent_fwd import fused_recurrent_gdr_fwd + +__all__ = ["fused_recurrent_gdr_fwd"] diff --git a/flash_qla/ops/gated_delta_rule/fused_recurrent/hopper/fused_recurrent_fwd.py b/flash_qla/ops/gated_delta_rule/fused_recurrent/hopper/fused_recurrent_fwd.py new file mode 100644 index 00000000..be60039d --- /dev/null +++ b/flash_qla/ops/gated_delta_rule/fused_recurrent/hopper/fused_recurrent_fwd.py @@ -0,0 +1,357 @@ +# Copyright (c) 2026 The Qwen team, Alibaba Group. +# Licensed under The MIT License [see LICENSE for details] +import torch +import tilelang +import tilelang.language as T + +MULTI_PROCESSOR_COUNT = torch.cuda.get_device_properties().multi_processor_count +TARGET_NUM_CTAS = int(MULTI_PROCESSOR_COUNT * 0.7) + + +@tilelang.jit(pass_configs={tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True}) +def tilelang_fused_recurrent_gdr_fwd( + H, + Hg, + DK, + DV, + scale, + accum_dtype, + qkva_dtype, + g_dtype, + b_dtype, + h0_dtype, + ht_dtype, + o_dtype, + seqlen_dtype, + use_initial_state, + store_final_state, + has_seqlens, + block_DV=128, + threads=128, +): + batch_size = T.dynamic("batch_size") + num_tokens = T.dynamic("num_tokens") # = q_len (D); D fixed per call + q_shape = (batch_size, num_tokens, Hg, DK) + v_shape = (batch_size, num_tokens, H, DV) + g_shape = (batch_size, num_tokens, H) + s_shape = (batch_size, H, DK, DV) + n_vt = (DV + block_DV - 1) // block_DV + + @T.prim_func + def kernel( + q: T.Tensor(q_shape, qkva_dtype), + k: T.Tensor(q_shape, qkva_dtype), + v: T.Tensor(v_shape, qkva_dtype), + g: T.Tensor(g_shape, g_dtype), + b: T.Tensor(g_shape, b_dtype), + h0: T.Tensor(s_shape, h0_dtype), + seqlens: T.Tensor([batch_size], seqlen_dtype), + o: T.Tensor(v_shape, o_dtype), + ht: T.Tensor(s_shape, ht_dtype), + ): + # Gemm-free, memory-bound decode recurrence. One CTA owns (sequence bb, V-head bh, + # V-column-tile bv). State S is kept [block_DV, DK] (V-major rows) in an fp32 fragment; + # the two GEMVs are reductions over the last (DK) dim and the rank-1 is a T.Parallel + # outer product. No tensor-core gemm -> no M-padding, no warp-partition constraint. + with T.Kernel(n_vt * batch_size * H, threads=threads) as (bbhv,): + bbh = bbhv // n_vt + bv = bbhv % n_vt + bb = bbh // H + bh = bbh % H + bhg = bh // (H // Hg) + v0 = bv * block_DV + + L = T.alloc_var("int32") + L = seqlens[bb] if has_seqlens else num_tokens + + S = T.alloc_fragment((block_DV, DK), accum_dtype) # state [V-tile, K], fp32 + prod = T.alloc_fragment((block_DV, DK), accum_dtype) + q_s = T.alloc_shared((1, DK), qkva_dtype) # 2-D staging (1-D slice copies fail layout) + k_s = T.alloc_shared((1, DK), qkva_dtype) + v_s = T.alloc_shared((1, block_DV), qkva_dtype) + o_sh = T.alloc_shared((1, block_DV), o_dtype) + kS = T.alloc_fragment((block_DV,), accum_dtype) + oo = T.alloc_fragment((block_DV,), accum_dtype) + vnew = T.alloc_fragment((block_DV,), accum_dtype) + decay = T.alloc_fragment((1,), accum_dtype) + bt = T.alloc_fragment((1,), accum_dtype) + + if use_initial_state: + # h0 is K-major [B,H,DK,DV]; load transposed into S[v, dk] = h0[bb,bh,dk,v0+v] + for j_v, j_k in T.Parallel(block_DV, DK): + S[j_v, j_k] = h0[bb, bh, j_k, v0 + j_v] + else: + T.clear(S) + + for t in T.serial(L): + T.copy(q[bb, t : t + 1, bhg, 0:DK], q_s) + T.copy(k[bb, t : t + 1, bhg, 0:DK], k_s) + T.copy(v[bb, t : t + 1, bh, v0 : v0 + block_DV], v_s) + decay[0] = T.exp2(g[bb, t, bh] * 1.442695) # raw g, exp2 + bt[0] = b[bb, t, bh] + for j_v, j_k in T.Parallel(block_DV, DK): + S[j_v, j_k] *= decay[0] + # kS[v] = sum_dk k[dk] * S[v, dk] + for j_v, j_k in T.Parallel(block_DV, DK): + prod[j_v, j_k] = k_s[0, j_k] * S[j_v, j_k] + T.reduce_sum(prod, kS, dim=1) + # v_new[v] = beta * (v[v] - kS[v]) + for j_v in T.Parallel(block_DV): + vnew[j_v] = bt[0] * (v_s[0, j_v] - kS[j_v]) + # rank-1: S[v, dk] += k[dk] * v_new[v] + for j_v, j_k in T.Parallel(block_DV, DK): + S[j_v, j_k] += k_s[0, j_k] * vnew[j_v] + # o[v] = scale * sum_dk q[dk] * S[v, dk] (post-update read) + for j_v, j_k in T.Parallel(block_DV, DK): + prod[j_v, j_k] = q_s[0, j_k] * S[j_v, j_k] + T.reduce_sum(prod, oo, dim=1) + for j_v in T.Parallel(block_DV): + o_sh[0, j_v] = oo[j_v] * scale + T.copy(o_sh, o[bb, t : t + 1, bh, v0 : v0 + block_DV]) + + if store_final_state: + # ht is K-major [B,H,DK,DV]; store transposed ht[bb,bh,dk,v0+v] = S[v, dk] + for j_v, j_k in T.Parallel(block_DV, DK): + ht[bb, bh, j_k, v0 + j_v] = S[j_v, j_k] + + return kernel + + +@tilelang.jit(pass_configs={tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True}) +def tilelang_fused_recurrent_gdr_fwd_hb( + H, + Hg, + DK, + DV, + scale, + accum_dtype, + qkva_dtype, + g_dtype, + b_dtype, + h0_dtype, + ht_dtype, + o_dtype, + seqlen_dtype, + use_initial_state, + store_final_state, + has_seqlens, + block_DV=64, + threads=256, +): + """Head-batched GQA specialization of the gemm-free decode kernel. One CTA owns + (sequence bb, K/Q head-group hg, V-col-tile bv) and processes ALL grp = H//Hg V-heads + h = hg*grp + i that share K/Q head hg, loading q/k ONCE (dedup). State is row-stacked + S[grp*block_DV, DK]: row gv -> head-band i = gv//block_DV, v-row jv = gv%block_DV. + Layout-inference rule learned on H100: every [M] fragment (S/prod/decay_f/.../oo) MUST be + accessed over the FULL Parallel(M,...) range -- a partial Parallel(block_DV) write at offset + i*block_DV makes the fragment's affine map non-invertible (TVM InverseAffineIterMap check + fails). So ALL per-head divergence is routed through GLOBAL reads/writes indexed by the + derived head hg*grp + gv//block_DV and channel v0 + gv%block_DV (global needs no fragment + inversion); only q/k stage through shared (loaded once for the group). + threads = grp*128 keeps per-thread register footprint identical to the per-head kernel.""" + grp = H // Hg + M = grp * block_DV + batch_size = T.dynamic("batch_size") + num_tokens = T.dynamic("num_tokens") + q_shape = (batch_size, num_tokens, Hg, DK) + v_shape = (batch_size, num_tokens, H, DV) + g_shape = (batch_size, num_tokens, H) + s_shape = (batch_size, H, DK, DV) + n_vt = (DV + block_DV - 1) // block_DV + + @T.prim_func + def kernel( + q: T.Tensor(q_shape, qkva_dtype), + k: T.Tensor(q_shape, qkva_dtype), + v: T.Tensor(v_shape, qkva_dtype), + g: T.Tensor(g_shape, g_dtype), + b: T.Tensor(g_shape, b_dtype), + h0: T.Tensor(s_shape, h0_dtype), + seqlens: T.Tensor([batch_size], seqlen_dtype), + o: T.Tensor(v_shape, o_dtype), + ht: T.Tensor(s_shape, ht_dtype), + ): + with T.Kernel(n_vt * batch_size * Hg, threads=threads) as (bbhv,): + bbh = bbhv // n_vt # flattened (bb, hg) + bv = bbhv % n_vt + bb = bbh // Hg + hg = bbh % Hg + v0 = bv * block_DV + + L = T.alloc_var("int32") + L = seqlens[bb] if has_seqlens else num_tokens + + S = T.alloc_fragment((M, DK), accum_dtype) # row-stacked state [grp*V-tile, K] + prod = T.alloc_fragment((M, DK), accum_dtype) + q_s = T.alloc_shared((1, DK), qkva_dtype) # shared across the grp heads + k_s = T.alloc_shared((1, DK), qkva_dtype) + decay_f = T.alloc_fragment((M,), accum_dtype) # row-aligned bands (full-M access) + b_f = T.alloc_fragment((M,), accum_dtype) + kS = T.alloc_fragment((M,), accum_dtype) + oo = T.alloc_fragment((M,), accum_dtype) + vnew = T.alloc_fragment((M,), accum_dtype) + + if use_initial_state: + # full-M gather; K-major h0[B,H,DK,DV] -> S[gv, K] with derived head/channel + for gv, j_k in T.Parallel(M, DK): + S[gv, j_k] = h0[bb, hg * grp + gv // block_DV, j_k, v0 + gv % block_DV] + else: + T.clear(S) + + for t in T.serial(L): + T.copy(q[bb, t : t + 1, hg, 0:DK], q_s) # loaded ONCE for the whole group + T.copy(k[bb, t : t + 1, hg, 0:DK], k_s) + # per-band decay/beta into FULL-M fragments (global g/beta at derived head). A + # shared [grp] band read by gv//block_DV does NOT lower in TileLang (the Parallel + # layout inferencer rejects the shared derived-index read), so we materialize the + # full M; the redundant exp2 is cheap vs the FMA floor on this memory-bound kernel. + for gv in T.Parallel(M): + decay_f[gv] = T.exp2(g[bb, t, hg * grp + gv // block_DV] * 1.442695) + b_f[gv] = b[bb, t, hg * grp + gv // block_DV] + for gv, j_k in T.Parallel(M, DK): + S[gv, j_k] *= decay_f[gv] + for gv, j_k in T.Parallel(M, DK): + prod[gv, j_k] = k_s[0, j_k] * S[gv, j_k] + T.reduce_sum(prod, kS, dim=1) + # v read straight from global (derived head/channel); no shared staging + for gv in T.Parallel(M): + vnew[gv] = b_f[gv] * ( + v[bb, t, hg * grp + gv // block_DV, v0 + gv % block_DV] - kS[gv] + ) + for gv, j_k in T.Parallel(M, DK): + S[gv, j_k] += k_s[0, j_k] * vnew[gv] + for gv, j_k in T.Parallel(M, DK): + prod[gv, j_k] = q_s[0, j_k] * S[gv, j_k] + T.reduce_sum(prod, oo, dim=1) + for gv in T.Parallel(M): # o written straight to global (derived head/channel) + o[bb, t, hg * grp + gv // block_DV, v0 + gv % block_DV] = oo[gv] * scale + + if store_final_state: + for gv, j_k in T.Parallel(M, DK): + ht[bb, hg * grp + gv // block_DV, j_k, v0 + gv % block_DV] = S[gv, j_k] + + return kernel + + +def _fused_recurrent_gdr_fwd_hb( + q, k, v, g, beta, scale, initial_state, output_final_state, seqlens, grp +): + """Head-batched dispatch helper. Grid collapses to n_vt*B*Hg; block_DV keys off the + POST-collapse supply (B*Hg); threads = grp*128 (cap 512).""" + B, Tq, Hg, K = k.shape + _, _, H, V = v.shape + + # block_DV from the head-batched grid (B*Hg CTAs, not B*H): the collapse already cost + # CTAs, so the V-split must work harder to refill. threads scale with grp to hold the + # per-thread footprint constant (== the per-head kernel). + grid_base = B * Hg + block_DV = 64 if grid_base * 2 >= TARGET_NUM_CTAS else 32 + threads = min(512, grp * 128) + + use_initial_state = initial_state is not None + if initial_state is None: + initial_state = torch.empty((B, H, K, V), dtype=torch.float32, device=k.device) + final_state = torch.empty((B, H, K, V), dtype=torch.float32, device=k.device) + o = torch.empty_like(v) + + has_seqlens = seqlens is not None + if seqlens is None: + seqlens = torch.empty((B,), dtype=torch.int32, device=k.device) + + kern = tilelang_fused_recurrent_gdr_fwd_hb( + H, Hg, K, V, scale, + accum_dtype="float32", + qkva_dtype=q.dtype, + g_dtype=g.dtype, + b_dtype=beta.dtype, + h0_dtype=initial_state.dtype, + ht_dtype=final_state.dtype, + o_dtype=o.dtype, + seqlen_dtype=seqlens.dtype, + use_initial_state=use_initial_state, + store_final_state=output_final_state, + has_seqlens=has_seqlens, + block_DV=block_DV, + threads=threads, + ) + kern(q, k, v, g, beta, initial_state, seqlens, o, final_state) + return o, (final_state if output_final_state else None) + + +def fused_recurrent_gdr_fwd( + q, + k, + v, + g, + beta, + scale=None, + initial_state=None, + output_final_state=False, + seqlens=None, + head_batch=None, +): + B, Tq, Hg, K = k.shape + _, _, H, V = v.shape + assert K == V == 128 and H % Hg == 0 + scale = scale or K ** -0.5 + + # Head-batched GQA: one CTA processes all grp = H//Hg V-heads sharing a K/Q head, loading + # q/k once. Auto OFF here (the host-gated decode path saves only the q/k load ~ sub-1%, + # below noise); exposed as a forceable flag for benchmarking/tests. Restricted to grp in + # {2,4} (threads <= 512) -- grp>4 risks register spills / 1024-thread occupancy loss. + grp = H // Hg + if head_batch is None: + head_batch = False + if head_batch: + assert grp in (2, 4), f"head_batch supports grp=H//Hg in {{2,4}}, got grp={grp}" + return _fused_recurrent_gdr_fwd_hb( + q, k, v, g, beta, scale, initial_state, output_final_state, seqlens, grp + ) + + # Tile selection. When writing the final state (`output_final_state`, the decode default), the + # K-major transposed store `ht[bb,bh,jk,v0+jv] = S[jv,jk]` is the dominant cost: at block_DV<128 + # it is catastrophically uncoalesced (~0.4 TB/s), but at block_DV=128 (full-V tile, n_vt=1) the + # transpose coalesces -> measured ~2x faster at EVERY batch size (1.35x @ B=1 .. 3.0x @ B=8-16 + # .. 2.0x @ B=256; benchmark/probe_h2_blockdv_crossover.py, final_state bit-identical). So force + # block_DV=128 whenever we store the final state. (This inverts the verify kernel's V-major + # choice of 64 -- that store needs no transpose, so there occupancy wins; here the write does.) + # No final-state write: keep the occupancy ladder (64/32) -- block_DV is then perf-neutral. + grid_base = B * H + if output_final_state: + block_DV = 128 + else: + block_DV = 64 if grid_base * 2 >= TARGET_NUM_CTAS else 32 + + use_initial_state = initial_state is not None + if initial_state is None: + initial_state = torch.empty((B, H, K, V), dtype=torch.float32, device=k.device) + final_state = torch.empty((B, H, K, V), dtype=torch.float32, device=k.device) + o = torch.empty_like(v) + + has_seqlens = seqlens is not None + if seqlens is None: + seqlens = torch.empty((B,), dtype=torch.int32, device=k.device) + seqlen_dtype = seqlens.dtype + + kern = tilelang_fused_recurrent_gdr_fwd( + H, + Hg, + K, + V, + scale, + accum_dtype="float32", + qkva_dtype=q.dtype, + g_dtype=g.dtype, + b_dtype=beta.dtype, + h0_dtype=initial_state.dtype, + ht_dtype=final_state.dtype, + o_dtype=o.dtype, + seqlen_dtype=seqlen_dtype, + use_initial_state=use_initial_state, + store_final_state=output_final_state, + has_seqlens=has_seqlens, + block_DV=block_DV, + threads=128, + ) + kern(q, k, v, g, beta, initial_state, seqlens, o, final_state) + return o, (final_state if output_final_state else None) diff --git a/flash_qla/ops/gated_delta_rule/fused_recurrent/hopper/fused_recurrent_verify.py b/flash_qla/ops/gated_delta_rule/fused_recurrent/hopper/fused_recurrent_verify.py new file mode 100644 index 00000000..1bcacd8b --- /dev/null +++ b/flash_qla/ops/gated_delta_rule/fused_recurrent/hopper/fused_recurrent_verify.py @@ -0,0 +1,538 @@ +# Copyright (c) 2026 The Qwen team, Alibaba Group. +# Licensed under The MIT License [see LICENSE for details] +"""SGLang verify kernel: gemm-free GDN recurrence + paged V-major (bf16) state pool, +per-token intermediate states, no-commit, varlen cu_seqlens. Host-side gating (g/beta +pre-activated, q/k pre-l2normed by the wrapper). CUDA-graph safe (no host sync / no alloc +in the captured entry; all buffers caller-provided).""" +import os + +import torch +import tilelang +import tilelang.language as T + +MULTI_PROCESSOR_COUNT = torch.cuda.get_device_properties().multi_processor_count +TARGET_NUM_CTAS = int(MULTI_PROCESSOR_COUNT * 0.7) + + +@tilelang.jit(pass_configs={tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True}) +def tilelang_fused_recurrent_gdr_verify( + H, + Hg, + DK, + DV, + scale, + accum_dtype, + qkva_dtype, + g_dtype, + b_dtype, + pool_dtype, + o_dtype, + seqlen_dtype, + idx_dtype, + store_intermediate, + disable_state_update, + block_DV=128, + threads=128, +): + total_tokens = T.dynamic("total_tokens") + N = T.dynamic("N") # number of requests + num_slots = T.dynamic("num_slots") + num_cache_slots = T.dynamic("num_cache_slots") + cache_steps = T.dynamic("cache_steps") + q_shape = (1, total_tokens, Hg, DK) + v_shape = (1, total_tokens, H, DV) + g_shape = (1, total_tokens, H) + pool_shape = (num_slots, H, DV, DK) # V-major [., H, V, K] + ibuf_shape = (num_cache_slots, cache_steps, H, DV, DK) # V-major + n_vt = (DV + block_DV - 1) // block_DV + + @T.prim_func + def kernel( + q: T.Tensor(q_shape, qkva_dtype), + k: T.Tensor(q_shape, qkva_dtype), + v: T.Tensor(v_shape, qkva_dtype), + g: T.Tensor(g_shape, g_dtype), + b: T.Tensor(g_shape, b_dtype), + pool: T.Tensor(pool_shape, pool_dtype), + state_indices: T.Tensor([N], idx_dtype), + cu_seqlens: T.Tensor([N + 1], seqlen_dtype), + intermediate_state_indices: T.Tensor([N], idx_dtype), + o: T.Tensor(v_shape, o_dtype), + ibuf: T.Tensor(ibuf_shape, pool_dtype), + ): + with T.Kernel(n_vt * N * H, threads=threads) as (bbhv,): + bbh = bbhv // n_vt + bv = bbhv % n_vt + bb = bbh // H # request index + bh = bbh % H + bhg = bh // (H // Hg) + v0 = bv * block_DV + + slot = T.alloc_var("int32") + cslot = T.alloc_var("int32") + seq_start = T.alloc_var("int32") + seq_end = T.alloc_var("int32") + slot = state_indices[bb] + cslot = intermediate_state_indices[bb] + seq_start = cu_seqlens[bb] + seq_end = cu_seqlens[bb + 1] + + S = T.alloc_fragment((block_DV, DK), accum_dtype) # state [V-tile, K] fp32 + prod = T.alloc_fragment((block_DV, DK), accum_dtype) + q_s = T.alloc_shared((1, DK), qkva_dtype) + k_s = T.alloc_shared((1, DK), qkva_dtype) + v_s = T.alloc_shared((1, block_DV), qkva_dtype) + o_sh = T.alloc_shared((1, block_DV), o_dtype) + kS = T.alloc_fragment((block_DV,), accum_dtype) + oo = T.alloc_fragment((block_DV,), accum_dtype) + vnew = T.alloc_fragment((block_DV,), accum_dtype) + decay = T.alloc_fragment((1,), accum_dtype) + bt = T.alloc_fragment((1,), accum_dtype) + + # gather V-major pool[slot, bh, v0:v0+block_DV, :] directly into S[v, dk] + T.clear(S) + with T.If(slot >= 0): + with T.Then(): + for j_v, j_k in T.Parallel(block_DV, DK): + S[j_v, j_k] = pool[slot, bh, v0 + j_v, j_k] + + for t in T.serial(seq_end - seq_start): + tt = seq_start + t # absolute token position in the flattened layout + T.copy(q[0, tt : tt + 1, bhg, 0:DK], q_s) + T.copy(k[0, tt : tt + 1, bhg, 0:DK], k_s) + T.copy(v[0, tt : tt + 1, bh, v0 : v0 + block_DV], v_s) + decay[0] = T.exp2(g[0, tt, bh] * 1.442695) + bt[0] = b[0, tt, bh] + for j_v, j_k in T.Parallel(block_DV, DK): + S[j_v, j_k] *= decay[0] + for j_v, j_k in T.Parallel(block_DV, DK): + prod[j_v, j_k] = k_s[0, j_k] * S[j_v, j_k] + T.reduce_sum(prod, kS, dim=1) + for j_v in T.Parallel(block_DV): + vnew[j_v] = bt[0] * (v_s[0, j_v] - kS[j_v]) + for j_v, j_k in T.Parallel(block_DV, DK): + S[j_v, j_k] += k_s[0, j_k] * vnew[j_v] + for j_v, j_k in T.Parallel(block_DV, DK): + prod[j_v, j_k] = q_s[0, j_k] * S[j_v, j_k] + T.reduce_sum(prod, oo, dim=1) + for j_v in T.Parallel(block_DV): + o_sh[0, j_v] = oo[j_v] * scale + T.copy(o_sh, o[0, tt : tt + 1, bh, v0 : v0 + block_DV]) + # per-token intermediate (V-major), gated by the POOL slot mask + if store_intermediate: + with T.If(slot >= 0): + with T.Then(): + for j_v, j_k in T.Parallel(block_DV, DK): + ibuf[cslot, t, bh, v0 + j_v, j_k] = S[j_v, j_k] + + # commit final state to the pool unless no-commit (verify) + if not disable_state_update: + with T.If(slot >= 0): + with T.Then(): + for j_v, j_k in T.Parallel(block_DV, DK): + pool[slot, bh, v0 + j_v, j_k] = S[j_v, j_k] + + return kernel + + +def fused_recurrent_gdr_verify_fwd( + q, + k, + v, + g, + beta, + pool, + state_indices, + cu_seqlens, + intermediate_states_buffer, + intermediate_state_indices, + o, + scale=None, + disable_state_update=True, +): + """Graph-safe low-level verify entry. ALL buffers caller-preallocated (o, pool, ibuf); + no host sync, no allocation. g/beta pre-activated and q/k pre-l2normed host-side.""" + _, total_tokens, Hg, K = k.shape + _, _, H, V = v.shape + N = state_indices.shape[0] + assert K == V == 128 and H % Hg == 0 + scale = scale or K ** -0.5 + store_intermediate = intermediate_states_buffer is not None + + # block_DV=64 (2 V-tiles) @ threads=128 is the bandwidth sweet spot (autotuned, H100); + # 32 (4 V-tiles) for the low-CTA tail. block_DV=128 is occupancy-starved -> never used. + grid_base = N * H + block_DV = 64 if grid_base * 2 >= TARGET_NUM_CTAS else 32 + + kern = tilelang_fused_recurrent_gdr_verify( + H, + Hg, + K, + V, + scale, + accum_dtype="float32", + qkva_dtype=q.dtype, + g_dtype=g.dtype, + b_dtype=beta.dtype, + pool_dtype=pool.dtype, + o_dtype=o.dtype, + seqlen_dtype=cu_seqlens.dtype, + idx_dtype=state_indices.dtype, + store_intermediate=store_intermediate, + disable_state_update=disable_state_update, + block_DV=block_DV, + threads=128, + ) + kern(q, k, v, g, beta, pool, state_indices, cu_seqlens, + intermediate_state_indices, o, intermediate_states_buffer) + return o + + +@tilelang.jit(pass_configs={tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True}) +def tilelang_fused_recurrent_gdr_verify_gated( + H, Hg, DK, DV, scale, accum_dtype, qkva_dtype, ab_dtype, gate_dtype, pool_dtype, + o_dtype, seqlen_dtype, idx_dtype, store_intermediate, disable_state_update, + l2norm_eps=1e-6, softplus_thr=20.0, allow_neg_eigval=False, block_DV=128, threads=128, +): + """In-kernel fused-gating verify kernel (req #5): takes raw a,b,A_log,dt_bias and + computes g=-exp(A_log)*softplus(a+dt_bias), beta=sigmoid(b), and qk-l2norm in-kernel.""" + total_tokens = T.dynamic("total_tokens") + N = T.dynamic("N") + num_slots = T.dynamic("num_slots") + num_cache_slots = T.dynamic("num_cache_slots") + cache_steps = T.dynamic("cache_steps") + qk_shape = (1, total_tokens, Hg, DK) + v_shape = (1, total_tokens, H, DV) + ab_shape = (1, total_tokens, H) + pool_shape = (num_slots, H, DV, DK) + ibuf_shape = (num_cache_slots, cache_steps, H, DV, DK) + n_vt = (DV + block_DV - 1) // block_DV + beta_mul = 2.0 if allow_neg_eigval else 1.0 + + @T.prim_func + def kernel( + q: T.Tensor(qk_shape, qkva_dtype), + k: T.Tensor(qk_shape, qkva_dtype), + v: T.Tensor(v_shape, qkva_dtype), + a: T.Tensor(ab_shape, ab_dtype), + b: T.Tensor(ab_shape, ab_dtype), + A_log: T.Tensor([H], gate_dtype), + dt_bias: T.Tensor([H], gate_dtype), + pool: T.Tensor(pool_shape, pool_dtype), + state_indices: T.Tensor([N], idx_dtype), + cu_seqlens: T.Tensor([N + 1], seqlen_dtype), + intermediate_state_indices: T.Tensor([N], idx_dtype), + o: T.Tensor(v_shape, o_dtype), + ibuf: T.Tensor(ibuf_shape, pool_dtype), + ): + with T.Kernel(n_vt * N * H, threads=threads) as (bbhv,): + bbh = bbhv // n_vt + bv = bbhv % n_vt + bb = bbh // H + bh = bbh % H + bhg = bh // (H // Hg) + v0 = bv * block_DV + + slot = T.alloc_var("int32") + cslot = T.alloc_var("int32") + seq_start = T.alloc_var("int32") + seq_end = T.alloc_var("int32") + slot = state_indices[bb] + cslot = intermediate_state_indices[bb] + seq_start = cu_seqlens[bb] + seq_end = cu_seqlens[bb + 1] + a_log_h = T.alloc_var(accum_dtype) + dt_b_h = T.alloc_var(accum_dtype) + a_log_h = A_log[bh] + dt_b_h = dt_bias[bh] + + S = T.alloc_fragment((block_DV, DK), accum_dtype) + prod = T.alloc_fragment((block_DV, DK), accum_dtype) + q_s = T.alloc_shared((1, DK), qkva_dtype) + k_s = T.alloc_shared((1, DK), qkva_dtype) + q_n = T.alloc_shared((1, DK), qkva_dtype) + k_n = T.alloc_shared((1, DK), qkva_dtype) + v_s = T.alloc_shared((1, block_DV), qkva_dtype) + o_sh = T.alloc_shared((1, block_DV), o_dtype) + qsq = T.alloc_fragment((1, DK), accum_dtype) + ksq = T.alloc_fragment((1, DK), accum_dtype) + ssq = T.alloc_fragment((1,), accum_dtype) + kS = T.alloc_fragment((block_DV,), accum_dtype) + oo = T.alloc_fragment((block_DV,), accum_dtype) + vnew = T.alloc_fragment((block_DV,), accum_dtype) + decay = T.alloc_fragment((1,), accum_dtype) + bt = T.alloc_fragment((1,), accum_dtype) + + T.clear(S) + with T.If(slot >= 0): + with T.Then(): + for j_v, j_k in T.Parallel(block_DV, DK): + S[j_v, j_k] = pool[slot, bh, v0 + j_v, j_k] + + for t in T.serial(seq_end - seq_start): + tt = seq_start + t + T.copy(q[0, tt : tt + 1, bhg, 0:DK], q_s) + T.copy(k[0, tt : tt + 1, bhg, 0:DK], k_s) + T.copy(v[0, tt : tt + 1, bh, v0 : v0 + block_DV], v_s) + # in-kernel qk l2norm: x_n = x / sqrt(sum(x^2) + eps) + for _i, j_k in T.Parallel(1, DK): + qsq[0, j_k] = q_s[0, j_k] * q_s[0, j_k] + T.reduce_sum(qsq, ssq, dim=1) + for _i, j_k in T.Parallel(1, DK): + q_n[0, j_k] = q_s[0, j_k] * T.rsqrt(ssq[0] + l2norm_eps) + for _i, j_k in T.Parallel(1, DK): + ksq[0, j_k] = k_s[0, j_k] * k_s[0, j_k] + T.reduce_sum(ksq, ssq, dim=1) + for _i, j_k in T.Parallel(1, DK): + k_n[0, j_k] = k_s[0, j_k] * T.rsqrt(ssq[0] + l2norm_eps) + # in-kernel gating: g = -exp(A_log)*softplus(a+dt_bias); beta = sigmoid(b) + x = a[0, tt, bh] + dt_b_h + sp = T.if_then_else(x > softplus_thr, x, T.log(1.0 + T.exp(x))) + decay[0] = T.exp2((-T.exp(a_log_h) * sp) * 1.442695) + bt[0] = beta_mul * T.sigmoid(b[0, tt, bh]) + for j_v, j_k in T.Parallel(block_DV, DK): + S[j_v, j_k] *= decay[0] + for j_v, j_k in T.Parallel(block_DV, DK): + prod[j_v, j_k] = k_n[0, j_k] * S[j_v, j_k] + T.reduce_sum(prod, kS, dim=1) + for j_v in T.Parallel(block_DV): + vnew[j_v] = bt[0] * (v_s[0, j_v] - kS[j_v]) + for j_v, j_k in T.Parallel(block_DV, DK): + S[j_v, j_k] += k_n[0, j_k] * vnew[j_v] + for j_v, j_k in T.Parallel(block_DV, DK): + prod[j_v, j_k] = q_n[0, j_k] * S[j_v, j_k] + T.reduce_sum(prod, oo, dim=1) + for j_v in T.Parallel(block_DV): + o_sh[0, j_v] = oo[j_v] * scale + T.copy(o_sh, o[0, tt : tt + 1, bh, v0 : v0 + block_DV]) + if store_intermediate: + with T.If(slot >= 0): + with T.Then(): + for j_v, j_k in T.Parallel(block_DV, DK): + ibuf[cslot, t, bh, v0 + j_v, j_k] = S[j_v, j_k] + + if not disable_state_update: + with T.If(slot >= 0): + with T.Then(): + for j_v, j_k in T.Parallel(block_DV, DK): + pool[slot, bh, v0 + j_v, j_k] = S[j_v, j_k] + + return kernel + + +def fused_recurrent_gdr_verify_gated_fwd( + q, k, v, a, b, A_log, dt_bias, pool, state_indices, cu_seqlens, + intermediate_states_buffer, intermediate_state_indices, o, + scale=None, disable_state_update=True, allow_neg_eigval=False, +): + """In-kernel fused-gating verify entry: raw a,b,A_log,dt_bias; computes g/beta + qk-l2norm + inside the kernel. Graph-safe (all buffers caller-provided).""" + _, total_tokens, Hg, K = k.shape + _, _, H, V = v.shape + N = state_indices.shape[0] + assert K == V == 128 and H % Hg == 0 + scale = scale or K ** -0.5 + store_intermediate = intermediate_states_buffer is not None + + grid_base = N * H # bandwidth sweet spot (autotuned, H100): block_DV=64 @ threads=128 + block_DV = 64 if grid_base * 2 >= TARGET_NUM_CTAS else 32 + + kern = tilelang_fused_recurrent_gdr_verify_gated( + H, Hg, K, V, scale, + accum_dtype="float32", qkva_dtype=q.dtype, ab_dtype=a.dtype, gate_dtype=A_log.dtype, + pool_dtype=pool.dtype, o_dtype=o.dtype, seqlen_dtype=cu_seqlens.dtype, + idx_dtype=state_indices.dtype, store_intermediate=store_intermediate, + disable_state_update=disable_state_update, allow_neg_eigval=allow_neg_eigval, + block_DV=block_DV, threads=max(128, block_DV * 2), + ) + kern(q, k, v, a, b, A_log, dt_bias, pool, state_indices, cu_seqlens, + intermediate_state_indices, o, intermediate_states_buffer) + return o + + +# ---------------------------------------------------------------------------------------------- +# H1: gating + qk-l2norm DEDUP PRE-PASS. +# The in-kernel-gated kernel above recomputes g/beta + qk-l2norm INSIDE the per-token hot loop, +# once per (token, V-head, V-tile) CTA -> l2norm redundant grp*n_vt times per (token,K-head), +# gating redundant n_vt times per (token,V-head). This pre-pass computes each ONCE (grid = +# total_tokens*Hk, one CTA per (token, K-head)) and feeds the HOST-gated main kernel +# (fused_recurrent_gdr_verify_fwd) whose hot loop then has NO transcendentals. Measured ceiling +# (gated-vs-host-gated, H100): 14-16% at T=12 large batch. Regime-gated (see should_use_prepass): +# at single-request the second-launch tax exceeds the small ceiling, so variant A is kept there. +# ---------------------------------------------------------------------------------------------- + +# Regime gate tunables. CALIBRATED FOR THE CUDA-GRAPH (production) PATH (H100, +# benchmark/bench_prepass.py bench_graph, _time_graph with 50+ warmup). Prepass wins when (a) drafts +# are long enough that the per-token gating/l2norm recompute it dedups is a meaningful fraction +# (T_avg >= 4; at T=1 it is paid once -> nothing to dedup across tokens, measured neutral/loss even +# under graphs), AND (b) the main-kernel work N*H*(1+T_avg) clears a SMALL floor. +# +# KEY: under CUDA-graph REPLAY the prepass's 2nd launch shrinks to a graph node, so the eager +# second-launch tax (~13-15us) VANISHES and the dedup win dominates down to SINGLE-REQUEST. Measured +# (graph, Hk=16 Hv=32): N=1,T=12 work=416 -> 1.18x WIN (eager was 0.60x LOSS); N=2,T=12 -> 1.08x; +# N=4,T=12 -> 1.18x; every T=12 point N>=1 wins 1.08-1.24x. The only graph losses are work<=320 +# (N=1,T=4=160 0.96x; N=2,T=4=320 0.93x). So the gate window for the verify (always T=12) is +# work in (320, 416]; MIN_WORK=384 fires the prepass for the ENTIRE T=12 verify path incl. N=1, +# capturing the 1.08-1.24x the old eager-conservative 3000 was leaving on the table at small batch. +# TRADEOFF: this assumes the CUDA-graph deployment (qwen36 captures all batch sizes). Under EAGER +# (warmup / --disable-cuda-graph), small-N now pays the 2nd-launch tax (N=1,T=12 eager 0.60x) -- but +# steady-state production verify is always captured, and warmup is untimed. Metric scales with H, so +# small-head configs self-raise the batch needed (gate stays safe for untested H<32). +PREPASS_MIN_T = 4.0 +# graph-calibrated default 384 (was 3000, eager-conservative); floor below N=1/T=12 work=416. +# Env-overridable for controlled A/B (set FLASHQLA_PREPASS_MIN_WORK=3000 to reproduce the old gate). +PREPASS_MIN_WORK = int(os.environ.get("FLASHQLA_PREPASS_MIN_WORK", "384")) + + +def should_use_prepass(N, H, total_tokens): + """Static (capture-safe; shape-only) decision: run the dedup prepass + host-gated main kernel + (True) vs the single in-kernel-gated kernel A (False). Both produce identical output; this + only picks the faster path per regime. CALIBRATED FOR THE CUDA-GRAPH (production) path, where + the prepass's 2nd launch is a free graph node and it wins down to single-request T>=4 (the eager + second-launch tax that favored variant A at small batch does NOT exist under graph replay).""" + if N <= 0: + return False + t_avg = total_tokens / N + work = H * (N + total_tokens) # == N*H*(1 + t_avg), proportional to the main-kernel runtime + return (t_avg >= PREPASS_MIN_T) and (work >= PREPASS_MIN_WORK) + + +@tilelang.jit(pass_configs={tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True}) +def tilelang_gdr_verify_prepass( + Hk, Hv, DK, accum_dtype, qk_dtype, ab_dtype, gate_dtype, g_out_dtype, b_out_dtype, + l2norm_eps=1e-6, softplus_thr=20.0, allow_neg_eigval=False, threads=128, +): + """Dedup pre-pass: ONE CTA per token tt computes q_n/k_n = l2norm(q/k) for all Hk K-heads at + once via a parallel [Hk,DK] reduce (the proven reduce_sum(...,dim=1) idiom), and + g = -exp(A_log[h])*softplus(a+dt_bias) (RAW log-decay; the main kernel applies exp2) + + beta = beta_mul*sigmoid(b) for all Hv V-heads. No recurrence, no cu_seqlens (token-local): + the main kernel only consumes tokens within cu_seqlens, so normalizing trailing padding is + harmless. + + q/k are read into a fragment and q_n/k_n written DIRECTLY from a full-range T.Parallel (the + head-batch global-write idiom) -- no per-row [1,DK] T.copy, so no serial Hk loop (that was a + measured ~6x slowdown) and no Hopper copy-layout trap. The gate write is the full contiguous + [1,Hv] row (Hv>=4 -> >=128-bit fp32 extent; a per-K-head [1,grp] write fails layout inference + for grp<4). Grid = total_tokens balances occupancy (~total_tokens CTAs).""" + beta_mul = 2.0 if allow_neg_eigval else 1.0 + total_tokens = T.dynamic("total_tokens") + qk_shape = (1, total_tokens, Hk, DK) + ab_shape = (1, total_tokens, Hv) + + @T.prim_func + def kernel( + q: T.Tensor(qk_shape, qk_dtype), + k: T.Tensor(qk_shape, qk_dtype), + a: T.Tensor(ab_shape, ab_dtype), + b: T.Tensor(ab_shape, ab_dtype), + A_log: T.Tensor([Hv], gate_dtype), + dt_bias: T.Tensor([Hv], gate_dtype), + q_n: T.Tensor(qk_shape, qk_dtype), + k_n: T.Tensor(qk_shape, qk_dtype), + g_out: T.Tensor(ab_shape, g_out_dtype), + beta_out: T.Tensor(ab_shape, b_out_dtype), + ): + with T.Kernel(total_tokens, threads=threads) as (tt,): + xf = T.alloc_fragment((Hk, DK), accum_dtype) # raw q/k rows (reused q then k) + sq = T.alloc_fragment((Hk, DK), accum_dtype) # squared, reduce input + ssq = T.alloc_fragment((Hk,), accum_dtype) # per-K-head sum-of-squares + g_f = T.alloc_fragment((1, Hv), accum_dtype) + b_f = T.alloc_fragment((1, Hv), accum_dtype) + g_sh = T.alloc_shared((1, Hv), g_out_dtype) + b_sh = T.alloc_shared((1, Hv), b_out_dtype) + + # q l2norm: load [Hk,DK] -> square -> reduce over DK -> normalize + write (all parallel) + for i, j in T.Parallel(Hk, DK): + xf[i, j] = q[0, tt, i, j] + for i, j in T.Parallel(Hk, DK): + sq[i, j] = xf[i, j] * xf[i, j] + T.reduce_sum(sq, ssq, dim=1) + for i, j in T.Parallel(Hk, DK): + q_n[0, tt, i, j] = xf[i, j] * T.rsqrt(ssq[i] + l2norm_eps) + + # k l2norm (reuse xf/sq/ssq) + for i, j in T.Parallel(Hk, DK): + xf[i, j] = k[0, tt, i, j] + for i, j in T.Parallel(Hk, DK): + sq[i, j] = xf[i, j] * xf[i, j] + T.reduce_sum(sq, ssq, dim=1) + for i, j in T.Parallel(Hk, DK): + k_n[0, tt, i, j] = xf[i, j] * T.rsqrt(ssq[i] + l2norm_eps) + + # gating for all Hv V-heads (fragment idiom, inlined exprs, direct global reads), + # staged through [1,Hv] shared and written as one contiguous row (>=128-bit extent) + for _i, h in T.Parallel(1, Hv): + g_f[0, h] = -T.exp(A_log[h]) * T.if_then_else( + a[0, tt, h] + dt_bias[h] > softplus_thr, + a[0, tt, h] + dt_bias[h], + T.log(1.0 + T.exp(a[0, tt, h] + dt_bias[h])), + ) + b_f[0, h] = beta_mul * T.sigmoid(b[0, tt, h]) + for _i, h in T.Parallel(1, Hv): + g_sh[0, h] = g_f[0, h] + b_sh[0, h] = b_f[0, h] + T.copy(g_sh, g_out[0, tt : tt + 1, 0:Hv]) + T.copy(b_sh, beta_out[0, tt : tt + 1, 0:Hv]) + + return kernel + + +def fused_recurrent_gdr_verify_prepass( + q, k, a, b, A_log, dt_bias, + q_n=None, k_n=None, g_out=None, beta_out=None, allow_neg_eigval=False, +): + """Dedup pre-pass dispatch. Computes q_n=l2norm(q), k_n=l2norm(k), g=raw log-decay, + beta=beta_mul*sigmoid(b) ONCE, feeding the host-gated verify main kernel. Buffers are + caller-provided for CUDA-graph capture safety (like o/pool/ibuf); allocate-if-None is an + EAGER convenience for tests/bench only -- it must NOT run inside a captured region (no alloc + in capture). g/beta are fp32 (matching gdn_sigmoid_gate); q_n/k_n match q/k dtype.""" + _, total_tokens, Hk, K = q.shape + Hv = a.shape[2] + assert Hv % Hk == 0, f"num_v_heads {Hv} must be divisible by num_k_heads {Hk}" + assert total_tokens > 0 + if q_n is None: + q_n = torch.empty_like(q) + if k_n is None: + k_n = torch.empty_like(k) + if g_out is None: + g_out = torch.empty(1, total_tokens, Hv, device=q.device, dtype=torch.float32) + if beta_out is None: + beta_out = torch.empty(1, total_tokens, Hv, device=q.device, dtype=torch.float32) + + kern = tilelang_gdr_verify_prepass( + Hk, Hv, K, + accum_dtype="float32", qk_dtype=q.dtype, ab_dtype=a.dtype, gate_dtype=A_log.dtype, + g_out_dtype=g_out.dtype, b_out_dtype=beta_out.dtype, + allow_neg_eigval=allow_neg_eigval, threads=128, + ) + kern(q, k, a, b, A_log, dt_bias, q_n, k_n, g_out, beta_out) + return q_n, k_n, g_out, beta_out + + +# Module-level UNBOUNDED, never-evicting scratch cache for the prepass outputs. Unbounded is the +# capture-safe choice: each captured graph bakes the device pointers of its own (total_tokens,...) +# entry; an evicting LRU (e.g. utils.tensor_cache) could free a still-referenced storage and make +# a later replay read freed memory. One live entry per captured (N,T); entries are never freed. +_PREPASS_SCRATCH = {} + + +def get_prepass_scratch(total_tokens, Hk, Hv, device, qk_dtype): + """Persistent prepass scratch (q_n,k_n bf16 [1,T,Hk,128]; g,beta fp32 [1,T,Hv]). Allocated + lazily on the first (warmup) call and reused thereafter so addresses are stable across + CUDA-graph capture/replay. Raises (rather than allocating + aborting capture) on a cold miss + during capture -- warmup must exercise every (N,T) that will be captured.""" + key = (total_tokens, Hk, Hv, device, qk_dtype) + buf = _PREPASS_SCRATCH.get(key) + if buf is None: + if torch.cuda.is_available() and torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + f"prepass scratch not warmed for shape {key}; run an eager warmup pass before " + "CUDA-graph capture (the scratch must be allocated outside the captured region)." + ) + q_n = torch.empty(1, total_tokens, Hk, 128, device=device, dtype=qk_dtype) + k_n = torch.empty(1, total_tokens, Hk, 128, device=device, dtype=qk_dtype) + g_out = torch.empty(1, total_tokens, Hv, device=device, dtype=torch.float32) + beta_out = torch.empty(1, total_tokens, Hv, device=device, dtype=torch.float32) + buf = (q_n, k_n, g_out, beta_out) + _PREPASS_SCRATCH[key] = buf + return buf diff --git a/tests/probes/debug_decode.py b/tests/probes/debug_decode.py new file mode 100644 index 00000000..350e3cb6 --- /dev/null +++ b/tests/probes/debug_decode.py @@ -0,0 +1,39 @@ +# tests/probes/debug_decode.py +"""Debug the decode kernel vs reference on the simplest case (B=1,Hk=Hv=1,D=1).""" +import torch +from ref_gdr import decode_recur +from flash_qla import recurrent_gated_delta_rule +from flash_qla.utils import l2norm + +torch.manual_seed(0) +B, D, Hk, Hv = 1, 1, 1, 1 +q = l2norm(torch.randn(B, D, Hk, 128, device="cuda", dtype=torch.bfloat16)) +k = l2norm(torch.randn(B, D, Hk, 128, device="cuda", dtype=torch.bfloat16)) +v = torch.randn(B, D, Hv, 128, device="cuda", dtype=torch.bfloat16) +g = torch.nn.functional.logsigmoid(torch.randn(B, D, Hv, device="cuda")) / 16 +beta = torch.randn(B, D, Hv, device="cuda").sigmoid() + +o_ref, s_ref = decode_recur(q, k, v, g, beta, scale=128 ** -0.5) +o_qla, s_qla = recurrent_gated_delta_rule(q, k, v, g, beta, scale=128 ** -0.5, output_final_state=True) + +print("g:", g.item(), "beta:", beta.item()) +print("o_ref[0,0,0,:6]:", o_ref[0, 0, 0, :6].tolist()) +print("o_qla[0,0,0,:6]:", o_qla[0, 0, 0, :6].float().tolist()) +print("o rel err:", ((o_qla.float() - o_ref).abs().max() / o_ref.abs().max()).item()) +print() +print("s_ref[0,0,:3,:4]:\n", s_ref[0, 0, :3, :4]) +print("s_qla[0,0,:3,:4]:\n", s_qla[0, 0, :3, :4].float()) +print("s rel err:", ((s_qla - s_ref).abs().max() / s_ref.abs().max()).item()) +sd = (s_qla - s_ref).abs() +idx = sd.argmax() +dk_i, dv_i = (idx % (128 * 128)) // 128, idx % 128 +print(f"s max-err at [dk={dk_i}, dv={dv_i}]: qla={s_qla[0,0,dk_i,dv_i].item():.5f} ref={s_ref[0,0,dk_i,dv_i].item():.5f}") +print("s per-col-tile max err (cols 0-31,32-63,64-95,96-127):", + [round((s_qla[0,0,:,c:c+32]-s_ref[0,0,:,c:c+32]).abs().max().item(), 4) for c in range(0, 128, 32)]) + +# manual: S should be k (x) (beta*v); o = scale*(q.k)*(beta*v) +manual_S = torch.outer(k[0, 0, 0].float(), beta.item() * v[0, 0, 0].float()) +print("\nmanual_S[0,:4]:", manual_S[0, :4].tolist()) +print("s_ref[0,0,0,:4]:", s_ref[0, 0, 0, :4].tolist()) +qk = (q[0, 0, 0].float() * k[0, 0, 0].float()).sum().item() +print("q.k:", qk, "| manual o[0]:", 128 ** -0.5 * qk * (beta.item() * v[0, 0, 0, 0].float()).item()) diff --git a/tests/probes/probe_gemm_m1.py b/tests/probes/probe_gemm_m1.py new file mode 100644 index 00000000..4e7bc1bf --- /dev/null +++ b/tests/probes/probe_gemm_m1.py @@ -0,0 +1,47 @@ +# tests/probes/probe_gemm_m1.py +"""Gate 1: does gemm_v1 accept M=1? If not, M-pad to 16. Root gate for the whole engine.""" +import torch, tilelang +import tilelang.language as T + +DK = DV = 128 + + +def build(M): # M = padded token rows (1 to test the gate directly; 16 = fallback) + @tilelang.jit(pass_configs={tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True}) + def _k(): + @T.prim_func + def k(kq: T.Tensor([M, DK], "bfloat16"), + s: T.Tensor([DK, DV], "bfloat16"), + o: T.Tensor([M, DV], "float32")): + with T.Kernel(1, threads=256) as _: + ks = T.alloc_shared((M, DK), "bfloat16") + ss = T.alloc_shared((DK, DV), "bfloat16") + of = T.alloc_fragment((M, DV), "float32") + T.copy(kq, ks); T.copy(s, ss) + T.gemm_v1(ks, ss, of, clear_accum=True) + T.copy(of, o) + return k + return _k() + + +def run(M): + torch.manual_seed(0) + k = torch.randn(M, DK, device="cuda", dtype=torch.bfloat16) + s = torch.randn(DK, DV, device="cuda", dtype=torch.bfloat16) + o = torch.empty(M, DV, device="cuda", dtype=torch.float32) + try: + build(M)(k, s, o) + except Exception as e: + print(f"M={M}: COMPILE/RUN FAIL {type(e).__name__}: {e}") + return False + ref = (k.float() @ s.float()) + err = (o - ref).abs().max().item() / ref.abs().max().item() + ok = err < 0.02 + print(f"M={M}: rel_err={err:.4f} {'OK' if ok else 'FAIL'}") + return ok + + +if __name__ == "__main__": + m1 = run(1) + m16 = run(16) + print("\nDECISION: M=1 usable directly:", m1, "| M-pad-to-16 fallback usable:", m16) diff --git a/tests/probes/probe_gemm_shapes.py b/tests/probes/probe_gemm_shapes.py new file mode 100644 index 00000000..cd920317 --- /dev/null +++ b/tests/probes/probe_gemm_shapes.py @@ -0,0 +1,64 @@ +# tests/probes/probe_gemm_shapes.py +"""Find a threads + warp-policy config under which the 3 decode gemm shapes all compile. +Shapes (M-padded to 16): kS = [16,128]@[128,128]; o same; rank-1 = [16,128]^T @ [16,128] -> [128,128].""" +import inspect +import torch +import tilelang +import tilelang.language as T + +print("gemm_v1 signature:", inspect.signature(T.gemm_v1)) +pol = None +for cand in ["GemmWarpPolicy"]: + if hasattr(T, cand): + pol = getattr(T, cand) + print(f"T.{cand}:", [p for p in dir(pol) if not p.startswith("_")]) +try: + from tilelang import GemmWarpPolicy as GWP + print("tilelang.GemmWarpPolicy:", [p for p in dir(GWP) if not p.startswith("_")]) + pol = pol or GWP +except Exception as e: + print("no tilelang.GemmWarpPolicy:", e) + +DK = DV = 128 +MPAD = 16 + + +def build(kind, threads, policy): + @tilelang.jit(pass_configs={tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True}) + def _k(): + @T.prim_func + def k(a: T.Tensor([MPAD, DK], "bfloat16"), + bb: T.Tensor([DK if kind != "rank1" else MPAD, DV], "bfloat16"), + c: T.Tensor([MPAD if kind != "rank1" else DK, DV], "float32")): + with T.Kernel(1, threads=threads) as _: + as_ = T.alloc_shared(a.shape, "bfloat16") + bs_ = T.alloc_shared(bb.shape, "bfloat16") + cf = T.alloc_fragment(c.shape, "float32") + T.copy(a, as_); T.copy(bb, bs_) + kw = {} if policy is None else {"policy": policy} + if kind == "rank1": + T.gemm_v1(as_, bs_, cf, transpose_A=True, clear_accum=True, **kw) + else: + T.gemm_v1(as_, bs_, cf, clear_accum=True, **kw) + T.copy(cf, c) + return k + return _k() + + +policies = [None] +if pol is not None: + for name in ["Square", "FullRow", "FullCol"]: + if hasattr(pol, name): + policies.append((name, getattr(pol, name))) + +for kind in ["kS", "rank1"]: + for threads in [64, 128, 256]: + for p in policies: + pname = p[0] if isinstance(p, tuple) else "default" + pval = p[1] if isinstance(p, tuple) else p + try: + build(kind, threads, pval) + print(f" {kind} threads={threads} policy={pname}: COMPILE OK") + except Exception as e: + msg = str(e).splitlines()[-1][:80] + print(f" {kind} threads={threads} policy={pname}: FAIL {msg}") diff --git a/tests/probes/probe_nogemm.py b/tests/probes/probe_nogemm.py new file mode 100644 index 00000000..6625410f --- /dev/null +++ b/tests/probes/probe_nogemm.py @@ -0,0 +1,84 @@ +# tests/probes/probe_nogemm.py +"""Probe a gemm-free decode step: state [BV, DK], GEMVs as reductions over the last dim, +rank-1 as a T.Parallel outer product. Verify one step vs manual.""" +import inspect +import torch +import tilelang +import tilelang.language as T + +print("reduce fns on T:", [x for x in dir(T) if "reduce" in x.lower()]) +if hasattr(T, "reduce_sum"): + print("reduce_sum sig:", inspect.signature(T.reduce_sum)) + +DK = 128 +BV = 32 # a V-tile + + +def build(): + @tilelang.jit(pass_configs={tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True}) + def _k(): + @T.prim_func + def k( + q: T.Tensor([DK], "bfloat16"), + kk: T.Tensor([DK], "bfloat16"), + vv: T.Tensor([BV], "bfloat16"), + gg: T.Tensor([1], "float32"), + bb: T.Tensor([1], "float32"), + s_out: T.Tensor([BV, DK], "float32"), + o_out: T.Tensor([BV], "float32"), + ): + with T.Kernel(1, threads=64) as _: + q_s = T.alloc_shared([DK], "bfloat16") + k_s = T.alloc_shared([DK], "bfloat16") + v_s = T.alloc_shared([BV], "bfloat16") + S = T.alloc_fragment([BV, DK], "float32") + prod = T.alloc_fragment([BV, DK], "float32") + kS = T.alloc_fragment([BV], "float32") + oo = T.alloc_fragment([BV], "float32") + vnew = T.alloc_fragment([BV], "float32") + T.copy(q, q_s); T.copy(kk, k_s); T.copy(vv, v_s) + T.clear(S) + decay = gg[0] + # decay (S starts 0, so no-op here, but keep the op) + for jv, jk in T.Parallel(BV, DK): + S[jv, jk] *= T.exp2(decay * 1.442695) + # kS[v] = sum_dk k[dk]*S[v,dk] + for jv, jk in T.Parallel(BV, DK): + prod[jv, jk] = k_s[jk] * S[jv, jk] + T.reduce_sum(prod, kS, dim=1) + # vnew[v] = beta*(v[v] - kS[v]) + for jv in T.Parallel(BV): + vnew[jv] = bb[0] * (v_s[jv] - kS[jv]) + # rank-1: S[v,dk] += k[dk]*vnew[v] + for jv, jk in T.Parallel(BV, DK): + S[jv, jk] += k_s[jk] * vnew[jv] + # o[v] = sum_dk q[dk]*S[v,dk] + for jv, jk in T.Parallel(BV, DK): + prod[jv, jk] = q_s[jk] * S[jv, jk] + T.reduce_sum(prod, oo, dim=1) + T.copy(S, s_out) + T.copy(oo, o_out) + return k + return _k() + + +torch.manual_seed(0) +q = torch.nn.functional.normalize(torch.randn(DK, dtype=torch.float32), dim=0).bfloat16() +kk = torch.nn.functional.normalize(torch.randn(DK, dtype=torch.float32), dim=0).bfloat16() +vv = torch.randn(BV, dtype=torch.float32).bfloat16() +gg = torch.tensor([-0.05], dtype=torch.float32) +bb = torch.tensor([0.4], dtype=torch.float32) +s_out = torch.empty(BV, DK, dtype=torch.float32, device="cuda") +o_out = torch.empty(BV, dtype=torch.float32, device="cuda") +try: + build()(q.cuda(), kk.cuda(), vv.cuda(), gg.cuda(), bb.cuda(), s_out, o_out) + # manual: S[v,dk] = k[dk]*beta*v[v]; o[v] = sum_dk q[dk]*S[v,dk] = (q.k)*beta*v[v] + S_ref = torch.outer(bb * vv.float(), kk.float()).cuda() # [BV, DK] + o_ref = (S_ref * q.float().cuda()[None, :]).sum(-1) + print("S rel err:", ((s_out - S_ref).abs().max() / S_ref.abs().max()).item()) + print("o rel err:", ((o_out - o_ref).abs().max() / o_ref.abs().max()).item()) + print("o_out[:4]:", o_out[:4].tolist(), "o_ref[:4]:", o_ref[:4].tolist()) +except Exception as e: + import traceback + traceback.print_exc() + print("FAIL", type(e).__name__, str(e).splitlines()[-1][:120]) diff --git a/tests/probes/probe_serial_runtime_l.py b/tests/probes/probe_serial_runtime_l.py new file mode 100644 index 00000000..91f565aa --- /dev/null +++ b/tests/probes/probe_serial_runtime_l.py @@ -0,0 +1,34 @@ +# tests/probes/probe_serial_runtime_l.py +"""Gate 2: single-role threads=256 kernel with a runtime per-CTA loop bound L=lens[bb].""" +import torch, tilelang +import tilelang.language as T + + +def build(): + @tilelang.jit(pass_configs={tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True}) + def _k(): + B = T.dynamic("B") + @T.prim_func + def k(lens: T.Tensor([B], "int32"), out: T.Tensor([B], "float32")): + with T.Kernel(B, threads=256) as (bb,): + Lv = T.alloc_var("int32"); Lv = lens[bb] + acc = T.alloc_fragment((1,), "float32"); acc[0] = 0.0 + for _t in T.serial(Lv): + acc[0] += 1.0 + out[bb] = acc[0] + return k + return _k() + + +if __name__ == "__main__": + lens = torch.tensor([1, 5, 12, 8], device="cuda", dtype=torch.int32) + out = torch.empty(4, device="cuda", dtype=torch.float32) + try: + build()(lens, out) + ok = torch.allclose(out, lens.float()) + print("out:", out.tolist(), "expected:", lens.tolist(), "->", "OK" if ok else "FAIL") + print("DECISION: runtime-L T.serial in single-role form:", + "USABLE" if ok else "FALLBACK to T.serial(D)+if t=L (spec 11.B)") diff --git a/tests/probes/probe_tilelang_prims.py b/tests/probes/probe_tilelang_prims.py new file mode 100644 index 00000000..44174061 --- /dev/null +++ b/tests/probes/probe_tilelang_prims.py @@ -0,0 +1,45 @@ +# tests/probes/probe_tilelang_prims.py +"""Gate 4: which TileLang math intrinsics exist and lower on SM90. +Decides in-kernel gating feasibility (softplus needs log/log2; l2norm needs rsqrt).""" +import tilelang +import tilelang.language as T + +NAMES = ["exp2", "exp", "log", "log2", "log1p", "rsqrt", "sqrt", "sigmoid", "tanh", "pow", "abs"] + + +def report_attrs(): + have = {n: hasattr(T, n) for n in NAMES} + print("attr presence:", have) + return have + + +def lower_smoke(name): + """Try to actually lower a 1-op kernel using T.; return True if it compiles.""" + fn = getattr(T, name, None) + if fn is None: + return False + + @tilelang.jit(pass_configs={tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True}) + def _k(): + @T.prim_func + def k(x: T.Tensor([128], "float32"), y: T.Tensor([128], "float32")): + with T.Kernel(1, threads=128) as _: + for i in T.Parallel(128): + y[i] = fn(x[i]) + return k + + try: + _k() # JIT/compile + return True + except Exception as e: + print(f" lower {name}: FAIL {type(e).__name__}: {e}") + return False + + +if __name__ == "__main__": + have = report_attrs() + print("lowering:") + lowered = {n: (lower_smoke(n) if have[n] else False) for n in NAMES} + print("lowered:", lowered) + print("\nDECISION: in-kernel gating feasible iff log2(or log)+rsqrt both lower:", + (lowered.get("log2") or lowered.get("log")) and lowered.get("rsqrt")) diff --git a/tests/probes/probe_v_first_store.py b/tests/probes/probe_v_first_store.py new file mode 100644 index 00000000..f31322e4 --- /dev/null +++ b/tests/probes/probe_v_first_store.py @@ -0,0 +1,35 @@ +# tests/probes/probe_v_first_store.py +"""Gate 6: does T.copy / indexed store from a [DK,DV] fragment into a V-major [DV,DK] slice +emit a correct strided store? Use DK!=DV so a transpose bug is NOT numerically silent.""" +import torch, tilelang +import tilelang.language as T + +DK, DV = 128, 64 # deliberately non-square + + +def build(): + @tilelang.jit(pass_configs={tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True}) + def _k(): + @T.prim_func + def k(src: T.Tensor([DK, DV], "float32"), + dst: T.Tensor([DV, DK], "float32")): # V-major destination + with T.Kernel(1, threads=256) as _: + f = T.alloc_fragment((DK, DV), "float32") + T.copy(src, f) + for i, j in T.Parallel(DK, DV): + dst[j, i] = f[i, j] # transposed store + return k + return _k() + + +if __name__ == "__main__": + src = torch.randn(DK, DV, device="cuda", dtype=torch.float32) + dst = torch.empty(DV, DK, device="cuda", dtype=torch.float32) + try: + build()(src, dst) + ok = torch.allclose(dst, src.t(), atol=1e-5) + print("transpose store:", "OK" if ok else "FAIL (max diff %.3e)" % (dst - src.t()).abs().max()) + print("DECISION: explicit transposed-index store works:", bool(ok), + "| if FAIL, need an SMEM transpose stage before the store") + except Exception as e: + print(f"COMPILE/RUN FAIL {type(e).__name__}: {e}") diff --git a/tests/ref_gdr.py b/tests/ref_gdr.py index fde5335f..be80742d 100644 --- a/tests/ref_gdr.py +++ b/tests/ref_gdr.py @@ -698,3 +698,114 @@ def chunk_gated_delta_rule_bwd( dk = torch.sum(dk.reshape(B, T, Hg, -1, K), dim=3) dg = torch_cumsum(dg, chunk_size=64, reverse=True, cu_seqlens=cu_seqlens) return dq, dk, dv, db, dg, dh0 + + +def decode_recur( + q, + k, + v, + g, + beta, # q,k:[B,T,Hk,128] v:[B,T,Hv,128] g,beta:[B,T,Hv] + scale=None, + initial_state=None, # initial_state: [B,Hv,128,128] fp32 or None + seqlens=None, # [B] int32 accepted lengths (default: all T) + read_pre_update=False, # NEGATIVE CONTROL: read o before the rank-1 (wrong) + gqa_mod=False, # NEGATIVE CONTROL: hg = h % Hk (wrong GQA mapping) + band_perm=False, # NEGATIVE CONTROL: cyclically swap V-heads WITHIN each GQA group (wrong) +): + """Ground-truth GDN decode recurrence (spec 2): per (b,h), per token tband + mis-mapping bug. (Use a per-head DISTINCT gate so the permuted result is far from correct.)""" + B, T, Hk, K = k.shape + _, _, Hv, V = v.shape + assert K == V == 128 and Hv % Hk == 0 + scale = scale if scale is not None else K ** -0.5 + grp = Hv // Hk + dev = k.device + S = ( + initial_state.clone().float() + if initial_state is not None + else torch.zeros(B, Hv, K, V, device=dev, dtype=torch.float32) + ) + o = torch.zeros(B, T, Hv, V, device=dev, dtype=torch.float32) + if seqlens is None: + seqlens = torch.full((B,), T, device=dev, dtype=torch.int32) + for b in range(B): + L = int(seqlens[b]) + for t in range(L): + for h in range(Hv): + hg = (h % Hk) if gqa_mod else (h // grp) + qt = q[b, t, hg].float() + kt = k[b, t, hg].float() + vt = v[b, t, h].float() + decay = torch.exp(g[b, t, h].float()) + Sh = S[b, h] * decay # [K,V] + kS = kt @ Sh # [V] + v_new = beta[b, t, h].float() * (vt - kS) # [V] + if read_pre_update: + o[b, t, h] = scale * (qt @ Sh) # WRONG: pre rank-1 + Sh = Sh + torch.outer(kt, v_new) # [K,V] + S[b, h] = Sh + if not read_pre_update: + o[b, t, h] = scale * (qt @ Sh) # [V] (post-update, correct) + if band_perm and grp > 1: # cyclic within-group head swap -> wrong (head-batch band control) + perm = torch.arange(Hv, device=dev) + for hgi in range(Hk): + base = hgi * grp + perm[base : base + grp] = base + (torch.arange(grp, device=dev) + 1) % grp + o = o[:, :, perm, :].contiguous() + S = S[:, perm, :, :].contiguous() + return o, S # o:[B,T,Hv,V] (only [:, :L_b] valid per b), final_state S:[B,Hv,K,V] + + +def verify_ref( + q, + k, + v, + g, + beta, # q,k:[1,T,Hk,128] v:[1,T,Hv,128] g,beta:[1,T,Hv] + pool, # [num_slots, Hv, V=128, K=128] V-major (state_v_first) + state_indices, # [N] int32; <0 => skip (gather + commit) + cu_seqlens, # [N+1] int32; request bb owns flattened tokens [cu_seqlens[bb], cu_seqlens[bb+1]) + intermediate_states_buffer, # [num_cache_slots, cache_steps, Hv, V, K] V-major + intermediate_state_indices, # [N] int32 dense (never -1); ibuf write gated by pool slot + scale=None, + disable_state_update=True, # no-commit +): + """Reference for the SGLang verify kernel: paged V-major bf16 pool, per-token + intermediate states, varlen cu_seqlens, no-commit. Mirrors decode_recur per request.""" + K = V = 128 + Hk, Hv = k.shape[2], v.shape[2] + grp = Hv // Hk + scale = scale if scale is not None else K ** -0.5 + N = len(cu_seqlens) - 1 + dev = k.device + o = torch.zeros(1, q.shape[1], Hv, V, device=dev, dtype=torch.float32) + pool_out = pool.clone() + ibuf_out = intermediate_states_buffer.clone() + for bb in range(N): + slot = int(state_indices[bb]) + cslot = int(intermediate_state_indices[bb]) + s0, s1 = int(cu_seqlens[bb]), int(cu_seqlens[bb + 1]) + for h in range(Hv): + hg = h // grp + # gather V-major pool[slot,h] [V,K] -> decode state S [K,V] + S = (pool[slot, h].float().t().clone() if slot >= 0 + else torch.zeros(K, V, device=dev, dtype=torch.float32)) + for ti, t in enumerate(range(s0, s1)): + decay = torch.exp(g[0, t, h].float()) + S = S * decay + kS = k[0, t, hg].float() @ S + v_new = beta[0, t, h].float() * (v[0, t, h].float() - kS) + S = S + torch.outer(k[0, t, hg].float(), v_new) + o[0, t, h] = scale * (q[0, t, hg].float() @ S) + if slot >= 0: # per-token intermediate, V-major [V,K]; gated by POOL slot mask + ibuf_out[cslot, ti, h] = S.t().to(ibuf_out.dtype) + if slot >= 0 and not disable_state_update: + pool_out[slot, h] = S.t().to(pool_out.dtype) + return o, pool_out, ibuf_out diff --git a/tests/test_decode_gdr.py b/tests/test_decode_gdr.py new file mode 100644 index 00000000..1652c15f --- /dev/null +++ b/tests/test_decode_gdr.py @@ -0,0 +1,137 @@ +# tests/test_decode_gdr.py +import pytest +import torch + +from ref_gdr import decode_recur +from ref_gdr import chunk_gated_delta_rule_fwd as chunk_fwd_ref +from flash_qla.utils import l2norm + +CUDA = pytest.mark.skipif(not torch.cuda.is_available(), reason="needs CUDA") + + +def _mk(B, T, Hk, Hv, seed=0, dtype=torch.float32): + torch.manual_seed(seed) + q = l2norm(torch.randn(B, T, Hk, 128, device="cuda", dtype=dtype)) + k = l2norm(torch.randn(B, T, Hk, 128, device="cuda", dtype=dtype)) + v = torch.randn(B, T, Hv, 128, device="cuda", dtype=dtype) + g = torch.nn.functional.logsigmoid(torch.randn(B, T, Hv, device="cuda")) / 16 + beta = torch.randn(B, T, Hv, device="cuda").sigmoid() + return q, k, v, g, beta + + +def _ref_bf16_inputs(B, T, Hk, Hv, seed=0): + return _mk(B, T, Hk, Hv, seed=seed, dtype=torch.bfloat16) + + +@CUDA +def test_decode_recur_matches_chunk_at_cs64(): + # A length-L single sequence: decode_recur must match the chunk reference on the L-prefix. + B, T, H = 1, 50, 8 + q, k, v, g, beta = _mk(B, T, H, H, dtype=torch.float32) + o_dec, s_dec = decode_recur(q, k, v, g, beta) + g_ref, o_ref, A_ref, h_ref, s_ref = chunk_fwd_ref( + q=q.double(), k=k.double(), v=v.double(), g=g.double(), beta=beta.double(), + scale=128 ** -0.5, initial_state=None, cu_seqlens=None) + assert (o_dec - o_ref.float()).abs().max() / o_ref.abs().max() < 1e-3 + assert (s_dec - s_ref.float()).abs().max() / s_ref.abs().max() < 1e-3 + + +@CUDA +@pytest.mark.parametrize("D", [1, 8]) +@pytest.mark.parametrize("Hk,Hv", [(8, 8), (2, 8), (1, 8)]) +@pytest.mark.parametrize("use_h0", [False, True]) +def test_kernel_matches_reference(D, Hk, Hv, use_h0): + from flash_qla import recurrent_gated_delta_rule + B = 1 + q, k, v, g, beta = _ref_bf16_inputs(B, D, Hk, Hv) + h0 = (torch.randn(B, Hv, 128, 128, device="cuda", dtype=torch.float32) if use_h0 else None) + o_ref, s_ref = decode_recur(q, k, v, g, beta, scale=128 ** -0.5, initial_state=h0) + o_qla, s_qla = recurrent_gated_delta_rule( + q, k, v, g, beta, scale=128 ** -0.5, initial_state=h0, output_final_state=True) + assert (o_qla.float() - o_ref).abs().max() / o_ref.abs().max().clamp_min(1e-6) <= 0.02 + assert (s_qla - s_ref).abs().max() / s_ref.abs().max().clamp_min(1e-6) <= 0.02 + + +@CUDA +def test_kernel_g0_swa_heads(): + from flash_qla import recurrent_gated_delta_rule + B, D, H = 1, 8, 8 + q, k, v, g, beta = _ref_bf16_inputs(B, D, H, H) + g[:, :, :H // 2] = 0.0 # half the heads have no decay + o_ref, s_ref = decode_recur(q, k, v, g, beta, scale=128 ** -0.5) + o_qla, _ = recurrent_gated_delta_rule(q, k, v, g, beta, scale=128 ** -0.5) + assert (o_qla.float() - o_ref).abs().max() / o_ref.abs().max() <= 0.02 + + +@CUDA +def test_kernel_ragged_seqlens(): + from flash_qla import recurrent_gated_delta_rule + B, D, H = 3, 8, 8 + q, k, v, g, beta = _ref_bf16_inputs(B, D, H, H) + seqlens = torch.tensor([1, 5, 8], device="cuda", dtype=torch.int32) + o_ref, s_ref = decode_recur(q, k, v, g, beta, scale=128 ** -0.5, seqlens=seqlens) + o_qla, s_qla = recurrent_gated_delta_rule( + q, k, v, g, beta, scale=128 ** -0.5, seqlens=seqlens, output_final_state=True) + for b in range(B): + L = int(seqlens[b]) + assert (o_qla[b, :L].float() - o_ref[b, :L]).abs().max() / o_ref[b, :L].abs().max() <= 0.02 + assert (s_qla[b] - s_ref[b]).abs().max() / s_ref[b].abs().max() <= 0.02 + + +@CUDA +def test_kernel_low_occupancy_vsplit(): + from flash_qla import recurrent_gated_delta_rule + # B*H small => wrapper picks block_DV in {64,32}; result must still match. + B, D, H = 1, 4, 4 + q, k, v, g, beta = _ref_bf16_inputs(B, D, H, H) + o_ref, _ = decode_recur(q, k, v, g, beta, scale=128 ** -0.5) + o_qla, _ = recurrent_gated_delta_rule(q, k, v, g, beta, scale=128 ** -0.5) + assert (o_qla.float() - o_ref).abs().max() / o_ref.abs().max() <= 0.02 + + +@CUDA +@pytest.mark.parametrize("B,H", [(8, 8), (16, 8)]) # B*H=64 -> block_DV=64; 128 -> block_DV=128 +def test_kernel_high_occupancy(B, H): + from flash_qla import recurrent_gated_delta_rule + D = 4 + q, k, v, g, beta = _ref_bf16_inputs(B, D, H, H) + o_ref, s_ref = decode_recur(q, k, v, g, beta, scale=128 ** -0.5) + o_qla, s_qla = recurrent_gated_delta_rule( + q, k, v, g, beta, scale=128 ** -0.5, output_final_state=True) + assert (o_qla.float() - o_ref).abs().max() / o_ref.abs().max() <= 0.02 + assert (s_qla - s_ref).abs().max() / s_ref.abs().max() <= 0.02 + + +@CUDA +def test_negctrl_step_order(): + # discriminating: kernel matches the post-update reference and DIFFERS from a pre-update one + from flash_qla import recurrent_gated_delta_rule + B, D, H = 1, 6, 8 + q, k, v, g, beta = _ref_bf16_inputs(B, D, H, H) + o_post, _ = decode_recur(q, k, v, g, beta, scale=128 ** -0.5) + o_pre, _ = decode_recur(q, k, v, g, beta, scale=128 ** -0.5, read_pre_update=True) + o_qla, _ = recurrent_gated_delta_rule(q, k, v, g, beta, scale=128 ** -0.5) + assert (o_qla.float() - o_post).abs().max() / o_post.abs().max() <= 0.02 + assert (o_qla.float() - o_pre).abs().max() / o_pre.abs().max() > 0.2 # must NOT match wrong order + + +@CUDA +def test_negctrl_gqa_mapping(): + # discriminating: kernel uses hg=h//grp and must DIFFER from the hg=h%Hk mapping + from flash_qla import recurrent_gated_delta_rule + B, D, Hk, Hv = 1, 6, 2, 8 + q, k, v, g, beta = _ref_bf16_inputs(B, D, Hk, Hv) + o_div, _ = decode_recur(q, k, v, g, beta, scale=128 ** -0.5) + o_mod, _ = decode_recur(q, k, v, g, beta, scale=128 ** -0.5, gqa_mod=True) + o_qla, _ = recurrent_gated_delta_rule(q, k, v, g, beta, scale=128 ** -0.5) + assert (o_qla.float() - o_div).abs().max() / o_div.abs().max() <= 0.02 + assert (o_qla.float() - o_mod).abs().max() / o_mod.abs().max() > 0.2 # must NOT match wrong GQA + + +def test_signature_contract(): + import inspect + from flash_qla import recurrent_gated_delta_rule + sig = inspect.signature(recurrent_gated_delta_rule) + for p in ["q", "k", "v", "g", "beta", "scale", "initial_state", + "output_final_state", "use_qk_l2norm_in_kernel", "seqlens"]: + assert p in sig.parameters diff --git a/tests/test_head_batch_gdr.py b/tests/test_head_batch_gdr.py new file mode 100644 index 00000000..ac880736 --- /dev/null +++ b/tests/test_head_batch_gdr.py @@ -0,0 +1,135 @@ +# tests/test_head_batch_gdr.py +# Head-batched GQA specialization of the gemm-free decode kernel: one CTA processes all +# grp = Hv//Hk V-heads sharing a K/Q head. Validates correctness + a within-group band-swap +# negative control + direct equality with the per-head path. Forces head_batch=True so the +# new code is actually exercised (auto stays OFF for the host-gated decode path). +import pytest +import torch + +from ref_gdr import decode_recur +from flash_qla.utils import l2norm + +CUDA = pytest.mark.skipif(not torch.cuda.is_available(), reason="needs CUDA") + + +def _mk(B, T, Hk, Hv, seed=0, distinct_gate=True): + """bf16 inputs. distinct_gate gives each V-head its own decay magnitude so a within-group + head-band swap produces a far-off result (mandatory for the band-swap negative control).""" + torch.manual_seed(seed) + q = l2norm(torch.randn(B, T, Hk, 128, device="cuda", dtype=torch.bfloat16)) + k = l2norm(torch.randn(B, T, Hk, 128, device="cuda", dtype=torch.bfloat16)) + v = torch.randn(B, T, Hv, 128, device="cuda", dtype=torch.bfloat16) + g = torch.nn.functional.logsigmoid(torch.randn(B, T, Hv, device="cuda")) / 16 + if distinct_gate: + g = g * (1 + torch.arange(Hv, device="cuda").float()[None, None, :]) + beta = torch.randn(B, T, Hv, device="cuda").sigmoid() + return q, k, v, g, beta + + +def _rel(a, b): + return (a.float() - b.float()).abs().max() / b.float().abs().max().clamp_min(1e-6) + + +@CUDA +@pytest.mark.parametrize("Hk,Hv", [(4, 8), (2, 8)]) # grp = 2, 4 +def test_head_batch_compiles_and_runs(Hk, Hv): + # smoke: the factory must JIT-build for both grp values (isolates a lowering failure). + from flash_qla import recurrent_gated_delta_rule + q, k, v, g, beta = _mk(1, 8, Hk, Hv) + o, s = recurrent_gated_delta_rule( + q, k, v, g, beta, scale=128 ** -0.5, output_final_state=True, head_batch=True) + assert o.shape == (1, 8, Hv, 128) and s.shape == (1, Hv, 128, 128) + assert torch.isfinite(o.float()).all() + + +@CUDA +@pytest.mark.parametrize("D", [1, 8, 12]) +@pytest.mark.parametrize("Hk,Hv", [(4, 8), (2, 8)]) # grp = 2, 4 +@pytest.mark.parametrize("use_h0", [False, True]) +def test_head_batch_matches_reference(D, Hk, Hv, use_h0): + from flash_qla import recurrent_gated_delta_rule + B = 1 + q, k, v, g, beta = _mk(B, D, Hk, Hv) + h0 = (torch.randn(B, Hv, 128, 128, device="cuda", dtype=torch.float32) if use_h0 else None) + o_ref, s_ref = decode_recur(q, k, v, g, beta, scale=128 ** -0.5, initial_state=h0) + o_hb, s_hb = recurrent_gated_delta_rule( + q, k, v, g, beta, scale=128 ** -0.5, initial_state=h0, + output_final_state=True, head_batch=True) + assert _rel(o_hb, o_ref) <= 0.02 + assert _rel(s_hb, s_ref) <= 0.02 + + +@CUDA +@pytest.mark.parametrize("Hk,Hv", [(4, 8), (2, 8)]) +def test_head_batch_equals_per_head(Hk, Hv): + # the two paths must agree on identical inputs (cheap regression catch; no fp32 ref). + from flash_qla import recurrent_gated_delta_rule + q, k, v, g, beta = _mk(1, 8, Hk, Hv) + o_hb, s_hb = recurrent_gated_delta_rule( + q, k, v, g, beta, scale=128 ** -0.5, output_final_state=True, head_batch=True) + o_ph, s_ph = recurrent_gated_delta_rule( + q, k, v, g, beta, scale=128 ** -0.5, output_final_state=True, head_batch=False) + assert _rel(o_hb, o_ph) <= 0.02 + assert _rel(s_hb, s_ph) <= 0.02 + + +@CUDA +def test_head_batch_low_occupancy_block_dv32(): + # B*Hg small -> head-batch wrapper picks block_DV=32 (the low-CTA tail); must still match. + from flash_qla import recurrent_gated_delta_rule + q, k, v, g, beta = _mk(1, 8, 2, 8) # Hg=2 -> grid_base=2 -> block_DV=32 + o_ref, _ = decode_recur(q, k, v, g, beta, scale=128 ** -0.5) + o_hb, _ = recurrent_gated_delta_rule(q, k, v, g, beta, scale=128 ** -0.5, head_batch=True) + assert _rel(o_hb, o_ref) <= 0.02 + + +@CUDA +def test_head_batch_high_occupancy_block_dv64(): + # B*Hg large -> head-batch wrapper picks block_DV=64; must still match. + from flash_qla import recurrent_gated_delta_rule + B, Hk, Hv = 16, 4, 8 # Hg=4 -> grid_base=64 -> block_DV=64, grp=2 + q, k, v, g, beta = _mk(B, 4, Hk, Hv) + o_ref, s_ref = decode_recur(q, k, v, g, beta, scale=128 ** -0.5) + o_hb, s_hb = recurrent_gated_delta_rule( + q, k, v, g, beta, scale=128 ** -0.5, output_final_state=True, head_batch=True) + assert _rel(o_hb, o_ref) <= 0.02 + assert _rel(s_hb, s_ref) <= 0.02 + + +@CUDA +def test_head_batch_ragged_seqlens(): + from flash_qla import recurrent_gated_delta_rule + B, Hk, Hv = 3, 2, 8 + q, k, v, g, beta = _mk(B, 8, Hk, Hv) + seqlens = torch.tensor([1, 5, 8], device="cuda", dtype=torch.int32) + o_ref, s_ref = decode_recur(q, k, v, g, beta, scale=128 ** -0.5, seqlens=seqlens) + o_hb, s_hb = recurrent_gated_delta_rule( + q, k, v, g, beta, scale=128 ** -0.5, seqlens=seqlens, + output_final_state=True, head_batch=True) + for b in range(B): + L = int(seqlens[b]) + assert _rel(o_hb[b, :L], o_ref[b, :L]) <= 0.02 + assert _rel(s_hb[b], s_ref[b]) <= 0.02 + + +@CUDA +def test_negctrl_head_band_swap(): + # discriminating: head-batch must match the correct band layout and DIFFER from a + # within-group V-head swap. gqa_mod can't express this (it only re-routes the K/Q source). + from flash_qla import recurrent_gated_delta_rule + B, D, Hk, Hv = 1, 8, 2, 8 # grp=4 + q, k, v, g, beta = _mk(B, D, Hk, Hv) # distinct per-head gate (mandatory) + o_ok, _ = decode_recur(q, k, v, g, beta, scale=128 ** -0.5) + o_swap, _ = decode_recur(q, k, v, g, beta, scale=128 ** -0.5, band_perm=True) + o_hb, _ = recurrent_gated_delta_rule(q, k, v, g, beta, scale=128 ** -0.5, head_batch=True) + assert _rel(o_hb, o_ok) <= 0.02 + assert _rel(o_hb, o_swap) > 0.2 # must NOT match a swapped-band layout + + +@CUDA +def test_head_batch_rejects_unsupported_grp(): + # grp not in {2,4} when forced must raise (thread/register cap), not silently mis-run. + from flash_qla import recurrent_gated_delta_rule + q, k, v, g, beta = _mk(1, 4, 1, 8) # grp=8 + with pytest.raises(AssertionError): + recurrent_gated_delta_rule(q, k, v, g, beta, scale=128 ** -0.5, head_batch=True) diff --git a/tests/test_prepass_gdr.py b/tests/test_prepass_gdr.py new file mode 100644 index 00000000..fa1824a9 --- /dev/null +++ b/tests/test_prepass_gdr.py @@ -0,0 +1,153 @@ +# tests/test_prepass_gdr.py +# H1: the gating + qk-l2norm PRE-PASS kernel must produce g/beta/q_n/k_n bit-equivalent (bf16 +# noise) to the reference gdn_sigmoid_gate + l2norm that the host-gated main verify kernel +# consumes. This is the correctness gate for the dedup pre-pass (replaces the in-hot-loop +# recompute of variant A). +import pytest +import torch + +from flash_qla.utils import l2norm +from flash_qla.ops.gated_delta_rule.fused_recurrent import gdn_sigmoid_gate + +CUDA = pytest.mark.skipif(not torch.cuda.is_available(), reason="needs CUDA") + + +def _rel(a, b): + return ((a.float() - b.float()).abs().max() / b.float().abs().max().clamp_min(1e-6)).item() + + +@CUDA +@pytest.mark.parametrize("Hk,Hv", [(8, 8), (2, 8), (16, 32), (4, 16)]) +@pytest.mark.parametrize("neg", [False, True]) +def test_prepass_matches_reference(Hk, Hv, neg): + from flash_qla.ops.gated_delta_rule.fused_recurrent.hopper.fused_recurrent_verify import ( + fused_recurrent_gdr_verify_prepass, + ) + N, D = 4, 5 + total = N * D + torch.manual_seed(7) + q = torch.randn(1, total, Hk, 128, device="cuda", dtype=torch.bfloat16) + k = torch.randn(1, total, Hk, 128, device="cuda", dtype=torch.bfloat16) + a = torch.randn(1, total, Hv, device="cuda", dtype=torch.bfloat16) + b = torch.randn(1, total, Hv, device="cuda", dtype=torch.bfloat16) + A_log = torch.randn(Hv, device="cuda").abs().log() + dt_bias = torch.randn(Hv, device="cuda") + + q_n, k_n, g, beta = fused_recurrent_gdr_verify_prepass( + q, k, a, b, A_log, dt_bias, allow_neg_eigval=neg) + g_ref, beta_ref = gdn_sigmoid_gate(A_log, a, dt_bias, b, allow_neg_eigval=neg) + + assert _rel(q_n, l2norm(q)) <= 0.02, "q l2norm mismatch" + assert _rel(k_n, l2norm(k)) <= 0.02, "k l2norm mismatch" + assert _rel(g, g_ref) <= 0.02, "g (raw log-decay) mismatch" + assert _rel(beta, beta_ref) <= 0.02, "beta mismatch" + + +@CUDA +def test_prepass_negctrl_raw_g_not_exp2(): + # discriminating: g must be the RAW log-decay (g<=0), NOT pre-exp2'd (which would be in (0,1]). + # The main kernel applies decay=exp2(g*1.442695); a pre-exp2'd g would be silently wrong. + from flash_qla.ops.gated_delta_rule.fused_recurrent.hopper.fused_recurrent_verify import ( + fused_recurrent_gdr_verify_prepass, + ) + N, D, Hk, Hv = 3, 4, 4, 8 + total = N * D + torch.manual_seed(11) + q = torch.randn(1, total, Hk, 128, device="cuda", dtype=torch.bfloat16) + k = torch.randn(1, total, Hk, 128, device="cuda", dtype=torch.bfloat16) + a = torch.randn(1, total, Hv, device="cuda", dtype=torch.bfloat16) + b = torch.randn(1, total, Hv, device="cuda", dtype=torch.bfloat16) + A_log = torch.randn(Hv, device="cuda").abs().log() + dt_bias = torch.randn(Hv, device="cuda") + _, _, g, _ = fused_recurrent_gdr_verify_prepass(q, k, a, b, A_log, dt_bias) + assert (g <= 1e-4).all(), "g must be raw log-decay (<=0), not pre-exp2'd" + + +@CUDA +@pytest.mark.parametrize("Hk,Hv", [(16, 32), (2, 8)]) +def test_prepass_plus_hostgated_matches_in_kernel_gated(Hk, Hv): + # end-to-end: prepass + host-gated main kernel must match the in-kernel-gated kernel on the + # SAME raw inputs (both compute the identical recurrence; only WHERE gating/l2norm happen + # differs). Within in-kernel-gating tolerance (0.03). + from flash_qla.ops.gated_delta_rule.fused_recurrent.hopper.fused_recurrent_verify import ( + fused_recurrent_gdr_verify_prepass, + fused_recurrent_gdr_verify_fwd, + fused_recurrent_gdr_verify_gated_fwd, + ) + N, D = 4, 6 + total = N * D + torch.manual_seed(13) + q = torch.randn(1, total, Hk, 128, device="cuda", dtype=torch.bfloat16) + k = torch.randn(1, total, Hk, 128, device="cuda", dtype=torch.bfloat16) + v = torch.randn(1, total, Hv, 128, device="cuda", dtype=torch.bfloat16) + a = torch.randn(1, total, Hv, device="cuda", dtype=torch.bfloat16) + b = torch.randn(1, total, Hv, device="cuda", dtype=torch.bfloat16) + A_log = torch.randn(Hv, device="cuda").abs().log() + dt_bias = torch.randn(Hv, device="cuda") + pool = torch.randn(N, Hv, 128, 128, device="cuda", dtype=torch.bfloat16) + cu = torch.arange(0, total + 1, D, dtype=torch.int32, device="cuda") + si = torch.arange(N, dtype=torch.int32, device="cuda") + ci = torch.arange(N, dtype=torch.int32, device="cuda") + ibuf_a = torch.zeros(N + 1, D, Hv, 128, 128, device="cuda", dtype=torch.bfloat16) + ibuf_b = torch.zeros(N + 1, D, Hv, 128, 128, device="cuda", dtype=torch.bfloat16) + o_gated = torch.empty(1, total, Hv, 128, device="cuda", dtype=torch.bfloat16) + o_pp = torch.empty(1, total, Hv, 128, device="cuda", dtype=torch.bfloat16) + + # in-kernel-gated (the current path) + fused_recurrent_gdr_verify_gated_fwd( + q, k, v, a, b, A_log, dt_bias, pool, si, cu, ibuf_a, ci, o_gated, disable_state_update=True) + # prepass + host-gated (the H1 path) + q_n, k_n, g, beta = fused_recurrent_gdr_verify_prepass(q, k, a, b, A_log, dt_bias) + fused_recurrent_gdr_verify_fwd( + q_n, k_n, v, g, beta, pool, si, cu, ibuf_b, ci, o_pp, disable_state_update=True) + + assert _rel(o_pp, o_gated) <= 0.03, "prepass+host-gated o must match in-kernel-gated o" + assert _rel(ibuf_b, ibuf_a) <= 0.03, "prepass+host-gated ibuf must match in-kernel-gated ibuf" + + +@CUDA +def test_prepass_path_cuda_graph(): + # the auto-prepass path (large-batch regime) must be CUDA-graph capturable: the persistent + # scratch is warmed by the eager warmup runs, so capture sees no allocation, and replay + # reuses the same buffers -> must match eager bit-for-bit. + from flash_qla import recurrent_gated_delta_rule_verify + N, D, Hk, Hv = 64, 4, 8, 8 # N*Hv*n_vt comfortably saturates SMs -> should_use_prepass True + total = N * D + torch.manual_seed(5) + q = torch.randn(1, total, Hk, 128, device="cuda", dtype=torch.bfloat16) + k = torch.randn(1, total, Hk, 128, device="cuda", dtype=torch.bfloat16) + v = torch.randn(1, total, Hv, 128, device="cuda", dtype=torch.bfloat16) + a = torch.randn(1, total, Hv, device="cuda", dtype=torch.bfloat16) + b = torch.randn(1, total, Hv, device="cuda", dtype=torch.bfloat16) + A_log = torch.randn(Hv, device="cuda").abs().log() + dt_bias = torch.randn(Hv, device="cuda") + pool = torch.randn(N, Hv, 128, 128, device="cuda", dtype=torch.bfloat16) + cu = torch.arange(0, total + 1, D, dtype=torch.int32, device="cuda") + si = torch.arange(N, dtype=torch.int32, device="cuda") + ci = torch.arange(N, dtype=torch.int32, device="cuda") + ibuf = torch.zeros(N + 1, D, Hv, 128, 128, device="cuda", dtype=torch.bfloat16) + o = torch.empty(1, total, Hv, 128, device="cuda", dtype=torch.bfloat16) + + def run(): + recurrent_gated_delta_rule_verify( + A_log, a, dt_bias, q, k, v, b, pool, si, cu, ibuf, ci, + o=o, fuse_gating=True, prepass=True, disable_state_update=True) + + s = torch.cuda.Stream() + s.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(s): + for _ in range(3): # warmup: JIT-compile prepass+main and allocate persistent scratch + run() + torch.cuda.current_stream().wait_stream(s) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): # would raise if the prepass allocated/synced inside capture + run() + graph.replay() + torch.cuda.synchronize() + o_graph = o.clone() + + o_eager = torch.empty_like(o) + recurrent_gated_delta_rule_verify( + A_log, a, dt_bias, q, k, v, b, pool, si, cu, ibuf, ci, + o=o_eager, fuse_gating=True, prepass=True, disable_state_update=True) + assert torch.equal(o_graph, o_eager), "prepass-path graph replay must match eager" diff --git a/tests/test_verify_gdr.py b/tests/test_verify_gdr.py new file mode 100644 index 00000000..a72f41bc --- /dev/null +++ b/tests/test_verify_gdr.py @@ -0,0 +1,243 @@ +# tests/test_verify_gdr.py +import pytest +import torch + +from ref_gdr import verify_ref +from flash_qla.ops.gated_delta_rule.fused_recurrent.hopper.fused_recurrent_verify import ( + fused_recurrent_gdr_verify_fwd, +) +from flash_qla.utils import l2norm + +CUDA = pytest.mark.skipif(not torch.cuda.is_available(), reason="needs CUDA") + + +def _mk(N, D, Hk, Hv, num_slots, seed=0, ragged=False, distinct_gate=False): + torch.manual_seed(seed) + Ls = (torch.randint(1, D + 1, (N,)) if ragged else torch.full((N,), D)).to(torch.int32) + cu = torch.zeros(N + 1, dtype=torch.int32) + cu[1:] = torch.cumsum(Ls, 0) + total = int(cu[-1]) + q = l2norm(torch.randn(1, total, Hk, 128, device="cuda", dtype=torch.bfloat16)) + k = l2norm(torch.randn(1, total, Hk, 128, device="cuda", dtype=torch.bfloat16)) + v = torch.randn(1, total, Hv, 128, device="cuda", dtype=torch.bfloat16) + if distinct_gate: # per-head distinct decay (catches a V-major transpose bug) + g = (torch.nn.functional.logsigmoid(torch.randn(1, total, Hv, device="cuda")) / 16 + * (1 + torch.arange(Hv, device="cuda").float()[None, None, :])) + else: + g = torch.nn.functional.logsigmoid(torch.randn(1, total, Hv, device="cuda")) / 16 + beta = torch.randn(1, total, Hv, device="cuda").sigmoid() + pool = torch.randn(num_slots, Hv, 128, 128, device="cuda", dtype=torch.bfloat16) # V-major + ibuf = torch.zeros(N + 1, D, Hv, 128, 128, device="cuda", dtype=torch.bfloat16) + si = torch.arange(N, dtype=torch.int32, device="cuda") + ci = torch.arange(N, dtype=torch.int32, device="cuda") + return q, k, v, g, beta, pool, si, cu.cuda(), ibuf, ci + + +def _rel(a, b): + return ((a.float() - b.float()).abs().max() / b.float().abs().max().clamp_min(1e-6)).item() + + +@CUDA +@pytest.mark.parametrize("ragged", [False, True]) +@pytest.mark.parametrize("Hk,Hv", [(8, 8), (2, 8)]) +def test_verify_nocommit(ragged, Hk, Hv): + N, D = 4, 4 + q, k, v, g, beta, pool, si, cu, ibuf, ci = _mk(N, D, Hk, Hv, num_slots=N, ragged=ragged) + o = torch.empty(1, q.shape[1], Hv, 128, device="cuda", dtype=torch.bfloat16) + pool0 = pool.clone() + o_ref, pool_ref, ibuf_ref = verify_ref(q, k, v, g, beta, pool, si, cu, ibuf, ci, disable_state_update=True) + fused_recurrent_gdr_verify_fwd(q, k, v, g, beta, pool, si, cu, ibuf, ci, o, disable_state_update=True) + assert _rel(o, o_ref) <= 0.02 + assert _rel(ibuf, ibuf_ref) <= 0.02 + assert torch.equal(pool, pool0), "no-commit must not touch the pool" + + +@CUDA +def test_verify_commit(): + N, D, H = 4, 4, 8 + q, k, v, g, beta, pool, si, cu, ibuf, ci = _mk(N, D, H, H, num_slots=N) + o = torch.empty(1, q.shape[1], H, 128, device="cuda", dtype=torch.bfloat16) + pool_k = pool.clone() + o_ref, pool_ref, ibuf_ref = verify_ref(q, k, v, g, beta, pool, si, cu, ibuf, ci, disable_state_update=False) + fused_recurrent_gdr_verify_fwd(q, k, v, g, beta, pool_k, si, cu, ibuf, ci, o, disable_state_update=False) + assert _rel(o, o_ref) <= 0.02 + assert _rel(pool_k, pool_ref) <= 0.02, "committed final state mismatch" + + +@CUDA +def test_verify_distinct_gate_vmajor(): + # per-head-distinct gates: a wrong V-major store would diverge from the reference + N, D, H = 4, 4, 8 + q, k, v, g, beta, pool, si, cu, ibuf, ci = _mk(N, D, H, H, num_slots=N, distinct_gate=True) + o = torch.empty(1, q.shape[1], H, 128, device="cuda", dtype=torch.bfloat16) + o_ref, pool_ref, ibuf_ref = verify_ref(q, k, v, g, beta, pool, si, cu, ibuf, ci, disable_state_update=True) + fused_recurrent_gdr_verify_fwd(q, k, v, g, beta, pool, si, cu, ibuf, ci, o, disable_state_update=True) + assert _rel(o, o_ref) <= 0.02 + assert _rel(ibuf, ibuf_ref) <= 0.02 + + +@CUDA +def test_verify_wrapper_gating(): + # high-level wrapper: host-side sigmoid gating + l2norm must match a hand-built reference + from flash_qla import recurrent_gated_delta_rule_verify + from flash_qla.ops.gated_delta_rule.fused_recurrent import gdn_sigmoid_gate + + N, D, H = 4, 4, 8 + total = N * D + torch.manual_seed(1) + q = torch.randn(1, total, H, 128, device="cuda", dtype=torch.bfloat16) + k = torch.randn(1, total, H, 128, device="cuda", dtype=torch.bfloat16) + v = torch.randn(1, total, H, 128, device="cuda", dtype=torch.bfloat16) + a = torch.randn(1, total, H, device="cuda", dtype=torch.bfloat16) + b = torch.randn(1, total, H, device="cuda", dtype=torch.bfloat16) + A_log = torch.randn(H, device="cuda").abs().log() + dt_bias = torch.randn(H, device="cuda") + pool = torch.randn(N, H, 128, 128, device="cuda", dtype=torch.bfloat16) + cu = torch.arange(0, total + 1, D, dtype=torch.int32, device="cuda") + si = torch.arange(N, dtype=torch.int32, device="cuda") + ci = torch.arange(N, dtype=torch.int32, device="cuda") + ibuf = torch.zeros(N + 1, D, H, 128, 128, device="cuda", dtype=torch.bfloat16) + + g_ref, beta_ref = gdn_sigmoid_gate(A_log, a, dt_bias, b) + o_ref, _, ibuf_ref = verify_ref( + l2norm(q), l2norm(k), v, g_ref, beta_ref, pool, si, cu, ibuf, ci, disable_state_update=True) + o = recurrent_gated_delta_rule_verify( + A_log, a, dt_bias, q, k, v, b, pool, si, cu, ibuf, ci, disable_state_update=True) + assert _rel(o, o_ref) <= 0.02 + assert _rel(ibuf, ibuf_ref) <= 0.02 + + +@CUDA +@pytest.mark.parametrize("Hk,Hv", [(8, 8), (2, 8)]) +def test_verify_in_kernel_gating(Hk, Hv): + # in-kernel g/beta/l2norm must match the host-gating reference on the same raw inputs + from flash_qla.ops.gated_delta_rule.fused_recurrent import gdn_sigmoid_gate + from flash_qla.ops.gated_delta_rule.fused_recurrent.hopper.fused_recurrent_verify import ( + fused_recurrent_gdr_verify_gated_fwd, + ) + + N, D = 4, 4 + total = N * D + torch.manual_seed(2) + q = torch.randn(1, total, Hk, 128, device="cuda", dtype=torch.bfloat16) + k = torch.randn(1, total, Hk, 128, device="cuda", dtype=torch.bfloat16) + v = torch.randn(1, total, Hv, 128, device="cuda", dtype=torch.bfloat16) + a = torch.randn(1, total, Hv, device="cuda", dtype=torch.bfloat16) + b = torch.randn(1, total, Hv, device="cuda", dtype=torch.bfloat16) + A_log = torch.randn(Hv, device="cuda").abs().log() + dt_bias = torch.randn(Hv, device="cuda") + pool = torch.randn(N, Hv, 128, 128, device="cuda", dtype=torch.bfloat16) + cu = torch.arange(0, total + 1, D, dtype=torch.int32, device="cuda") + si = torch.arange(N, dtype=torch.int32, device="cuda") + ci = torch.arange(N, dtype=torch.int32, device="cuda") + ibuf = torch.zeros(N + 1, D, Hv, 128, 128, device="cuda", dtype=torch.bfloat16) + o = torch.empty(1, total, Hv, 128, device="cuda", dtype=torch.bfloat16) + + g_ref, beta_ref = gdn_sigmoid_gate(A_log, a, dt_bias, b) + o_ref, _, ibuf_ref = verify_ref( + l2norm(q), l2norm(k), v, g_ref, beta_ref, pool, si, cu, ibuf, ci, disable_state_update=True) + fused_recurrent_gdr_verify_gated_fwd( + q, k, v, a, b, A_log, dt_bias, pool, si, cu, ibuf, ci, o, disable_state_update=True) + assert _rel(o, o_ref) <= 0.03 + assert _rel(ibuf, ibuf_ref) <= 0.03 + + +@CUDA +def test_verify_cuda_graph(): + # the low-level entry must be CUDA-graph capturable (no host sync / no alloc) and replay correctly + N, D, H = 4, 4, 8 + q, k, v, g, beta, pool, si, cu, ibuf, ci = _mk(N, D, H, H, num_slots=N) + o = torch.empty(1, q.shape[1], H, 128, device="cuda", dtype=torch.bfloat16) + + def run(): + fused_recurrent_gdr_verify_fwd(q, k, v, g, beta, pool, si, cu, ibuf, ci, o, disable_state_update=True) + + s = torch.cuda.Stream() + s.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(s): + for _ in range(3): + run() + torch.cuda.current_stream().wait_stream(s) + + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + run() + graph.replay() + torch.cuda.synchronize() + + o_eager = torch.empty_like(o) + fused_recurrent_gdr_verify_fwd(q, k, v, g, beta, pool, si, cu, ibuf, ci, o_eager, disable_state_update=True) + assert torch.equal(o, o_eager), "graph replay must match eager" + + +@CUDA +def test_verify_negctrl_vmajor_transpose(): + # discriminating: the V-major ibuf must DIFFER from a transposed-major reference + N, D, H = 4, 4, 8 + q, k, v, g, beta, pool, si, cu, ibuf, ci = _mk(N, D, H, H, num_slots=N, distinct_gate=True) + o = torch.empty(1, q.shape[1], H, 128, device="cuda", dtype=torch.bfloat16) + o_ref, _, ibuf_ref = verify_ref(q, k, v, g, beta, pool, si, cu, ibuf, ci, disable_state_update=True) + fused_recurrent_gdr_verify_fwd(q, k, v, g, beta, pool, si, cu, ibuf, ci, o, disable_state_update=True) + assert _rel(ibuf, ibuf_ref) <= 0.02 + ibuf_wrong = ibuf_ref.transpose(-1, -2).contiguous() # K/V swapped (the silent-bug case) + assert _rel(ibuf, ibuf_wrong) > 0.2 # the test discriminates a wrong major-order + + +@CUDA +def test_verify_gated_cuda_graph(): + # the fully in-kernel gated path (no PyTorch gating/l2norm) is the capture-safe SGLang entry + from flash_qla.ops.gated_delta_rule.fused_recurrent.hopper.fused_recurrent_verify import ( + fused_recurrent_gdr_verify_gated_fwd, + ) + N, D, H = 4, 4, 8 + total = N * D + torch.manual_seed(3) + q = torch.randn(1, total, H, 128, device="cuda", dtype=torch.bfloat16) + k = torch.randn(1, total, H, 128, device="cuda", dtype=torch.bfloat16) + v = torch.randn(1, total, H, 128, device="cuda", dtype=torch.bfloat16) + a = torch.randn(1, total, H, device="cuda", dtype=torch.bfloat16) + b = torch.randn(1, total, H, device="cuda", dtype=torch.bfloat16) + A_log = torch.randn(H, device="cuda").abs().log() + dt_bias = torch.randn(H, device="cuda") + pool = torch.randn(N, H, 128, 128, device="cuda", dtype=torch.bfloat16) + cu = torch.arange(0, total + 1, D, dtype=torch.int32, device="cuda") + si = torch.arange(N, dtype=torch.int32, device="cuda") + ci = torch.arange(N, dtype=torch.int32, device="cuda") + ibuf = torch.zeros(N + 1, D, H, 128, 128, device="cuda", dtype=torch.bfloat16) + o = torch.empty(1, total, H, 128, device="cuda", dtype=torch.bfloat16) + + def run(): + fused_recurrent_gdr_verify_gated_fwd( + q, k, v, a, b, A_log, dt_bias, pool, si, cu, ibuf, ci, o, disable_state_update=True) + + s = torch.cuda.Stream() + s.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(s): + for _ in range(3): + run() + torch.cuda.current_stream().wait_stream(s) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + run() + graph.replay() + torch.cuda.synchronize() + o_eager = torch.empty_like(o) + fused_recurrent_gdr_verify_gated_fwd( + q, k, v, a, b, A_log, dt_bias, pool, si, cu, ibuf, ci, o_eager, disable_state_update=True) + assert torch.equal(o, o_eager), "in-kernel gated graph replay must match eager" + + +@CUDA +def test_verify_skip_slot(): + # a -1 pool slot must skip gather/commit/intermediate writes for that request + N, D, H = 3, 4, 8 + q, k, v, g, beta, pool, si, cu, ibuf, ci = _mk(N, D, H, H, num_slots=N) + si[1] = -1 # request 1 has no pool slot + o = torch.empty(1, q.shape[1], H, 128, device="cuda", dtype=torch.bfloat16) + ibuf0 = ibuf.clone() + o_ref, pool_ref, ibuf_ref = verify_ref(q, k, v, g, beta, pool, si, cu, ibuf, ci, disable_state_update=True) + fused_recurrent_gdr_verify_fwd(q, k, v, g, beta, pool, si, cu, ibuf, ci, o, disable_state_update=True) + # request 1's tokens: o still produced (from zero state), ibuf row left untouched + s0, s1 = int(cu[1]), int(cu[2]) + assert _rel(o[:, s0:s1], o_ref[:, s0:s1]) <= 0.02 + assert torch.equal(ibuf[int(ci[1])], ibuf0[int(ci[1])]), "skipped slot must not write ibuf"