Skip to content

Commit e0317d1

Browse files
sdpa fwd (frost): THD templates compile with dynamic batch and head extents -- one kernel per layout class (#1039)
A compile() of an SDPA THD template pinned b, qh, kh and every stride in its fakes, so FlashInfer's serving shapes minted one 2.6 s kernel per (b, qh, kh, strides): 126 forward kernels in FlashInfer's cuDNN attention tests, 4 in the FlashInfer-shaped suite alone. The kernels never needed it: _host reads B / QH / KH from problem_size at run time and the token totals were already sym_int; only the fakes pinned them. compile(dynamic_bhk=True) (THD only) rebinds b, qh, kh -- and a padded LSE's s_max -- to cute.sym_int() right after the cache key is taken, keeps plain ints for the problem_size fake, and gives the metadata and O-descriptor fakes fresh symbols (SymInt has no __add__). A packed declared stride is passed as None so the compact fake derives it from the dynamic extents; a padded LSE in a compact dim order passes that order (lse_padded_order); anything else keeps its static key. The adapter canonicalizes the key (b = qh = kh = 0, lse_padded_rows = 1) whenever the module offers dynamic_bhk, and passes the bound K/V views' strides as None when they are packed. Static still: d, dtypes, masks, the GQA ratio, paged pool strides, dense (non-THD) shapes. Two DSL facts: `and` is staged, so the PackGQA head check nests its constant test outside (`const_expr(CFG.PACK_GQA and q.shape[2] != ...)` tried to const_expr a runtime value); and a const_expr on a dynamic extent is an error, which is why the padded-Stats store selects on the LSE fake's rank. FlashInfer-shaped SDPA suite: 16 passed, 4 kernels compiled (one per d / has_lse / padded combination) where every shape used to mint its own. Review fix folded in: make_fake_compact_tensor's stride_order is per AXIS -- the rank of that axis's stride, 0 = fastest (the DSL reads it as stride_order.index(rank)); the padded-LSE order helper returned the axes sorted by stride, the inverse permutation. FlashInfer's (b, s_max, h) and the (3, 2, 1, 0) default are their own inverses, which is why every existing case passed; 4 of the 6 (b, h, s_max) storage orders compiled to another layout. The helper returns the inverse; a GPU-free test checks all six orders against the strides CuTe derives, and a GPU test runs the (h, s_max, b) order against the contiguous layout row for row. Co-authored-by: Claude Fable 5.1 <noreply@anthropic.com>
1 parent 8ed0a0d commit e0317d1

24 files changed

Lines changed: 620 additions & 114 deletions

‎docs/utilities/python_graph_and_execution_backends.md‎

Lines changed: 22 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -674,6 +674,28 @@ only to decline is why `closed_under` existed.
674674
backend-unlowerable (`serialize()` and `key()` refuse it), and it surfaces as
675675
a graph fact the capability rows gate on.
676676

677+
### One kernel per layout class, not per shape (SDPA THD)
678+
679+
A `compile()` of an SDPA template used to pin batch, head extents and every
680+
stride in its fakes, so FlashInfer's serving shapes minted one 2.6 s kernel per
681+
`(b, qh, kh, strides)`: 126 forward kernels in FlashInfer's cuDNN attention
682+
tests. The kernels never needed that — `_host` reads `B / QH / KH` from
683+
`problem_size` at run time and the THD token totals were already `sym_int`.
684+
Under `compile(dynamic_bhk=True)` (THD only) the template rebinds `b`, `qh`,
685+
`kh` (and a padded LSE's `s_max`) to `cute.sym_int()` right after the cache
686+
key is taken, keeps plain ints for the `problem_size` fake, and gives the
687+
metadata / O-descriptor fakes fresh symbols (`SymInt` has no `__add__`); a
688+
packed declared stride is passed as `None` so the compact fake derives it from
689+
the dynamic extents, a padded LSE in a compact dim order passes that order
690+
(`lse_padded_order`), and anything else keeps its static key. The adapter
691+
canonicalizes the key (`b = qh = kh = 0`, `lse_padded_rows = 1`) when the
692+
module offers `dynamic_bhk`. What stays static: `d`, dtypes, masks, the GQA
693+
ratio (`CFG.QH_PER_KH`), paged pool strides, dense (non-THD) shapes. Two DSL
694+
facts shaped this: `and` is staged, so `const_expr(CFG.PACK_GQA and
695+
q.shape[2] != ...)` must nest its constant test outside; a `const_expr` on a
696+
dynamic extent or stride is an error, which is why the padded-Stats store
697+
selects on the fake's RANK (rank-4) and not on `shape[0] > 1`.
698+
677699
### Accept means run
678700

679701
For a python plan, `check_support()` accepted ⇒ `build_plans()` and

‎python/cudnn/sdpa/frost/SUPPORT_MATRIX_TRACKER.md‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -546,4 +546,4 @@ only — head dim innermost, then heads, then tokens (`graph_analyzer.packed_lay
546546
The batch stride is not gated: every sequence base comes from the ragged offsets and
547547
the lowering binds the batch axis at extent 1, so its declared value is never read.
548548
FlashInfer declares it equal to the token stride (`h * d`), which the previous
549-
all-four-axes check refused at `b > 1`. Stats under THD are written in the caller's declared layout: packed `(T, H)` rows, head-major `(1, QH, head_stride)`, or -- a Stats tensor **without** ragged offsets -- the per-batch padded form -- the graph's logical `[b, h, s_max, 1]` Stats view over FlashInfer's physical, contiguous `(b, s_max, h)` `return_lse` buffer (declared strides `[s_max*h, 1, h, 1]`; the adapter rebuilds the view with `as_strided`, nothing is allocated in the logical order), stored per batch through the declared strides on every THD row (SM100 / SM107 / SM120, `Capabilities.thd_padded_stats`); the adapter seeds that buffer with `-inf` on the launch stream first, so the rows past a sequence's length read the backend's value.
549+
all-four-axes check refused at `b > 1`. Stats under THD are written in the caller's declared layout: packed `(T, H)` rows, head-major `(1, QH, head_stride)`, or -- a Stats tensor **without** ragged offsets -- the per-batch padded form -- the graph's logical `[b, h, s_max, 1]` Stats view over FlashInfer's physical, contiguous `(b, s_max, h)` `return_lse` buffer (declared strides `[s_max*h, 1, h, 1]`; the adapter rebuilds the view with `as_strided`, nothing is allocated in the logical order), stored per batch through the declared strides on every THD row (SM100 / SM107 / SM120, `Capabilities.thd_padded_stats`); the adapter seeds that buffer with `-inf` on the launch stream first, so the rows past a sequence's length read the backend's value. On SM100 / SM107 the THD templates compile with DYNAMIC batch and head extents (`compile(dynamic_bhk=True)`): one artifact per layout class (d, dtypes, masks, GQA ratio, packed vs declared strides), not per shape.

‎python/cudnn/sdpa/fwd/api_dsl.py‎

Lines changed: 79 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -2102,6 +2102,45 @@ def execute(
21022102
O_view.copy_(O_scratch)
21032103
self._logger.debug("execute completed")
21042104

2105+
def _thd_dynamic_bhk(self) -> bool:
2106+
"""Whether this THD plan compiles its batch and head extents dynamic:
2107+
the kernel module offers it, and a padded LSE (if any) is in a compact
2108+
order the fake can express. Decided once per plan: this is asked on
2109+
every execute (the compile kwargs are rebuilt per call) and
2110+
``inspect.signature`` alone cost 60 us a call."""
2111+
cached = getattr(self, "_thd_dynamic_bhk_cached", None)
2112+
if cached is not None:
2113+
return cached
2114+
import inspect
2115+
2116+
dyn = bool(self.thd and self._k_mod is not None and "dynamic_bhk" in inspect.signature(self._k_mod.compile).parameters)
2117+
if dyn and self.lse_desc is not None and self.thd_stats_padded and self._thd_padded_lse_order() is None:
2118+
dyn = False
2119+
self._thd_dynamic_bhk_cached = dyn
2120+
return dyn
2121+
2122+
def _thd_padded_lse_order(self):
2123+
"""CuTe's ``stride_order`` for a padded (b, h, s_max, 1) LSE whose declared
2124+
strides are compact in SOME dim order, else None: per AXIS, the rank of
2125+
that axis's stride (0 = fastest) -- ``make_fake_compact_tensor`` reads it
2126+
as ``stride_order.index(rank)``. FlashInfer's (b, s_max, h) is (3, 1, 2, 0).
2127+
Memoized: asked per execute."""
2128+
if self._lse_stride is None:
2129+
return None
2130+
cached = getattr(self, "_thd_padded_lse_order_cached", ())
2131+
if cached != ():
2132+
return cached
2133+
shape = (self.batch_size, self.h_q, self.s_q_max, 1)
2134+
st = (*self._lse_stride, 1)
2135+
axes = sorted(range(4), key=lambda i: (st[i], -i)) # fastest axis first
2136+
expect, acc = [0] * 4, 1
2137+
for i in axes:
2138+
expect[i] = acc
2139+
acc *= shape[i]
2140+
compact = all(shape[i] == 1 or expect[i] == st[i] for i in range(4))
2141+
self._thd_padded_lse_order_cached = tuple(axes.index(i) for i in range(4)) if compact else None
2142+
return self._thd_padded_lse_order_cached
2143+
21052144
def _thd_compile_kwargs(self) -> dict:
21062145
"""The THD compile key — PLAN-TIME-ONLY by contract (issue #552).
21072146
@@ -2119,10 +2158,16 @@ def _key(desc):
21192158
return (0, ts, hs, es)
21202159

21212160
has_lse = self.lse_desc is not None
2161+
# Dynamic batch / head extents (the kernels read them at run time; only the
2162+
# fakes pinned them): one artifact per LAYOUT class. A packed declared
2163+
# stride derives from the dynamic extents and is passed as None; a
2164+
# padded LSE in a compact order passes that order. Anything else keeps
2165+
# the static key for that operand.
2166+
dyn = self._thd_dynamic_bhk()
21222167
kwargs = dict(
2123-
b=self.batch_size,
2124-
qh=self.h_q,
2125-
kh=self.h_kv,
2168+
b=0 if dyn else self.batch_size,
2169+
qh=0 if dyn else self.h_q,
2170+
kh=0 if dyn else self.h_kv,
21262171
# The Stats layout is a per-shape specialization (like d_qk/d_v):
21272172
# has_lse=False compiles the store out; token-major binds the
21282173
# packed rank-2 (T, H) view; head-major carries the caller-declared
@@ -2131,26 +2176,41 @@ def _key(desc):
21312176
lse_head_major=has_lse and self.thd_stats_head_major,
21322177
lse_head_stride=(self.thd_stats_head_stride if (has_lse and self.thd_stats_head_major) else 0),
21332178
# padded per-batch Stats: the (B, QH, s_max) fake in the declared strides
2134-
lse_padded_rows=(self.s_q_max if (has_lse and self.thd_stats_padded) else 0),
2135-
lse_stride=(self._lse_stride if (has_lse and self.thd_stats_padded) else None),
2179+
lse_padded_rows=((1 if dyn else self.s_q_max) if (has_lse and self.thd_stats_padded) else 0),
2180+
lse_stride=(self._lse_stride if (has_lse and self.thd_stats_padded and not (dyn and self._thd_padded_lse_order() is not None)) else None),
21362181
)
2182+
if dyn:
2183+
kwargs["dynamic_bhk"] = True
2184+
if has_lse and self.thd_stats_padded and self._thd_padded_lse_order() is not None:
2185+
kwargs["lse_padded_order"] = self._thd_padded_lse_order()
21372186
if self._fp8:
21382187
# FP8/MXFP8 THD serves only native packed contracts. Flavor
21392188
# selection chooses the exact kernel module, so no stride or
21402189
# head-dim entries are needed in this per-module compile key.
21412190
return kwargs
2191+
2192+
def _stride_key(desc):
2193+
key, packed = self._thd_declared(desc)
2194+
return None if (dyn and packed) else (0, *key)
2195+
21422196
kwargs.update(
21432197
d_qk=self.head_dim_qk,
21442198
d_v=self.head_dim_v,
2145-
q_stride=_key(self.q_desc),
2146-
# Paged pools bind as declared (strides in the kernel's
2147-
# [num_pages, page_size, H_kv, D] order); no token-stride key.
2148-
k_stride=self._paged_pool_stride(self.k_desc) if self.paged else _key(self.k_desc),
2149-
v_stride=self._paged_pool_stride(self.v_desc) if self.paged else _key(self.v_desc),
2150-
o_stride=_key(self.o_desc),
2199+
q_stride=_stride_key(self.q_desc),
2200+
k_stride=self._paged_pool_stride(self.k_desc) if self.paged else _stride_key(self.k_desc),
2201+
v_stride=self._paged_pool_stride(self.v_desc) if self.paged else _stride_key(self.v_desc),
2202+
o_stride=_stride_key(self.o_desc),
21512203
)
21522204
if self.paged:
2153-
kwargs.update(block_table_stride=self.paged_table_stride, block_table_v_stride=self.paged_table_v_stride)
2205+
# A compact (max_pages, 1) table derives from the dynamic extents under
2206+
# dynamic_bhk (max_pages varies per graph; it is the paged path's last
2207+
# per-shape key); a declared non-compact table keeps its strides.
2208+
n_pages = -(-int(self.paged_max_seq_len_kv) // int(self.paged_page_size)) if self.paged_max_seq_len_kv else None
2209+
2210+
def _table_key(declared):
2211+
return None if (dyn and declared is not None and n_pages is not None and tuple(declared) == (n_pages, 1)) else declared
2212+
2213+
kwargs.update(block_table_stride=_table_key(self.paged_table_stride), block_table_v_stride=_table_key(self.paged_table_v_stride))
21542214
return kwargs
21552215

21562216
def _thd_unit_envelope(self) -> int:
@@ -2392,7 +2452,13 @@ def _execute_thd(
23922452
# key (a runtime value the kernel rebuilds symbolically).
23932453
kwargs = self._thd_compile_kwargs()
23942454
if not self.paged:
2395-
kwargs.update(k_stride=(0, *pack.K.stride()[1:]), v_stride=(0, *pack.V.stride()[1:]))
2455+
# the bound views' strides; packed ones derive from the dynamic extents
2456+
# under dynamic_bhk and stay out of the key
2457+
def _view_key(t, h):
2458+
st = tuple(int(x) for x in t.stride()[1:])
2459+
return None if (kwargs.get("dynamic_bhk") and st == (h * t.shape[-1], t.shape[-1], 1)) else (0, *st)
2460+
2461+
kwargs.update(k_stride=_view_key(pack.K, self.h_kv), v_stride=_view_key(pack.V, self.h_kv))
23962462
fn = self._k_mod.compile(**kwargs)
23972463
self._seed_padded_lse(LSE, current_stream)
23982464
# Paged pools: the block tables follow the THD length slots in the ABI.

‎python/cudnn/sdpa/fwd/kernels/sm100/prefill_d128_f16.py‎

Lines changed: 25 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -2393,8 +2393,11 @@ def _host(
23932393
# K is split along seq under cga2 — box rows are per-CTA. V is split along d_v
23942394
# (plumbed via num_iters//CTA_MMA). O's TMA box inner dim follows O's swizzle, NOT V's.
23952395
_O_GRANU_ELEMS = CFG.O_SWZ_BYTES // CFG.BPE
2396-
if cutlass.const_expr(CFG.PACK_GQA and q_tensor.shape[2] != k_tensor.shape[2] * CFG.QH_PER_KH):
2397-
raise ValueError(f"CFG.QH_PER_KH ({CFG.QH_PER_KH}) does not match tensor head extents H_q={q_tensor.shape[2]}, H_kv={k_tensor.shape[2]}")
2396+
if cutlass.const_expr(CFG.PACK_GQA):
2397+
# nested: `and` is staged by the DSL, and under dynamic_bhk the head extents
2398+
# are runtime values -- PackGQA is dense-only, so this branch never sees them
2399+
if cutlass.const_expr(q_tensor.shape[2] != k_tensor.shape[2] * CFG.QH_PER_KH):
2400+
raise ValueError(f"CFG.QH_PER_KH ({CFG.QH_PER_KH}) does not match tensor head extents H_q={q_tensor.shape[2]}, H_kv={k_tensor.shape[2]}")
23982401
qk_box_q = (1, CFG.TILE_M // HEADS_PER_TILE, HEADS_PER_TILE, TMA_QK_GRANU_ELEMS)
23992402
# Paged KV: K/V are [num_pages, page_size, H_kv, D] views of the page pool
24002403
# (batch -> page, seq -> row-in-page) and a tile is a stack of K_BOXES /
@@ -2551,6 +2554,8 @@ def compile( # noqa: A001
25512554
lse_head_major: bool = False,
25522555
lse_head_stride: int = 0,
25532556
lse_padded_rows: int = 0,
2557+
lse_padded_order: tuple = (3, 2, 1, 0),
2558+
dynamic_bhk: bool = False,
25542559
q_stride: Optional[tuple] = None,
25552560
k_stride: Optional[tuple] = None,
25562561
v_stride: Optional[tuple] = None,
@@ -2597,6 +2602,20 @@ def compile( # noqa: A001
25972602
global stride must be a 16-byte multiple; the compact BSHD H-stride is
25982603
d * BPE, so d must be a multiple of 8 at 2 bytes/elem."""
25992604
_cache_key = _template_key(globals(), locals(), "compile")
2605+
_b0, _qh0, _kh0 = b, qh, kh # the problem_size fake: runtime scalars, values immaterial
2606+
if dynamic_bhk:
2607+
# Batch and head extents compile DYNAMIC: one artifact per layout class,
2608+
# not per (b, qh, kh) -- serving shapes vary in all three. The kernel
2609+
# already reads B / QH / KH from problem_size at run time; only the fakes
2610+
# pinned them. A packed stride (None) derives from the dynamic extents;
2611+
# a declared stride stays the fixed number it is.
2612+
if not CFG.THD_VARLEN:
2613+
raise ValueError("dynamic_bhk is THD-only (dense shapes still pin the fakes)")
2614+
b = cute.sym_int(divisibility=1)
2615+
qh = cute.sym_int(divisibility=1)
2616+
kh = cute.sym_int(divisibility=1)
2617+
if lse_padded_rows:
2618+
lse_padded_rows = cute.sym_int(divisibility=1)
26002619
if not (0 < d_qk <= CFG.TILE_K and 0 < d_v <= CFG.TILE_O):
26012620
raise ValueError(f"d128 envelope: need 0 < d_qk <= {CFG.TILE_K} and 0 < d_v <= {CFG.TILE_O}; got ({d_qk}, {d_v})")
26022621
if (d_qk * CFG.BPE) % 16 != 0 or (d_v * CFG.BPE_O) % 16 != 0:
@@ -2696,7 +2715,7 @@ def _fake_table(stride):
26962715
fake_lse = (
26972716
cute.runtime.make_fake_tensor(cutlass.Float32, (b, qh, lse_padded_rows, 1), (*lse_stride, 1), assumed_align=4)
26982717
if lse_stride
2699-
else cute.runtime.make_fake_compact_tensor(cutlass.Float32, (b, qh, lse_padded_rows, 1), stride_order=(3, 2, 1, 0), assumed_align=4)
2718+
else cute.runtime.make_fake_compact_tensor(cutlass.Float32, (b, qh, lse_padded_rows, 1), stride_order=lse_padded_order, assumed_align=4)
27002719
)
27012720
elif lse_head_major:
27022721
# head_stride covering t_q is validated at execute (t_q is a
@@ -2740,7 +2759,7 @@ def _fake_table(stride):
27402759
# seq_kv_lens always part of the ABI; read only when CFG.SEQ_KV_LENS_PRESENT == 1
27412760
# (compile-time fold). THD overloads it as the [seq_kv_lens(B)|cu_q(B+1)|
27422761
# cu_k(B+1)|batch_remap(B)|live|ctr] metadata buffer (length 4B+4).
2743-
_skv_len = (4 * b + 4) if CFG.THD_VARLEN else b
2762+
_skv_len = cute.sym_int(divisibility=1) if dynamic_bhk else ((4 * b + 4) if CFG.THD_VARLEN else b)
27442763
fake_seq_kv_lens = cute.runtime.make_fake_compact_tensor(
27452764
cutlass.Int32,
27462765
(_skv_len,),
@@ -2764,7 +2783,7 @@ def _fake_table(stride):
27642783
# Per-batch O TMA-descriptor array (16 int64 = 128 B each) + 1 pad slot
27652784
# + 2 slots for the packed-total-clamped K/V runtime descriptors the setup
27662785
# kernel writes (issue #624); dummy 1-elem when THD off (never read).
2767-
_odesc_len = (b * _TENSOR_MAP_QWORDS + 3 * _TENSOR_MAP_QWORDS) if CFG.THD_VARLEN else 1
2786+
_odesc_len = cute.sym_int(divisibility=1) if dynamic_bhk else ((b * _TENSOR_MAP_QWORDS + 3 * _TENSOR_MAP_QWORDS) if CFG.THD_VARLEN else 1)
27682787
fake_o_desc = cute.runtime.make_fake_compact_tensor(
27692788
cutlass.Int64,
27702789
(_odesc_len,),
@@ -2796,7 +2815,7 @@ def _fake_table(stride):
27962815
fake_o_desc,
27972816
# THD: the packed totals are runtime values carried by the (dynamic)
27982817
# tensor extents — _host reads them from the views' shapes.
2799-
(b, qh, kh, 0, 0, 0) if CFG.THD_VARLEN else (b, qh, kh, sq, skv, 0),
2818+
(_b0, _qh0, _kh0, 0, 0, 0) if CFG.THD_VARLEN else (_b0, _qh0, _kh0, sq, skv, 0),
28002819
cutlass.Float32(0.0),
28012820
cutlass.Int32(0),
28022821
fake_seq_q_lens,

0 commit comments

Comments
 (0)