From 31d8b35cf0ddd7ef180538c44861082910ed2d23 Mon Sep 17 00:00:00 2001 From: Emil Gilliam Date: Tue, 29 Sep 2026 12:24:40 -0700 Subject: [PATCH] frost(sdpa_bwd_sm80): serve THD ports at their own head stride The ragged sweeps fuzz head-axis stride gaps (PR #960): each ragged port may be a head-interleaved record, head stride D + k*16 bytes. The SM80 backward's prepared kernel already addresses heads through each operand's stride, but the plan-time rule and the prepared/staged operand geometry still pinned the head stride to D, so every such graph was declined and, on A100, ran nowhere. - Rule (engine mismatch and adapter check_support): element stride 1, head stride >= D and token stride >= H * head_stride, each a multiple of 8 elements so every head base stays 16-byte aligned. Gated on the new Capabilities.thd_head_stride, set on the SM80 row only; the SM100 row keeps head stride == D. - The adapter records each port's head stride next to its token stride; the prepared and staged THD geometry use it instead of D. - The bounded dQ cast and dK/dV fold address the caller's head through the output's head stride instead of the head dim. Tests: head-gapped graph cases (every port, head and token gaps together, GQA, padded head dim, causal + deterministic) with NaN gap cells asserted never read or written; accept, misaligned and overlapping probes; the SM100 row still declines a head gap. Tracker footnote k updated. Co-Authored-By: Claude Opus 5.5 --- python/cudnn/sdpa/bwd/api_dsl.py | 30 +-- python/cudnn/sdpa/bwd/engines.py | 33 ++-- .../cudnn/sdpa/bwd/kernels/sm80/bprop_f16.py | 7 +- python/cudnn/sdpa/bwd/prepared_sm80.py | 2 +- python/cudnn/sdpa/bwd/staged_sm80.py | 4 +- .../sdpa/frost/SUPPORT_MATRIX_TRACKER.md | 16 +- .../sdpa/frost/test_sdpa_bwd_thd_sm80.py | 176 ++++++++++++------ 7 files changed, 176 insertions(+), 92 deletions(-) diff --git a/python/cudnn/sdpa/bwd/api_dsl.py b/python/cudnn/sdpa/bwd/api_dsl.py index d2cb6205ed..0abd482aef 100644 --- a/python/cudnn/sdpa/bwd/api_dsl.py +++ b/python/cudnn/sdpa/bwd/api_dsl.py @@ -1086,6 +1086,7 @@ def _initialize_implementation(self) -> None: self._thd_lse_token_major: bool = False self._thd_lse_head_stride: int = 0 self._thd_token_strides: dict = {} # port role -> the caller's packed token stride (plan-time) + self._thd_head_strides: dict = {} # port role -> the caller's head stride (plan-time; D when compact) @staticmethod def _thd_total(capacity: int, declared: Optional[int]) -> int: @@ -1107,13 +1108,14 @@ def thd_total_kv(self) -> Optional[int]: @staticmethod def _packed_bshd(desc: TensorDesc) -> bool: """True when a logical-BHSD desc sits on PACKED BSHD rows the kernels can - bind directly: head stride D, element stride 1, and a token stride that - covers the row (``>= H*D``) and keeps every row 16-byte aligned (a - multiple of 8 fp16/bf16 elements -- the cp.async loads move 16-byte - chunks). The token stride need not be compact: the kernels address a - row as ``token * token_stride + head * D`` with the port's own plan-time - stride, so a view into a wider per-token record (a K/V slice of an - interleaved ``[T, 2, H, D]`` buffer) is served at that stride. The + bind directly: element stride 1, head stride ``>= D`` and token stride + ``>= H * head_stride``, each a multiple of 8 fp16/bf16 elements so every + head base stays 16-byte aligned (the cp.async loads move 16-byte + chunks). Neither need be compact: the kernels address a row as + ``token * token_stride + head * head_stride`` with the port's own + plan-time strides, so a view into a wider per-token record (a K/V slice + of an interleaved ``[T, 2, H, D]`` buffer) or a head-interleaved record + is served at those strides. The batch stride is not consulted -- a ragged port's sequences start at the ragged offsets, and the packed view rebuilds that axis from the token extent. The token stride is always checked (the packed view walks every @@ -1121,7 +1123,9 @@ def _packed_bshd(desc: TensorDesc) -> bool: strides wildcard on a size-1 extent (the analyzer's convention).""" _, h, _, d = (int(x) for x in desc.shape) st = tuple(int(x) for x in desc.stride) - return (int(desc.shape[1]) == 1 or st[1] == d) and st[2] >= h * d and st[2] % 8 == 0 and (d == 1 or st[3] == 1) + hs = st[1] if h > 1 else d + head_ok = h == 1 or (hs >= d and hs % 8 == 0) + return head_ok and st[2] >= h * hs and st[2] % 8 == 0 and (d == 1 or st[3] == 1) def _checked_lse_view(self, lse_tensor: torch.Tensor) -> torch.Tensor: """Validate a caller-provided Stats/LSE buffer and return the @@ -1233,9 +1237,10 @@ def check_support(self) -> bool: # No staging leg: the kernels bind the caller's packed rows directly, # each port at its own plan-time token stride (the compiled fake # carries it). Anything the row arithmetic cannot express -- a - # head stride other than D, a token stride below the row or off - # 16-byte alignment -- is a decline here, not a silent mis-bind. + # head or token stride below its extent or off 16-byte alignment -- + # is a decline here, not a silent mis-bind. self._thd_token_strides = {} + self._thd_head_strides = {} for role, desc in ( ("q", self.q_desc), ("k", self.k_desc), @@ -1248,10 +1253,11 @@ def check_support(self) -> bool: ): self._value_error_if( not self._packed_bshd(desc), - f"SM80 bwd THD: {desc.name} must be packed BSHD rows (head stride D, element stride 1, " - f"token stride >= H*D and a multiple of 8 elements); got {tuple(desc.stride)} (the packed path has no staging copy)", + f"SM80 bwd THD: {desc.name} must be packed BSHD rows (element stride 1, head stride >= D and " + f"token stride >= H * head stride, each a multiple of 8 elements); got {tuple(desc.stride)} (the packed path has no staging copy)", ) self._thd_token_strides[role] = int(desc.stride[2]) + self._thd_head_strides[role] = int(desc.stride[1]) if int(desc.shape[1]) > 1 else int(desc.shape[3]) self._t_q_cap = self._thd_total(int(b) * int(s_qo), self.max_total_seq_len_q) self._t_kv_cap = self._thd_total(int(b) * int(s_kv), self.max_total_seq_len_kv) self._value_error_if(self._t_q_cap <= 0 or self._t_kv_cap <= 0, "SM80 bwd THD: the packed token capacities must be > 0") diff --git a/python/cudnn/sdpa/bwd/engines.py b/python/cudnn/sdpa/bwd/engines.py index 5fbc7c26b6..97cbd0be63 100644 --- a/python/cudnn/sdpa/bwd/engines.py +++ b/python/cudnn/sdpa/bwd/engines.py @@ -241,6 +241,9 @@ class Capabilities: # threads the real length keeps 1. Appended last: Capabilities is # positional-append-only. bottom_right_s_q_multiple: int = 1 + # THD ports may be head-interleaved (head stride >= D, a multiple of 8 + # elements); rows without it require head stride == D. + thd_head_stride: bool = False def mismatch(capabilities: Capabilities, facts: "ga.SdpaGraphFacts", requested: Any = None) -> Optional[str]: @@ -364,14 +367,14 @@ def mismatch(capabilities: Capabilities, facts: "ga.SdpaGraphFacts", requested: if not facts.bshd_layout: return "Q/K/V/O/dO/dQ/dK/dV must be BSHD-physical under THD (the packed path has no staging copy)" # PACKED rows, not necessarily compact: the kernels address a row as - # ``token * token_stride + head * D + col`` with each port's own - # plan-time token stride, so a port declared as a view into a wider - # per-token record (a K or V slice of an interleaved [T, 2, H, D] - # buffer, token stride 2*H*D) is served at that stride. Required: head - # stride D (no head stride is read), element stride 1, and a token - # stride that covers the row (>= H*D) and keeps every row 16-byte - # aligned (a multiple of 8 fp16/bf16 elements -- the cp.async loads - # move 16-byte chunks). The batch stride is not consulted (a ragged + # ``token * token_stride + head * head_stride + col`` with each port's + # own plan-time strides, so a view into a wider per-token record (a K + # or V slice of an interleaved [T, 2, H, D] buffer, token stride + # 2*H*D) or a head-interleaved record (head stride D + gap) is served + # at those strides. Required: element stride 1, head stride >= D and + # token stride >= H * head_stride, each a multiple of 8 fp16/bf16 + # elements so every head base stays 16-byte aligned (the cp.async + # loads move 16-byte chunks). The batch stride is not consulted (a ragged # port's sequences start at its ragged offsets). The TOKEN stride is # always checked: the packed view walks every token of every sequence # with it, so it is load-bearing even when the envelope S_max is 1. @@ -381,13 +384,15 @@ def mismatch(capabilities: Capabilities, facts: "ga.SdpaGraphFacts", requested: for name, dim, stride in facts.port_layouts: _, h, _, d = (int(x) for x in dim) ts = int(stride[2]) - head_ok = int(dim[1]) == 1 or int(stride[1]) == d - token_ok = ts >= h * d and ts % 8 == 0 + hs = int(stride[1]) if h > 1 else d + head_ok = h == 1 or (hs == d if not capabilities.thd_head_stride else (hs >= d and hs % 8 == 0)) + token_ok = ts >= h * hs and ts % 8 == 0 elem_ok = d == 1 or int(stride[3]) == 1 if not (head_ok and token_ok and elem_ok): return ( - f"{name} must be packed BSHD rows under THD (head stride D, element stride 1, " - f"token stride >= H*D and a multiple of 8 elements); got stride {tuple(stride)}" + f"{name} must be packed BSHD rows under THD (element stride 1, " + f"{'head stride >= D and ' if capabilities.thd_head_stride else 'head stride D, '}" + f"token stride >= H * head stride, each a multiple of 8 elements); got stride {tuple(stride)}" ) # A packed graph has no [B, H, S_q, S_kv] bias: the per-sequence score # rectangles are different sizes, and the kernels' bias read is the @@ -818,7 +823,8 @@ def _sm80_spec() -> EngineSpec: opt-in SMEM, which the sm86/sm89 parts do not have. THD / ragged: served on the packed ``[1, T, H, D]`` path (BSHD ports each - at its own token stride -- a K/V slice of an interleaved record included -- + at its own token and head stride -- a K/V slice of an interleaved record + and a head-interleaved record included -- per-batch ``seq_len_q/kv`` turned into device ``cu_seqlens`` by a setup launch, Stats read in either packed packing, declared totals sizing the carved scratch -- hence ``thd_declared_totals``). Deterministic dQ, @@ -855,6 +861,7 @@ def _sm80_spec() -> EngineSpec: # build time; B * S_max would be the fallback and is far larger # than any packed buffer. thd_declared_totals=True, + thd_head_stride=True, decode=False, # prefill kernels only layouts=frozenset({"bshd", "dense_flex"}), # The kernels READ the declared stats strides natively (the diff --git a/python/cudnn/sdpa/bwd/kernels/sm80/bprop_f16.py b/python/cudnn/sdpa/bwd/kernels/sm80/bprop_f16.py index e4306ec7db..9257aaf952 100644 --- a/python/cudnn/sdpa/bwd/kernels/sm80/bprop_f16.py +++ b/python/cudnn/sdpa/bwd/kernels/sm80/bprop_f16.py @@ -1473,7 +1473,10 @@ def _cast_thd_kernel( src = cutlass.make_array_view(dQ_acc).data_ptr() + (cutlass.Int64(row) * cutlass.Int64(H) + cutlass.Int64(h)) * cutlass.Int64(fd) + c # The caller's dQ at ITS token stride (compact folds to H*d_out). dst = ( - cutlass.make_array_view(dQ_out).data_ptr() + cutlass.Int64(row) * cutlass.Int64(dQ_out.stride[1]) + cutlass.Int64(h) * cutlass.Int64(d_out) + c + cutlass.make_array_view(dQ_out).data_ptr() + + cutlass.Int64(row) * cutlass.Int64(dQ_out.stride[1]) + + cutlass.Int64(h) * cutlass.Int64(dQ_out.stride[2]) + + c ) v = Pointer(src, dtype=cutlass.Float32).load(count=2) Pointer(dst, dtype=cutlass.Int32).store(fp32_to_fp16(v[0], v[1], dtype=io_dtype), alignment=4) @@ -1565,7 +1568,7 @@ def _dkv_reduce_thd_kernel( for g in cutlass.range_constexpr(ratio): acc = acc + Pointer(src + in_base + cutlass.Int32(g * d_in), dtype=io_dtype).load().to(cutlass.Float32) # The caller's dK/dV at ITS token stride (compact folds to Hk*d_out). - out_off = cutlass.Int64(row) * cutlass.Int64(OUT.stride[1]) + cutlass.Int64(hk) * cutlass.Int64(d_out) + di + out_off = cutlass.Int64(row) * cutlass.Int64(OUT.stride[1]) + cutlass.Int64(hk) * cutlass.Int64(OUT.stride[2]) + di Pointer(cutlass.make_array_view(OUT).data_ptr() + out_off, dtype=io_dtype).store(acc.to(io_dtype)) diff --git a/python/cudnn/sdpa/bwd/prepared_sm80.py b/python/cudnn/sdpa/bwd/prepared_sm80.py index 88efcb3a43..c0721eb1de 100644 --- a/python/cudnn/sdpa/bwd/prepared_sm80.py +++ b/python/cudnn/sdpa/bwd/prepared_sm80.py @@ -75,7 +75,7 @@ def build_spec(api, d64_module, *, staged=False): tokens = api._t_kv_cap if role in ("k", "v", "dk", "dv") else api._t_q_cap token_stride = api._thd_token_strides[role] shape = (1, shape[1], tokens, shape[3]) - strides = (tokens * token_stride, shape[3], token_stride, 1) + strides = (tokens * token_stride, api._thd_head_strides[role], token_stride, 1) elif role == "stats": if api._thd_lse_token_major: shape, strides = (api._t_q_cap, api.h_q), (api.h_q, 1) diff --git a/python/cudnn/sdpa/bwd/staged_sm80.py b/python/cudnn/sdpa/bwd/staged_sm80.py index d34cd9bc28..a13bea1209 100644 --- a/python/cudnn/sdpa/bwd/staged_sm80.py +++ b/python/cudnn/sdpa/bwd/staged_sm80.py @@ -30,6 +30,7 @@ def layout_for(api): plan = SimpleNamespace(**vars(api)) plan.head_dim_qk, plan.head_dim_v = api.flavor_d_qk, api.flavor_d_v plan._thd_token_strides = dict(api._thd_token_strides) + plan._thd_head_strides = dict(api._thd_head_strides) copies, offset = [], 0 for role in ROLES[:5] + ROLES[6:9]: desc = getattr(api, role + "_desc") @@ -55,6 +56,7 @@ def layout_for(api): setattr(plan, role + "_desc", cooked) if api.thd: plan._thd_token_strides[role] = heads * dim + plan._thd_head_strides[role] = dim copies.append((role, offset, shape)) offset += ws_align(math.prod(shape) * 2) # Auxiliary accumulators stay in the chain workspace. Prepared casts @@ -110,7 +112,7 @@ def run_staged(api, tensors, workspace, stream, scale, rope_freqs): if api.thd: tokens = api._t_kv_cap if role in ("k", "v", "dk", "dv") else api._t_q_cap shape = (1, shape[1], tokens, shape[3]) - strides = (tokens * api._thd_token_strides[role], shape[3], api._thd_token_strides[role], 1) + strides = (tokens * api._thd_token_strides[role], api._thd_head_strides[role], api._thd_token_strides[role], 1) if tensor.dtype != api.dtype or f.shape != shape or any(n > 1 and actual != expected for n, actual, expected in zip(shape, f.strides, strides)): raise ValueError(f"sdpa_bwd_sm80: {role} must match the declared shape, dtype and strides") stats = original["stats"] diff --git a/python/cudnn/sdpa/frost/SUPPORT_MATRIX_TRACKER.md b/python/cudnn/sdpa/frost/SUPPORT_MATRIX_TRACKER.md index f18a149f74..b8f60aa803 100644 --- a/python/cudnn/sdpa/frost/SUPPORT_MATRIX_TRACKER.md +++ b/python/cudnn/sdpa/frost/SUPPORT_MATRIX_TRACKER.md @@ -1174,13 +1174,15 @@ The SM80 backward additionally has a dedicated plain-dense **d=64 fast path** feature-free d=64 graph. ᵏ **THD / ragged backward** (`sdpa_bwd_sm80`). Q/K/V/O/dO and the gradients are -PACKED `[1, T, H, D]` **BSHD rows, each at its own token stride**: head stride -`D`, element stride 1, token stride `>= H*D` and a multiple of 8 elements -(16-byte rows for the `cp.async` loads). A compact port is the common case; a -K/V view into an interleaved `[T, 2, H, D]` record (token stride `2*H*D`, the -fused-KV slicing layout) is served at that stride, the gap columns never read or -written. Stepped strides are plan-time specialization; packed token capacities -remain dynamic in the compiled host. +PACKED `[1, T, H, D]` **BSHD rows, each at its own token and head stride**: +element stride 1, head stride `>= D`, token stride `>= H * head_stride`, each a +multiple of 8 elements (every head base 16-byte aligned for the `cp.async` +loads). A compact port is the common case; a K/V view into an interleaved +`[T, 2, H, D]` record (token stride `2*H*D`, the fused-KV slicing layout) or a +head-interleaved record (head stride `D + gap`, what the ragged sweeps fuzz) is +served at those strides, the gap cells never read or written. Stepped strides +are plan-time specialization; packed token capacities remain dynamic in the +compiled host. Lengths arrive as the graph's per-batch `seq_len_q/kv` (`use_padding_mask=True`) and become `cu_seqlens` on device in a one-warp setup launch. Like every FROST THD row, the packed addressing is `prefix(lengths) × token stride`: the bound diff --git a/test/python/sdpa/frost/test_sdpa_bwd_thd_sm80.py b/test/python/sdpa/frost/test_sdpa_bwd_thd_sm80.py index f4065984ba..9340860648 100644 --- a/test/python/sdpa/frost/test_sdpa_bwd_thd_sm80.py +++ b/test/python/sdpa/frost/test_sdpa_bwd_thd_sm80.py @@ -326,34 +326,43 @@ def _plan_index(g, name=_ENGINE): _GAP_ROLES = ("q", "k", "v", "o", "do", "dq", "dk", "dv") -def _gapped(x, gap, fill=float("nan")): +def _gapped(x, gap, fill=float("nan"), head_gap=0): """``x`` ([1, cap, nh, dd] compact) re-homed in a per-token record ``gap`` - elements wider: the returned view has token stride ``nh*dd + gap`` and the - gap columns hold ``fill`` (NaN by default -- a read of a gap column would - poison the result, a write would show up in the storage). ``gap <= 0`` - returns ``x`` itself (a negative gap declares OVERLAPPING rows, a + elements wider, each head ``head_gap`` elements wider: the returned view has + head stride ``dd + head_gap`` and token stride ``nh * (dd + head_gap) + gap``, + and the gap cells hold ``fill`` (NaN by default -- a read of a gap cell + would poison the result, a write would show up in the storage). No gaps + returns ``x`` itself (a negative gap declares OVERLAPPING rows / heads, a probe-only decline that never binds data).""" - if gap <= 0: + if gap <= 0 and head_gap <= 0: return x + gap, head_gap = max(gap, 0), max(head_gap, 0) _, cap, nh, dd = x.shape - ts = nh * dd + gap + hs = dd + head_gap + ts = nh * hs + gap stor = torch.full((cap, ts), fill, device=x.device, dtype=x.dtype) - view = stor.as_strided((1, cap, nh, dd), (cap * ts, ts, dd, 1)) + view = stor.as_strided((1, cap, nh, dd), (cap * ts, ts, hs, 1)) view.copy_(x) return view def _gap_columns(view): - """The gap columns of a ``_gapped`` view's per-token records (empty for a - compact tensor).""" + """The gap cells of a ``_gapped`` view's per-token records -- every storage + element of the record span the view does not cover (token gap columns and + head gap columns alike; empty for a compact tensor).""" _, cap, nh, dd = view.shape ts = view.stride(1) if ts == nh * dd: return view.new_empty(0) - return view.as_strided((cap, ts - nh * dd), (ts, 1), view.storage_offset() + nh * dd) + n = cap * ts + covered = torch.arange(n, device=view.device).as_strided((1, cap, nh, dd), view.stride()).reshape(-1) + mask = torch.ones(n, dtype=torch.bool, device=view.device) + mask[covered] = False + stor = view.as_strided((n,), (1,), view.storage_offset()) + return stor[mask] -def _build_thd_bwd_graph(case, *, stats_layout="head_major", declare_totals=True, gaps=None, sink=None, **sdpa_kwargs): +def _build_thd_bwd_graph(case, *, stats_layout="head_major", declare_totals=True, gaps=None, head_gaps=None, sink=None, **sdpa_kwargs): """A ragged backward graph over ``case``'s packed buffers. Everything is declared as the ENVELOPE (B, H, S_max, D) plus a per-tensor @@ -361,25 +370,36 @@ def _build_thd_bwd_graph(case, *, stats_layout="head_major", declare_totals=True role (``q k v o do dq dk dv``) to extra elements on its token stride: the port is then a view into a wider per-token record (``k``/``v`` at ``H_kv * D`` is the fused-KV interleaved layout), which the SM80 packed - path serves at that stride. The bound input buffers are re-homed in such - records with NaN in the gap columns. ``sink`` adds the ``sink_token`` / + path serves at that stride. ``head_gaps`` likewise widens a port's HEAD + stride to ``D + gap`` (a head-interleaved record); the token stride grows + to ``H * (D + gap)`` plus the token gap. The bound input buffers are + re-homed in such records with NaN in every gap cell. ``sink`` adds the ``sink_token`` / ``dSink_token`` ports ((1, H, 1, 1) fp32); the graph's ``_thd_test_ports["dsink"]`` handle finds the dSink buffer in the pack. """ b, h, hkv, d, d_v, dev = case.b, case.h, case.hkv, case.d, case.d_v, "cuda" gaps = dict(gaps or {}) + head_gaps = dict(head_gaps or {}) assert set(gaps) <= set(_GAP_ROLES), gaps + assert set(head_gaps) <= set(_GAP_ROLES), head_gaps + + def _strides(s_max, nh, dd, role): + """Envelope strides of ``role``'s record: head stride ``dd + head gap``, + token stride ``nh * head_stride + token gap``.""" + hs = dd + head_gaps.get(role, 0) + ts = nh * hs + gaps.get(role, 0) + return [s_max * ts, hs, ts, 1] + io = cudnn.data_type.HALF if case.dtype == torch.float16 else cudnn.data_type.BFLOAT16 s_max_q, s_max_kv = max(max(case.lens_q), 1), max(max(case.lens_kv), 1) g = cudnn.pygraph(io_data_type=io, intermediate_data_type=cudnn.data_type.FLOAT, compute_data_type=cudnn.data_type.FLOAT) vp, t, geom = {}, {}, {} - def _port(name, s_max, nh, dd, cu, gap=0): - """Envelope (B, nh, s_max, dd) with packed BSHD strides (token stride - ``nh*dd`` widened by ``gap``) and an element ragged offset of cu * token stride.""" - ts = nh * dd + gap - stride = [s_max * ts, dd, ts, 1] - ro_t = (torch.tensor(cu, dtype=torch.int64, device=dev) * ts).view(b + 1, 1, 1, 1) + def _port(name, s_max, nh, dd, cu): + """Envelope (B, nh, s_max, dd) with packed BSHD strides (token and head + stride widened by the role's gaps) and an element ragged offset of cu * token stride.""" + stride = _strides(s_max, nh, dd, name) + ro_t = (torch.tensor(cu, dtype=torch.int64, device=dev) * stride[2]).view(b + 1, 1, 1, 1) geom[name] = (s_max, stride, nh, dd, ro_t) x = g.tensor(name=name, dim=[b, nh, s_max, dd], stride=stride, data_type=io) ro = g.tensor(name=f"{name}_ro", dim=[b + 1, 1, 1, 1], stride=[1, 1, 1, 1], data_type=cudnn.data_type.INT64) @@ -387,33 +407,15 @@ def _port(name, s_max, nh, dd, cu, gap=0): vp[ro] = ro_t return x - t["q"] = _port("q", s_max_q, h, d, case.cu_q, gaps.get("q", 0)) - t["o"] = _port("o", s_max_q, h, d_v, case.cu_q, gaps.get("o", 0)) - t["do"] = _port("do", s_max_q, h, d_v, case.cu_q, gaps.get("do", 0)) - t["k"] = _port("k", s_max_kv, hkv, d, case.cu_k, gaps.get("k", 0)) - t["v"] = _port("v", s_max_kv, hkv, d_v, case.cu_k, gaps.get("v", 0)) + t["q"] = _port("q", s_max_q, h, d, case.cu_q) + t["o"] = _port("o", s_max_q, h, d_v, case.cu_q) + t["do"] = _port("do", s_max_q, h, d_v, case.cu_q) + t["k"] = _port("k", s_max_kv, hkv, d, case.cu_k) + t["v"] = _port("v", s_max_kv, hkv, d_v, case.cu_k) # The gradients' own records (declared below, once the node exists). - geom["dq"] = ( - s_max_q, - [s_max_q * (h * d + gaps.get("dq", 0)), d, h * d + gaps.get("dq", 0), 1], - h, - d, - torch.tensor(case.cu_q, dtype=torch.int64, device=dev).view(b + 1, 1, 1, 1) * (h * d + gaps.get("dq", 0)), - ) - geom["dk"] = ( - s_max_kv, - [s_max_kv * (hkv * d + gaps.get("dk", 0)), d, hkv * d + gaps.get("dk", 0), 1], - hkv, - d, - torch.tensor(case.cu_k, dtype=torch.int64, device=dev).view(b + 1, 1, 1, 1) * (hkv * d + gaps.get("dk", 0)), - ) - geom["dv"] = ( - s_max_kv, - [s_max_kv * (hkv * d_v + gaps.get("dv", 0)), d_v, hkv * d_v + gaps.get("dv", 0), 1], - hkv, - d_v, - torch.tensor(case.cu_k, dtype=torch.int64, device=dev).view(b + 1, 1, 1, 1) * (hkv * d_v + gaps.get("dv", 0)), - ) + for role, s_max, nh, dd, cu in (("dq", s_max_q, h, d, case.cu_q), ("dk", s_max_kv, hkv, d, case.cu_k), ("dv", s_max_kv, hkv, d_v, case.cu_k)): + stride = _strides(s_max, nh, dd, role) + geom[role] = (s_max, stride, nh, dd, torch.tensor(cu, dtype=torch.int64, device=dev).view(b + 1, 1, 1, 1) * stride[2]) # Packed Stats in one of the layouts a forward emits. head_major is # (1, H, head_stride) with a 64-rounded token capacity (WIDER than the @@ -474,9 +476,10 @@ def _port(name, s_max, nh, dd, cu, gap=0): out.set_ragged_offset(ro) vp[ro] = ro_t # Inputs re-homed in their declared records (NaN gap columns: never read). - vp.update({t[r]: _gapped(getattr(case, r), gaps.get(r, 0)) for r in ("q", "k", "v", "o", "do")}) + vp.update({t[r]: _gapped(getattr(case, r), gaps.get(r, 0), head_gap=head_gaps.get(r, 0)) for r in ("q", "k", "v", "o", "do")}) g._thd_test_ports = t # the sink/dSink handles, for the sinks test g._thd_test_gaps = gaps # the gradients' records, for _run_graph + g._thd_test_head_gaps = head_gaps return g, vp, (dq_t, dk_t, dv_t) @@ -517,12 +520,12 @@ def _run_graph(lens_q, lens_kv, *, h=2, hkv=None, d=_D, d_v=None, dtype=torch.bf ) g, vp, (dq_t, dk_t, dv_t) = _build_thd_bwd_graph(case, stats_layout=stats_layout, **kw) _plan_graph(g) - gaps = g._thd_test_gaps - # NaN-filled gradient records at the declared token strides: a gap column - # that stops being NaN was written, a live row that stays NaN was skipped. - dq = _gapped(torch.full_like(case.q, float("nan")), gaps.get("dq", 0)) - dk = _gapped(torch.full((1, case.cap_kv, case.hkv, case.d), float("nan"), device="cuda", dtype=dtype), gaps.get("dk", 0)) - dv = _gapped(torch.full((1, case.cap_kv, case.hkv, case.d_v), float("nan"), device="cuda", dtype=dtype), gaps.get("dv", 0)) + gaps, hgaps = g._thd_test_gaps, g._thd_test_head_gaps + # NaN-filled gradient records at the declared token and head strides: a gap + # cell that stops being NaN was written, a live row that stays NaN was skipped. + dq = _gapped(torch.full_like(case.q, float("nan")), gaps.get("dq", 0), head_gap=hgaps.get("dq", 0)) + dk = _gapped(torch.full((1, case.cap_kv, case.hkv, case.d), float("nan"), device="cuda", dtype=dtype), gaps.get("dk", 0), head_gap=hgaps.get("dk", 0)) + dv = _gapped(torch.full((1, case.cap_kv, case.hkv, case.d_v), float("nan"), device="cuda", dtype=dtype), gaps.get("dv", 0), head_gap=hgaps.get("dv", 0)) vp.update({dq_t: dq, dk_t: dk, dv_t: dv}) ws = torch.empty(max(g.get_workspace_size(), 1), device="cuda", dtype=torch.uint8) g.execute(vp, ws) @@ -632,6 +635,31 @@ def test_graph_thd_gapped_token_strides(variant): _run_graph((300, 128, 200), (300, 128, 200), **_GAP_CASES[variant]) +_HEAD_GAP_CASES = { + # Every port head-interleaved by one 16-byte unit (what the ragged sweeps + # draw since their head-gap fuzzing): head stride D + 8, compact otherwise. + "all_ports": dict(head_gaps={r: 8 for r in _GAP_ROLES}), + # Head AND token gaps, each port on its own schedule. + "head_and_token": dict(head_gaps={r: 8 * (i % 3 + 1) for i, r in enumerate(_GAP_ROLES)}, gaps={r: 8 * (i + 1) for i, r in enumerate(_GAP_ROLES)}), + # GQA: the K/V loads and the dK/dV fold at the KV side's head stride. + "gqa": dict(h=4, hkv=2, head_gaps={"q": 8, "k": 16, "v": 24, "dk": 8, "dv": 16}), + # A head dim inside the envelope: the staging copy reads the head-gapped + # record, the cast and fold write it. + "padded_d": dict(d=96, head_gaps={r: 8 for r in _GAP_ROLES}), + # Causal + deterministic relay with a head-interleaved Q side. + "causal_det": dict(use_causal_mask=True, use_deterministic_algorithm=True, head_gaps={"q": 8, "o": 8, "do": 8, "dq": 8}), +} + + +@pytest.mark.parametrize("variant", sorted(_HEAD_GAP_CASES), ids=sorted(_HEAD_GAP_CASES)) +def test_graph_thd_gapped_head_strides(variant): + """Ports declared with a head stride wider than D (head-interleaved records) + are read and written at their own head stride: the reference matches per + sequence, the NaN gap cells of the inputs were never read and those of the + gradients never written.""" + _run_graph((300, 128, 200), (300, 128, 200), **_HEAD_GAP_CASES[variant]) + + def test_graph_thd_gapped_capacity_tail_left_untouched(): """The capacity-tail contract on gapped records: rows past the packed total and the gap columns both stay NaN, on the direct-bound MHA epilogue.""" @@ -913,8 +941,9 @@ def forbidden(*args, **kwargs): # --- probe: accepts and rejects ----------------------------------------------- -def _thd_mismatch(lens_q=(256, 128), lens_kv=(256, 128), *, h=2, hkv=None, stats_layout="head_major", d=_D, **kw): - """``mismatch()`` for a ragged backward graph on the SM80 row, or None if served.""" +def _thd_mismatch(lens_q=(256, 128), lens_kv=(256, 128), *, h=2, hkv=None, stats_layout="head_major", d=_D, caps=None, **kw): + """``mismatch()`` for a ragged backward graph on the SM80 row (``caps`` + overrides its capabilities), or None if served.""" from cudnn.sdpa import graph_analyzer as ga from cudnn.sdpa.bwd.engines import ENGINE_SPECS, mismatch @@ -928,7 +957,7 @@ def _thd_mismatch(lens_q=(256, 128), lens_kv=(256, 128), *, h=2, hkv=None, stats facts = ga.analyze(g) spec = next(s for s in ENGINE_SPECS if s.name == _ENGINE) assert facts is not None - return mismatch(spec.capabilities, facts) + return mismatch(caps(spec.capabilities) if caps else spec.capabilities, facts) def test_graph_thd_accepts_the_plain_case(): @@ -963,6 +992,41 @@ def test_accept_thd_gapped_token_strides(): assert _thd_mismatch(h=4, hkv=2, gaps={"k": 2 * _D, "v": 2 * _D, "dk": 8, "dv": 16}) is None +def test_accept_thd_gapped_head_strides(): + """A head stride wider than D (a multiple of 8 elements) is served, alone, + per port, and together with a token gap; the size-1 KV head axis wildcards.""" + assert _thd_mismatch(head_gaps={"q": 8}) is None + assert _thd_mismatch(head_gaps={r: 8 * (i % 3 + 1) for i, r in enumerate(_GAP_ROLES)}, gaps={"k": 16, "dq": 8}) is None + assert _thd_mismatch(h=4, hkv=1, head_gaps={"q": 8, "o": 16, "dq": 8}) is None + + +def test_thd_head_strides_are_an_sm80_capability(): + """The relaxed head-stride rule is gated on ``thd_head_stride``: a row + without it (the SM100 backward binds compact packed rows) still declines a + head gap.""" + import dataclasses + + reason = _thd_mismatch(head_gaps={"q": 8}, caps=lambda c: dataclasses.replace(c, thd_head_stride=False)) + assert reason is not None and "head stride D" in reason, reason + + +def test_reject_thd_misaligned_head_stride(): + """A head stride off 16-byte alignment (4 fp16 elements past the head) is a + typed plan-time decline: every head base must stay 16-byte aligned.""" + reason = _thd_mismatch(head_gaps={"k": 4}) + assert reason is not None and "multiple of 8" in reason, reason + reason = _thd_mismatch(head_gaps={"dv": 12}) + assert reason is not None and "multiple of 8" in reason, reason + + +def test_reject_thd_head_stride_below_the_head(): + """A head stride shorter than D (overlapping heads) never reaches the + packed-row rule: the layout envelope's non-overlap check (or the node) + declines it first.""" + reason = _thd_mismatch(head_gaps={"q": -8}) + assert reason is not None and ("refused by the node" in reason or "overlapping" in reason or ">= D" in reason), reason + + def test_reject_thd_misaligned_token_stride(): """A token stride off 16-byte alignment (4 fp16 elements past the row) is a typed plan-time decline: the cp.async loads move 16-byte chunks."""