SDPA-specific hard rules, on top of the package-wide Rules 1-5 in
../AGENTS.md. Numbered S1, S2, ... so reviews can cite them
without colliding with the package-wide numbering; the list grows — append,
never renumber.
Rule S1 — THD/packed Stats (LSE) must stay packed: token-major or head-major, never dense-padded.
- Consumers read Stats through the same
cu_seqlenpacking as Q/O: TE and Megatron take token-major(T, H)(cuDNN's TH1 recipe) natively, FA-style callers take head-major(H, head_stride). A dense-padded declaration (per-sequence stride) mis-addresses under that packing — it must be rejected at validation time, not silently accepted and mis-read. - Validate by stride, not by a
thd/packedflag: token-major isstride_h == 1 and stride_s == H; head-major isstride_s == 1withstride_hthe declared head stride, which must cover the packed token countT—stride_h >= His not the bound; atT > stride_hthe per-head slices alias.Tis a runtime total, so plan-time classification can only checkstride_s == 1 and stride_h >= 1. In the THD path the packed total is a device value — Rule 3 bans reading it back, sostride_h >= Tis caller contract (stated in the prepared THD binding contract), not something the adapter verifies:as_stridedbounds-checks storage capacity, never overlap. Do not "fix" this with a host-side length read; an in-kernel assert is the only legal detector. Classify withgraph_analyzer.thd_stats_packing(stride_h, stride_s, h_q)— the one classifier the fwd adapters, the bwd probe and the bwd lowering share; never re-implement the stride test inline (Rule 3's "suspect duplicated logic first"). - Covered by
test_fwd_probe_rejects_invalid_stats_metadataand thestats_layout-parametrized THD tests (test_dsl_sm100_thd_statsand siblings) intest/python/sdpa/frost/.
Rule S2 — A change to any FROST SDPA Capabilities row updates
python/cudnn/sdpa/frost/SUPPORT_MATRIX_TRACKER.md in the same commit.
- The matrix is the only place a reader can see what FROST serves across
arch x pass x head-dim x dtype without reading five
engines.pyrows. It is generated by hand from those rows, so it goes stale the moment one changes and nothing else notices. - Applies to every field that changes what a graph gets:
dtypes,out_dtypes,d_shapes/d_pad_multiple/d_envelope_floors/thd_d_shapes, the mask and feature booleans (causal,bottom_right,right_band_widening,swa,padded,padded_stats,sink,thd,cu_seq_len,bias,decode,layouts, ...), and adding or retiring anEngineSpecrow. A change confined to knob domains (tile_ms,sched_policies, ...) does not need a matrix edit — the matrix deliberately does not track knobs. - The matrix distinguishes native (a kernel for this geometry exists)
from envelope-served (the graph rides a larger flavor with TMA
zero-padding, at that flavor's MMA cost). When adding a head dim, say
which one it is — an envelope hit is a perf cliff a user cannot see from
a
Capabilitiesrow.d_pad_multiple=0means exact-only, so no envelope at all. - Reviewing: if the diff touches
fwd/engines.py,bwd/engines.pyorcudnn/engines/manifest.pyandpython/cudnn/sdpa/frost/SUPPORT_MATRIX_TRACKER.mdis untouched, ask why before approving.
Rule S3 — When P is aliased into a score slot's TMEM tail, the warpgroup that stores into columns another warpgroup still has to read must be ordered after that read.
- Some kernels have no spare TMEM and pack P into the tail of the S slot it
was computed from (
P_EVEN_OFF/P_ODD_OFFinside theS_ACC_*column range — e.g.sm100/prefill_d256_fp8.py,P_EVEN_OFF = 96in a 128-col slot). With a single softmax owner that is safe: all of S is in registers before any P store. It is a race the moment a row's keys are split across two softmax warpgroups: the half storing into the aliased columns overwrites the other half's unread S whenever that half lags a tile. GitHub #981 was exactly this — garbage weights on keys 96–111, masked paths only, sporadic and data-dependent (the no-mask fast path is single-owner). - The fix is an explicit release: the reading half arrives an mbarrier
after its
tcgen05.ldhas landed (tcgen05.wait::ldfirst — the softmax paths otherwise never wait on loads explicitly), and the storing half waits on it right before the aliased store (mb_softmax_hi_loadedinsm100/prefill_d256_fp8.py, mirroringmb_softmax_max). Do NOT try to fix it by swapping which half owns which keys: the alpha/stats value is aliased too (STATS_*_OFF == S_ACC_*_OFF, i.e. S column 0), so the half that stores alpha must also be the one that owns key 0 — the original ownership is forced, only the ordering was missing. - The same kernel family can hide the split behind a different role name.
sm100/prefill_d256_mxfp8.pyhas nomb_softmax_max: itsFUSED_CORR_SPLIT_Pschedule (mxfp8 + strict top-left causal) gives keys 64–127 to the correction warpgroup (_fused_p1_step), and half 0 stored P into cols 96–111 with nothing ordering it after that warpgroup's S-hi load. There the release already existed —mb_stat_emptyis arrived right after that load — so the fix is to wait on it before the aliased store instead of at the end of the iteration. The dense split-P schedule of the same kernel is safe only because both halves meet at a 256-thread named barrier (barrier_id=8) after their loads; if that barrier ever moves below the P store, this comes back. - Detector: every kernel with
P_EVEN_OFF: int =inside theS_ACC_*range (grep -l "P_EVEN_OFF: int = " python/cudnn/sdpa/fwd/kernels/*/*.py, today the six d256 kernels) whose config can selectSOFTMAX_WARPGROUPS == 2(grep -n "SOFTMAX_WARPGROUPS=2\|split_p\|fused_corr_split_p" python/cudnn/sdpa/fwd/config_*.py) — regardless of which role runs the second half. For each P store, name the warpgroup that still reads the target columns and the barrier that orders the store after that read; if there is none, it is a bug. Random data can hide it: the lagging half usually lags only when it rescales O (alpha ≠ 1), whichRESCALE_THRESHOLDmakes rare — force it with K scaled bygrowth ** (key // 128)(k_tile_growthintest/python/sdpa/frost/test_sdpa_fwd_mxfp8_sm100.py; the pre-fix mxfp8 kernel then fails every run with inf/NaN O). The symptom in a fuzz suite isref == 0, gpu != 0elements on masked configs withkey mod TILE_Nin the aliased range — recover the leaked keys by matchingO_gpu - O_refagainstVrows.
Rule S4 — Split-K partial Stats keep the combiner's log base.
stats_use_log2changes final Stats only. Partial LSE consumed by a natural-log combine stays natural-log, including direct kernel entry points; applylog2(e)exactly once at the final store.- Check both O and Stats against an independent reference with nonzero logits.
O-only checks miss a wrong Stats base; zero logits miss an unscaled row maximum.
See
test_fp8_graph_stats_use_log2and the split-KV Stats tests.test_sm120_direct_template_stats_basebypasses the adapter: the adapter clears the partial-log2 flag itself, so adapter-only tests cannot detect a missing guard in a directly called template.
Rule S5 — Strided outputs must retain their layout through the final store.
make_array_view(t)[b, s, h, :]returns a row pointer; indexing that pointer bydassumes a unit D stride. Use full indexing (view[b, s, h, d]) when accepting an arbitrary declared D stride, or explicitly require D-contiguous storage. Test padding canaries as well as numerical output; the detector istest_pointer_combine_strided_outputs_and_dead_splits.- TMA alignment checks use each operand's actual element width. FP8 Q/K/V
can produce half or FP8 O; treating O as always two bytes admits strides
aligned to eight elements that are illegal for one-byte O. Keep engine and
adapter admission in agreement; the detector is
test_prepared_fp8_output_stride_uses_output_element_width.
Stats stores obey the same full-indexing rule as O. Nonunit sequence stride
can otherwise leave half the rows unwritten while corrupting padding. The
SM107 detector is test_sm107_fp8_stats_nonunit_row_stride; MXFP8 also checks
rebound padded and batch-inner layouts under CUDA Graph replay.
Block-scaled SF_O uses byte addressing: its fake tensor extent, host geometry
arguments, device parameters, and every intermediate offset product must all
stay Int64. Widen operands before multiplication. The physical detector is
test_block_scaled_sf_plane_stride_above_int32: it writes two live SF planes
separated above 2**32 and checks capture/replay. A wide fake extent alone only
fixes binding; deliberately narrowing the plane stride must fail numerically.
Rule S6 — A kernel feature lands on every arch line's test file, and its other-arch lowerings are smoke-compiled from whatever GPU you have.
- Each FROST fwd arch line has its own test file and marker:
test_sdpa_fwd_fp8_sm100.pyruns underrequires_blackwell(SM 100..119, the Rubin lane included —_D128_ARCHpicks the kernel), the sm120 line underrequires_blackwell_geforce(120..129) intest_sdpa_fwd_fp8_sm120.py, sm80 in its own file. A test added to one file never runs on the other lanes, and_skip_on_rubinis a d192/d256-flavor statement, not a default decorator to copy. Detector:pytest --collect-only -q -k <feature>per file must list the cases (the block-scaled O review found SM120 FP4 declining itself and the Rubin lane skipping the epilogue entirely). - The DSL traces the kernel in Python before any arch-specific codegen, so a
lowering for an arch you do not have still fails or passes its trace here:
_load_sm120_kernel_module(None, TemplateParams(dtype_qkv=0, dtype_o=5), fp8=True).compile(compute_capability=(12, 0), b=1, qh=2, kh=2, sq=256, skv=256, d_qk=128, d_v=128)on an SM100 box reproduced the SM120 lane's'NoneType' object has no attribute 'iterator'exactly. Run it for every template variant you touched before pushing. - Inside a
defnested in a kernel body, do not touch a free variable (attribute access, store through it) inside a dynamicif: the DSL's region rewrite yields and rebinds the names it sees written there, which makes the free variable an unbound closure-local — it reads asNoneat trace time, orUnboundLocalErrorif you print it at the closure's top. Hoist what the closure needs into a local before thedef(o_ptr = o.iterator.raw_ptr(),sfo_base_ptr) and let the closure add offsets only.
Rule S7 — Promote indices before stride multiplication when the addressed span needs Int64.
A stride can fit in Int32 while index * stride does not. Promote the index before multiplication when the addressed span
requires Int64; casting the completed product preserves an overflow. Exercise
both a physical stride above 2**32 and a smaller stride whose last batch
offset exceeds 2**32, including input, Stats and gradient ports. Seed the
wrapped addresses inside allocated guard storage, so a deliberately narrowed
control fails numerically without an out-of-bounds access. See
TestPreparedSm120Bwd.test_physical_batch_stride_above_int32.
When removing wrapper-side output clears, verify that the prepared chain
overwrites every element, including masked rows and partial tiles. Poison
fresh auxiliary outputs with NaNs, forbid the removed Torch clear calls, and
replay after previously active rows become fully masked. The detector is
test_wrapper_aux_outputs_need_no_torch_clear for SM80 backward dBias/dSink.
When changing tile, packing, CGA or split candidates, spy on the chooser's inputs for both split and unsplit legs: physical CTA count can differ from public MMA width, and masked KV work depends on the candidate Q span and tile alignment. Compare masked bounds with an independent visible-key oracle and verify every alternative is rescored, deduplicated and within the candidate cap. An exact winning-rank golden alone does not detect stale model inputs.