Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
23 commits
Select commit Hold shift + click to select a range
49fb0dd
docs: add GDN decode (fused recurrent) kernel design spec
rifkybujana Jun 14, 2026
e78a0cf
docs: add SGLang verify spec; fix decode-spec rank-1 grounding + ragg…
rifkybujana Jun 15, 2026
8eba13f
docs(verify-spec): scope to DFlash linear chain; drop DDTree/tree sur…
rifkybujana Jun 15, 2026
6bac07a
docs(verify-spec): apply fork-grounded review (3 fixes + resolved dep…
rifkybujana Jun 15, 2026
adeed79
docs(plan): GDN decode feasibility gates + core kernel implementation…
rifkybujana Jun 15, 2026
c368815
feat(decode): gemm-free fused recurrent GDN decode kernel (Phase 1, v…
rifkybujana Jun 15, 2026
2dbe39e
feat(verify): paged V-major SGLang verify kernel (infra A core, valid…
rifkybujana Jun 15, 2026
28b8f38
feat(verify): high-level SGLang verify wrapper (host gating) + CUDA-g…
rifkybujana Jun 15, 2026
521f4f9
feat(verify): in-kernel fused sigmoid gating + qk-l2norm (req #5, val…
rifkybujana Jun 15, 2026
682f016
bench+docs: verify-kernel bandwidth benchmark + implementation-status…
rifkybujana Jun 15, 2026
6ca0c40
test+fix: address code review (negative controls, capture-safety, con…
rifkybujana Jun 15, 2026
75f799c
bench: FlashQLA vs FLA/Triton GDN verify speed comparison (H100)
rifkybujana Jun 15, 2026
f5b8744
perf(decode/verify): autotuned tile -> match/beat FLA at large batch …
rifkybujana Jun 15, 2026
be6da09
docs(verify-spec): record tuned FlashQLA >= FLA benchmark result
rifkybujana Jun 15, 2026
94a10e4
docs: add CLAUDE.md (repo guidance for Claude Code)
rifkybujana Jun 15, 2026
36400e2
docs(specs): reconcile design docs with as-built kernel
rifkybujana Jun 15, 2026
d78764a
feat(decode): head-batched GQA specialization (gemm-free row-stack)
rifkybujana Jun 15, 2026
dfed1cc
perf(verify): H1 gating+l2norm dedup pre-pass (+8-24% large-batch ver…
rifkybujana Jun 15, 2026
110e4d2
perf(decode): block_DV=128 for the final-state path (~2x at all batch…
rifkybujana Jun 15, 2026
27dc255
bench: high-batch FLA vs FlashQLA varA vs H1-prepass (3-way, parity-g…
rifkybujana Jun 15, 2026
69f0ca4
docs+probe: optimization sweep — attribution + recorded nulls
rifkybujana Jun 15, 2026
88027e1
perf(verify): graph-calibrate should_use_prepass (PREPASS_MIN_WORK 30…
rifkybujana Jun 15, 2026
066a214
feat(verify): env-override PREPASS_MIN_WORK for controlled A/B
rifkybujana Jun 16, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions .gitignore
Original file line number Diff line number Diff line change
@@ -1,3 +1,6 @@
__pycache__
*.egg-info
build/

# local modal test harness
_modal/
81 changes: 81 additions & 0 deletions CLAUDE.md
Original file line number Diff line number Diff line change
@@ -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 <name>` preset that loads `tests/settings/<name>.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).
54 changes: 54 additions & 0 deletions benchmark/bench_head_batch.py
Original file line number Diff line number Diff line change
@@ -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()
150 changes: 150 additions & 0 deletions benchmark/bench_prepass.py
Original file line number Diff line number Diff line change
@@ -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<MIN_T keeps prepass OFF; expect no win ==")
for N in (8, 32, 128):
bench_graph(N, 1, 16, 32)
print("\n== N=1 single-request EAGER sanity (must remain A / loss) ==")
for Hv in (64, 32, 16):
bench(1, 12, max(1, Hv // 4), Hv, "lat")


if __name__ == "__main__":
main()
Loading