Skip to content

Latest commit

 

History

History
203 lines (180 loc) · 12.2 KB

File metadata and controls

203 lines (180 loc) · 12.2 KB

python/cudnn/sdpa — Agent Guide

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.

Hard rules

Rule S1 — THD/packed Stats (LSE) must stay packed: token-major or head-major, never dense-padded.

  • Consumers read Stats through the same cu_seqlen packing 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/packed flag: token-major is stride_h == 1 and stride_s == H; head-major is stride_s == 1 with stride_h the declared head stride, which must cover the packed token count T — stride_h >= H is not the bound; at T > stride_h the per-head slices alias. T is a runtime total, so plan-time classification can only check stride_s == 1 and stride_h >= 1. In the THD path the packed total is a device value — Rule 3 bans reading it back, so stride_h >= T is caller contract (stated in the prepared THD binding contract), not something the adapter verifies: as_strided bounds-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 with graph_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_metadata and the stats_layout-parametrized THD tests (test_dsl_sm100_thd_stats and siblings) in test/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.py rows. 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 an EngineSpec row. 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 Capabilities row. d_pad_multiple=0 means exact-only, so no envelope at all.
  • Reviewing: if the diff touches fwd/engines.py, bwd/engines.py or cudnn/engines/manifest.py and python/cudnn/sdpa/frost/SUPPORT_MATRIX_TRACKER.md is 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_OFF inside the S_ACC_* column range — e.g. sm100/prefill_d256_fp8.py, P_EVEN_OFF = 96 in 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.ld has landed (tcgen05.wait::ld first — the softmax paths otherwise never wait on loads explicitly), and the storing half waits on it right before the aliased store (mb_softmax_hi_loaded in sm100/prefill_d256_fp8.py, mirroring mb_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.py has no mb_softmax_max: its FUSED_CORR_SPLIT_P schedule (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_empty is 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 the S_ACC_* range (grep -l "P_EVEN_OFF: int = " python/cudnn/sdpa/fwd/kernels/*/*.py, today the six d256 kernels) whose config can select SOFTMAX_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), which RESCALE_THRESHOLD makes rare — force it with K scaled by growth ** (key // 128) (k_tile_growth in test/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 is ref == 0, gpu != 0 elements on masked configs with key mod TILE_N in the aliased range — recover the leaked keys by matching O_gpu - O_ref against V rows.

Rule S4 — Split-K partial Stats keep the combiner's log base.

  • stats_use_log2 changes final Stats only. Partial LSE consumed by a natural-log combine stays natural-log, including direct kernel entry points; apply log2(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_log2 and the split-KV Stats tests. test_sm120_direct_template_stats_base bypasses 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 by d assumes 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 is test_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.py runs under requires_blackwell (SM 100..119, the Rubin lane included — _D128_ARCH picks the kernel), the sm120 line under requires_blackwell_geforce (120..129) in test_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_rubin is 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 def nested in a kernel body, do not touch a free variable (attribute access, store through it) inside a dynamic if: 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 as None at trace time, or UnboundLocalError if you print it at the closure's top. Hoist what the closure needs into a local before the def (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.

Output initialization regressions

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.

Heuristic geometry regressions

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.