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
40 changes: 26 additions & 14 deletions docs/fe-oss-apis/attention/sdpa_bwd_sm120.md
Original file line number Diff line number Diff line change
Expand Up @@ -145,16 +145,25 @@ warp-partition triple:
| 192 | 32 × 64 | double-buffered Q at the default config |
| 256 | 32 × 64 | SMEM-bound: single-buffered Q at the default config |

SMEM per CTA is `(Q_STAGES + 1)·M·D + N·D + max(N·D, 2·M·N)` elements against
the ~99 KB SM120 cap. The constructor tries `Q_STAGES = 2` and falls back to a
**single Q buffer** when it doesn't fit; among the default configs, this occurs
at D=256. In the single-buffer branch the iteration reorders GEMM5 *before*
GEMM4 (GEMM5 is sQ's last reader), so the Q refill for the next tile hides
behind GEMM4 and the dQ scatter instead of stalling the loop. Head dims that
are multiples of 8 but are not native sizes are zero-padded by the adapter
(pad columns contribute nothing anywhere in the chain). Explicit
`tile_m`/`tile_n` knobs override the Q/KV tile defaults; off-table combinations
derive their warp partitions from a largest-valid rule.
SMEM per CTA is `Q_STAGES·tile_q·d_qk + tile_q·d_v + tile_kv·d_qk +
max(tile_kv·d_v, 2·tile_q·tile_kv)` elements against the ~99 KB SM120 cap.
The constructor tries `Q_STAGES = 2` and falls back to a **single Q buffer**
when it does not fit; among the default configs, this occurs at D=256. In
the single-buffer branch the iteration reorders GEMM5 *before* GEMM4 (GEMM5
is sQ's last reader), so the Q refill for the next tile hides behind GEMM4
and the dQ scatter instead of stalling the loop. Head dims that are multiples
of 8 but are not native sizes are zero-padded by the adapter (pad columns
contribute nothing anywhere in the chain). Explicit `tile_m`/`tile_n` knobs
override the Q/KV tile defaults; off-table combinations derive their warp
partitions from a largest-valid rule.

The Q/K head dim may exceed the V head dim (MLA: DeepSeek-V3 and Kimi-K2.6
train at 192/128). `d_qk` sizes Q/K/dQ/dK and `d_v` sizes V/O/dO/dV: GEMM1
contracts over `d_qk`, GEMM2 over `d_v`, and dK/dV share one warp partition
with per-side column slices. Tile defaults come from `d_qk`. Unequal dims
must both be multiples of 64 (one smem page/swizzle); the adapter pads each
side to its own native kernel size and raises the VO side to at least 64
when the sizes differ.

### dQ scatter and the scrambled workspace

Expand Down Expand Up @@ -236,7 +245,9 @@ only the unused relay operand remains in the kernel ABI.
- SM120 and SM121, e.g. RTX 5090, RTX PRO 6000 Blackwell, and DGX Spark
- Dtypes: FP16 / BF16 (LSE fp32)
- Head dims: 32/64/128/192/256 natively; any other multiple of 8 up to 256 is
served by zero-padding D to the next supported size (`d_qk == d_v`)
served by zero-padding D to the next supported size. Rectangular
`d_qk >= d_v` (MLA, e.g. 192/128) is supported: each side pads to its own
native size (the VO side raises to >= 64 when the sizes differ)
- Masks: none, causal (top-left or bottom-right), right-band-widened causal
(`diagonal_band_right_bound` > 0, the causal diagonal shifted right by a
compile-time R), sliding window (left-window offset, with or without
Expand All @@ -251,6 +262,7 @@ only the unused relay operand remains in the kernel ABI.
- No dropout / bias / ALiBi / softcap / THD
- Workspace (carved from the caller's buffer): fp32 `delta` and `dq_accum`
scratch plus int32 relay-counter storage (reserved in both modes); GQA adds
the io-dtype `dk_ws`/`dv_ws` partials buffers (`B·S_kv·H_q·D_padded`
elements each, where `D_padded` is the adapter's zero-padded head
dimension); use `scratch_workspace_bytes()` for the exact total
the io-dtype `dk_ws`/`dv_ws` partials buffers (`B·S_kv·H_q·d_qk_padded` and
`B·S_kv·H_q·d_v_padded` elements, where `d_*_padded` are the adapter's
zero-padded head dimensions); use `scratch_workspace_bytes()` for the exact
total
124 changes: 73 additions & 51 deletions python/cudnn/sdpa/bwd/api_dsl.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@
SUPPORTED_HEAD_DIMS as _SM120_SUPPORTED_HEAD_DIMS,
TemplateParams as Sm120TemplateParams,
padded_head_dim as _sm120_padded_head_dim,
padded_head_dims as _sm120_padded_head_dims,
)
from cudnn.sdpa.fwd.api_dsl import WorkspaceCarver, _torch_stream_context, ws_align

Expand Down Expand Up @@ -111,7 +112,8 @@ def __init__(
self.s_k_max: Optional[int] = None
self.h_q: Optional[int] = None
self.h_kv: Optional[int] = None
self.head_dim: Optional[int] = None
self.head_dim_qk: Optional[int] = None
self.head_dim_v: Optional[int] = None
self.dtype: Optional[torch.dtype] = None
self._initialize_implementation()
self._logger.debug("__init__ completed")
Expand Down Expand Up @@ -157,8 +159,11 @@ def _initialize_implementation(self) -> None:
self.compute_capability: Optional[tuple[int, int]] = None
self._k_mod = None
self._sq_rounded: Optional[int] = None
# Kernel-facing head dim. When it differs from D, operands stage through zero-padded compact copies
self.head_dim_padded: Optional[int] = None
# Kernel-facing head dims per side (QK: Q/K/dQ/dK, V: V/O/dO/dV).
# When one differs from its D, that side's operands stage through
# zero-padded compact copies.
self.head_dim_qk_padded: Optional[int] = None
self.head_dim_v_padded: Optional[int] = None
# name -> staging number-of-elements for each non-BSHD-compact port
self._staging_numels: dict[str, int] = {}

Expand Down Expand Up @@ -189,29 +194,41 @@ def check_support(self) -> bool:

b, h_q, s_q, d_qk = self.q_desc.shape
_, h_kv, s_kv, _ = self.k_desc.shape
d_v = int(self.v_desc.shape[3])
self._check_tensor_shape(self.k_desc, (b, h_kv, s_kv, d_qk), name="K")
self._check_tensor_shape(self.v_desc, (b, h_kv, s_kv, d_qk), name="V")
self._check_tensor_shape(self.o_desc, (b, h_q, s_q, d_qk), name="O")
self._check_tensor_shape(self.do_desc, (b, h_q, s_q, d_qk), name="dO")
self._check_tensor_shape(self.v_desc, (b, h_kv, s_kv, d_v), name="V")
self._check_tensor_shape(self.o_desc, (b, h_q, s_q, d_v), name="O")
self._check_tensor_shape(self.do_desc, (b, h_q, s_q, d_v), name="dO")
self._check_tensor_shape(self.dq_desc, tuple(self.q_desc.shape), name="dQ")
self._check_tensor_shape(self.dk_desc, tuple(self.k_desc.shape), name="dK")
self._check_tensor_shape(self.dv_desc, tuple(self.v_desc.shape), name="dV")

for label, val in (("B", b), ("H_q", h_q), ("H_kv", h_kv), ("S_q", s_q), ("S_kv", s_kv), ("D", d_qk)):
for label, val in (("B", b), ("H_q", h_q), ("H_kv", h_kv), ("S_q", s_q), ("S_kv", s_kv), ("D_QK", d_qk), ("D_V", d_v)):
self._value_error_if(int(val) <= 0, f"{label} must be > 0; got {val}")
self._value_error_if(
h_q % h_kv != 0,
f"SM120 DSL SDPA backward requires H_q to be a multiple of H_kv (GQA / MQA); got H_q={h_q}, H_kv={h_kv}",
)
self.head_dim_padded = _sm120_padded_head_dim(int(d_qk)) if d_qk % 8 == 0 else None
self._value_error_if(
self.head_dim_padded is None,
f"D ({d_qk}) must be a multiple of 8 and <= {max(_SM120_SUPPORTED_HEAD_DIMS)}",
d_v > d_qk,
f"SM120 DSL SDPA backward requires D_QK >= D_V (MLA-style rectangular head dims); got D_QK={d_qk}, D_V={d_v}",
)
# All operands stage when D pads; otherwise only non-BSHD-compact ones.
for desc in (self.q_desc, self.k_desc, self.v_desc, self.o_desc, self.do_desc, self.dq_desc, self.dk_desc, self.dv_desc):
if self.head_dim_padded != d_qk or not self._bshd_physical_ok(desc):
self._staging_numels[desc.name] = math.prod(desc.shape[:-1]) * self.head_dim_padded
self._value_error_if(
d_qk % 8 != 0 or _sm120_padded_head_dim(int(d_qk)) is None,
f"D_QK ({d_qk}) must be a multiple of 8 and <= {max(_SM120_SUPPORTED_HEAD_DIMS)}",
)
self._value_error_if(
d_v % 8 != 0 or _sm120_padded_head_dim(int(d_v)) is None,
f"D_V ({d_v}) must be a multiple of 8 and <= {max(_SM120_SUPPORTED_HEAD_DIMS)}",
)
self.head_dim_qk_padded, self.head_dim_v_padded = _sm120_padded_head_dims(int(d_qk), int(d_v))
# A side's operands all stage when its D pads; otherwise only the non-BSHD-compact ones.
for desc in (self.q_desc, self.k_desc, self.dq_desc, self.dk_desc):
if self.head_dim_qk_padded != d_qk or not self._bshd_physical_ok(desc):
self._staging_numels[desc.name] = math.prod(desc.shape[:-1]) * self.head_dim_qk_padded
for desc in (self.v_desc, self.o_desc, self.do_desc, self.dv_desc):
if self.head_dim_v_padded != d_v or not self._bshd_physical_ok(desc):
self._staging_numels[desc.name] = math.prod(desc.shape[:-1]) * self.head_dim_v_padded

self._value_error_if(
self.stats_desc.ndim != 4 or tuple(self.stats_desc.shape) != (b, h_q, s_q, 1),
Expand Down Expand Up @@ -304,7 +321,8 @@ def check_support(self) -> bool:
self.s_k_max = int(s_kv)
self.h_q = int(h_q)
self.h_kv = int(h_kv)
self.head_dim = int(d_qk)
self.head_dim_qk = int(d_qk)
self.head_dim_v = int(d_v)
self._sq_rounded = _round_up(self.s_q_max, _SM120_ROW_ROUND)
self._is_supported = True

Expand Down Expand Up @@ -340,8 +358,9 @@ def compile(self) -> None:
qh=self.h_q,
sq=self.s_q_max,
skv=self.s_k_max,
d=self.head_dim_padded,
d_qk=self.head_dim_qk_padded,
kvh=self.h_kv,
d_v=self.head_dim_v_padded,
)
self._logger.debug("compile completed")

Expand All @@ -350,13 +369,14 @@ def _dq_sem_len(self) -> int:

return self.batch_size * self.h_q * _round_up(self.s_q_max, _SM120_MIN_Q_TILE) // _SM120_MIN_Q_TILE

def _dkv_ws_elems(self) -> int:
"""io-dtype elements of each GQA partials buffer (dk_ws / dv_ws);
0 for MHA, where they alias the dk/dv outputs."""
def _dkv_ws_elems(self) -> tuple[int, int]:
"""io-dtype elements of the GQA partials buffers (dk_ws, dv_ws);
(0, 0) for MHA, where they alias the dk/dv outputs."""

if self.h_q == self.h_kv:
return 0
return self.batch_size * self.s_k_max * self.h_q * self.head_dim_padded
return (0, 0)
rows = self.batch_size * self.s_k_max * self.h_q
return (rows * self.head_dim_qk_padded, rows * self.head_dim_v_padded)

def _checked_seq_lens(self, seq_lens: torch.Tensor, name: str) -> torch.Tensor:
"""Validate per-batch lengths and return a (B,) int32 view (never a copy/cast)."""
Expand All @@ -379,16 +399,17 @@ def _checked_seq_lens(self, seq_lens: torch.Tensor, name: str) -> torch.Tensor:
return seq_lens.reshape(-1)

def scratch_workspace_bytes(self) -> int:
"""delta (fp32 [B, H, SQ_r128]) + dq_accum (fp32 flat [B*SQ_r128*H*D])
"""delta (fp32 [B, H, SQ_r128]) + dq_accum (fp32 flat [B*SQ_r128*H*D_QK])
+ dq_sem (int32 flat [B*H*ceil(SQ/32)], deterministic relay counters)
+ dk_ws/dv_ws (io [B, SKV, H_q, D] each, per-q-head partials, GQA only)
+ one compact staging copy per non-BSHD-compact operand."""
+ dk_ws/dv_ws (io [B, SKV, H_q, D_QK] / [B, SKV, H_q, D_V], per-q-head
partials, GQA only) + one compact staging copy per non-BSHD-compact
operand."""

self._ensure_support_checked()
delta_bytes = ws_align(self.batch_size * self.h_q * self._sq_rounded * 4)
dq_accum_bytes = ws_align(self.batch_size * self._sq_rounded * self.h_q * self.head_dim_padded * 4)
dq_accum_bytes = ws_align(self.batch_size * self._sq_rounded * self.h_q * self.head_dim_qk_padded * 4)
dq_sem_bytes = ws_align(self._dq_sem_len() * 4)
dkv_ws_bytes = 2 * ws_align(self._dkv_ws_elems() * self.dtype.itemsize)
dkv_ws_bytes = sum(ws_align(elems * self.dtype.itemsize) for elems in self._dkv_ws_elems())
staging_bytes = sum(ws_align(numel * self.dtype.itemsize) for numel in self._staging_numels.values())
return delta_bytes + dq_accum_bytes + dq_sem_bytes + dkv_ws_bytes + staging_bytes

Expand Down Expand Up @@ -446,7 +467,7 @@ def execute(

carver = WorkspaceCarver(workspace, self.scratch_workspace_bytes(), "sdpa_bwd_sm120")
delta = carver.take(self.batch_size * self.h_q * self._sq_rounded, torch.float32).reshape(self.batch_size, self.h_q, self._sq_rounded)
dq_accum = carver.take(self.batch_size * self._sq_rounded * self.h_q * self.head_dim_padded, torch.float32)
dq_accum = carver.take(self.batch_size * self._sq_rounded * self.h_q * self.head_dim_qk_padded, torch.float32)
dq_sem = carver.take(self._dq_sem_len(), torch.int32)

if current_stream is None:
Expand All @@ -457,26 +478,27 @@ def execute(

import cutlass

# Non-compact operands (all operands when D pads) stage through
# workspace-carved compact copies with zero-filled pad columns.
d_pad = self.head_dim_padded
pads = d_pad != self.head_dim
# Non-compact operands (a whole side when its D pads) stage through
# workspace-carved compact copies with zero-filled pad columns. The QK
# side (Q/dQ/K/dK) and the VO side (V/O/dO/dV) pad independently.
d_qk, dqk_pad = self.head_dim_qk, self.head_dim_qk_padded
d_v, dv_pad = self.head_dim_v, self.head_dim_v_padded

def _staged_bshd(tensor: torch.Tensor) -> torch.Tensor:
def _staged_bshd(tensor: torch.Tensor, d_orig: int, d_pad: int) -> torch.Tensor:
view = tensor.transpose(1, 2)
if not pads and view.is_contiguous():
if d_pad == d_orig and view.is_contiguous():
return view
b, s, h, _ = view.shape
staged = carver.take(b * s * h * d_pad, self.dtype).view(b, s, h, d_pad)
if pads:
staged[..., self.head_dim :].zero_()
staged[..., : self.head_dim].copy_(view)
if d_pad != d_orig:
staged[..., d_orig:].zero_()
staged[..., :d_orig].copy_(view)
return staged

def _staged_out_bshd(tensor: torch.Tensor):
def _staged_out_bshd(tensor: torch.Tensor, d_orig: int, d_pad: int):
"""(kernel-facing compact BSHD buffer, user view to scatter back into or None)."""
view = tensor.transpose(1, 2)
if not pads and view.is_contiguous():
if d_pad == d_orig and view.is_contiguous():
return view, None
b, s, h, _ = view.shape
return carver.take(b * s * h * d_pad, self.dtype).view(b, s, h, d_pad), view
Expand All @@ -485,14 +507,14 @@ def _staged_out_bshd(tensor: torch.Tensor):
seq_kv_t = self._checked_seq_lens(seq_kv_lens, "seq_kv_lens") if seq_kv_lens is not None else None

with _torch_stream_context(current_stream, q_tensor.device):
q = _staged_bshd(q_tensor)
k = _staged_bshd(k_tensor)
v = _staged_bshd(v_tensor)
o = _staged_bshd(o_tensor)
do = _staged_bshd(do_tensor)
dq, dq_user = _staged_out_bshd(dq_tensor)
dk, dk_user = _staged_out_bshd(dk_tensor)
dv, dv_user = _staged_out_bshd(dv_tensor)
q = _staged_bshd(q_tensor, d_qk, dqk_pad)
k = _staged_bshd(k_tensor, d_qk, dqk_pad)
v = _staged_bshd(v_tensor, d_v, dv_pad)
o = _staged_bshd(o_tensor, d_v, dv_pad)
do = _staged_bshd(do_tensor, d_v, dv_pad)
dq, dq_user = _staged_out_bshd(dq_tensor, d_qk, dqk_pad)
dk, dk_user = _staged_out_bshd(dk_tensor, d_qk, dqk_pad)
dv, dv_user = _staged_out_bshd(dv_tensor, d_v, dv_pad)
lse = stats_tensor.reshape(self.batch_size, self.h_q, self.s_q_max)

kernels = self._compiled_kernel
Expand All @@ -501,9 +523,9 @@ def _staged_out_bshd(tensor: torch.Tensor):
if self.h_q == self.h_kv:
dk_ws, dv_ws = dk, dv
else:
ws_shape = (self.batch_size, self.s_k_max, self.h_q, self.head_dim_padded)
dk_ws = carver.take(self._dkv_ws_elems(), self.dtype).view(ws_shape)
dv_ws = carver.take(self._dkv_ws_elems(), self.dtype).view(ws_shape)
dkw_elems, dvw_elems = self._dkv_ws_elems()
dk_ws = carver.take(dkw_elems, self.dtype).view(self.batch_size, self.s_k_max, self.h_q, dqk_pad)
dv_ws = carver.take(dvw_elems, self.dtype).view(self.batch_size, self.s_k_max, self.h_q, dv_pad)

# Kernel chain (dot -> main -> [reduce] -> cvt)
kernels.dot(o, do, delta, dq_accum, dq_sem, current_stream)
Expand All @@ -530,9 +552,9 @@ def _staged_out_bshd(tensor: torch.Tensor):
kernels.cvt(dq_accum, dq, cutlass.Float32(scale_val), current_stream)
if kernels.dsink is not None:
kernels.dsink(lse, delta, sink_tensor.reshape(self.h_q), dsink_tensor.reshape(self.h_q), seq_q_t, current_stream)
for user_view, staged in ((dq_user, dq), (dk_user, dk), (dv_user, dv)):
for user_view, staged, d_orig in ((dq_user, dq, d_qk), (dk_user, dk, d_qk), (dv_user, dv, d_v)):
if user_view is not None:
user_view.copy_(staged[..., : self.head_dim])
user_view.copy_(staged[..., :d_orig])


def _tensor_signature(tensor: torch.Tensor) -> tuple:
Expand Down
15 changes: 14 additions & 1 deletion python/cudnn/sdpa/bwd/config_sm120.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,11 +15,24 @@


def padded_head_dim(d: int) -> "int | None":
"""Smallest native bin >= ``d``, or ``None`` when ``d`` exceeds every bin."""
"""Smallest native kernel head-dim size >= ``d``, or ``None`` when ``d`` exceeds them all."""

return min((b for b in SUPPORTED_HEAD_DIMS if b >= d), default=None)


def padded_head_dims(d_qk: int, d_v: int) -> "tuple[int, int] | None":
"""Native kernel head-dim sizes for a head-dim pair."""

d_qk_pad = padded_head_dim(d_qk)
d_v_pad = padded_head_dim(d_v)
if d_qk_pad is None or d_v_pad is None:
return None
# Unequal dims must both be multiples of 64 (one smem swizzle).
if d_v_pad != d_qk_pad:
d_v_pad = max(d_v_pad, 64)
return d_qk_pad, d_v_pad


@dataclass(frozen=True)
class TemplateParams:
"""Per-graph parameters that change the traced SM120 backward kernel.
Expand Down
3 changes: 3 additions & 0 deletions python/cudnn/sdpa/bwd/engines.py
Original file line number Diff line number Diff line change
Expand Up @@ -179,6 +179,8 @@ def mismatch(capabilities: Capabilities, facts: "ga.SdpaGraphFacts", requested:
return f"head dims (D_QK={facts.d_qk}, D_V={facts.d_v}) exceed the {max(capabilities.d)} envelope"
elif facts.d_qk not in capabilities.d:
return f"serves D in {sorted(capabilities.d)}; graph has D={facts.d_qk}"
elif facts.d_v not in capabilities.d:
return f"serves D_V in {sorted(capabilities.d)}; graph has D_V={facts.d_v}"
if facts.dtype not in capabilities.dtypes:
return f"dtype {facts.dtype} not in {sorted(str(d) for d in capabilities.dtypes)}"
if not facts.uniform_dtype:
Expand Down Expand Up @@ -261,6 +263,7 @@ def _sm120_spec() -> EngineSpec:
sm_hi=_BLACKWELL_GEFORCE[1],
# Any head size multipled of 8
d=frozenset(range(8, max(_SM120_HEAD_DIMS) + 1, 8)),
dqk_ge_dv=True,
dtypes=frozenset({cudnn.data_type.HALF, cudnn.data_type.BFLOAT16}),
gqa=True,
causal=True,
Expand Down
Loading