Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
30 changes: 18 additions & 12 deletions python/cudnn/sdpa/bwd/api_dsl.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -1107,21 +1108,24 @@ 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
token with it, even when the envelope S_max is 1); the head and element
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
Expand Down Expand Up @@ -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),
Expand All @@ -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")
Expand Down
33 changes: 20 additions & 13 deletions python/cudnn/sdpa/bwd/engines.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]:
Expand Down Expand Up @@ -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.
Expand All @@ -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
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand Down
7 changes: 5 additions & 2 deletions python/cudnn/sdpa/bwd/kernels/sm80/bprop_f16.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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))


Expand Down
2 changes: 1 addition & 1 deletion python/cudnn/sdpa/bwd/prepared_sm80.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
4 changes: 3 additions & 1 deletion python/cudnn/sdpa/bwd/staged_sm80.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand All @@ -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
Expand Down Expand Up @@ -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"]
Expand Down
16 changes: 9 additions & 7 deletions python/cudnn/sdpa/frost/SUPPORT_MATRIX_TRACKER.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Loading
Loading