You signed in with another tab or window. Reload to refresh your session.You signed out in another tab or window. Reload to refresh your session.You switched accounts on another tab or window. Reload to refresh your session.Dismiss alert
{{ message }}
Repository navigation
Commit e0317d1
Browse filesBrowse the repository at this point in the historyBrowse 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>
Copy file name to clipboardExpand all lines: python/cudnn/sdpa/frost/SUPPORT_MATRIX_TRACKER.md
+1-1Lines changed: 1 addition & 1 deletion
Display the source diff
Display the rich diff
Original file line number
Diff line number
Diff line change
@@ -546,4 +546,4 @@ only — head dim innermost, then heads, then tokens (`graph_analyzer.packed_lay
546
546
The batch stride is not gated: every sequence base comes from the ragged offsets and
547
547
the lowering binds the batch axis at extent 1, so its declared value is never read.
548
548
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.
0 commit comments