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
10 changes: 10 additions & 0 deletions benchmark/linear_attention/results/gdn2/gb200/gdn2_20260814.csv
Original file line number Diff line number Diff line change
Expand Up @@ -19,3 +19,13 @@ gdn2_h64,cudnn,gdn2,8,8192,64,64,128,1.718,16.113,210,67,0.000000,20,5.00,2.47
gdn2_h64,fla,gdn2,8,8192,64,64,128,11.849,45.914,30,24,0.000000,20,0.72,0.87
gdn2_h64,cudnn,gdn2,16,8192,64,64,128,3.001,28.796,240,75,0.000000,20,5.72,2.76
gdn2_h64,fla,gdn2,16,8192,64,64,128,23.265,87.351,31,25,0.000000,20,0.74,0.91
gdn2_hon,cudnn_state_on,gdn2,4,2048,64,64,128,0.384,1.763,117,77,0.000000,20,2.80,2.82
gdn2_hon,cudnn_state_on,gdn2,4,4096,64,64,128,0.744,3.494,121,77,0.000000,20,2.89,2.84
gdn2_hon,cudnn_state_on,gdn2,4,8192,64,64,128,1.468,6.990,123,77,0.000000,20,2.93,2.84
gdn2_hon,cudnn_state_on,gdn2,4,16384,64,64,128,2.925,13.954,123,78,0.000000,20,2.94,2.85
gdn2_hon,cudnn_state_on,gdn2,4,32768,64,64,128,5.832,27.875,124,78,0.000000,20,2.95,2.85
gdn2_hon,cudnn_state_on,gdn2,1,8192,64,64,128,0.691,3.115,65,43,0.000000,20,1.55,1.59
gdn2_hon,cudnn_state_on,gdn2,2,8192,64,64,128,0.718,3.457,126,78,0.000000,20,2.99,2.87
gdn2_hon,cudnn_state_on,gdn2,4,8192,64,64,128,1.467,6.990,123,77,0.000000,20,2.93,2.84
gdn2_hon,cudnn_state_on,gdn2,8,8192,64,64,128,2.919,12.448,124,87,0.000000,20,2.94,3.19
gdn2_hon,cudnn_state_on,gdn2,16,8192,64,64,128,5.677,24.345,127,89,0.000000,20,3.03,3.26
Loading
Sorry, something went wrong. Reload?
Sorry, we cannot display this file.
Sorry, this file is invalid so it cannot be displayed.
Loading
Sorry, something went wrong. Reload?
Sorry, we cannot display this file.
Sorry, this file is invalid so it cannot be displayed.
Loading
Sorry, something went wrong. Reload?
Sorry, we cannot display this file.
Sorry, this file is invalid so it cannot be displayed.
Loading
Sorry, something went wrong. Reload?
Sorry, we cannot display this file.
Sorry, this file is invalid so it cannot be displayed.
10 changes: 10 additions & 0 deletions benchmark/linear_attention/results/gdn2/gb300/gdn2_20260814.csv
Original file line number Diff line number Diff line change
Expand Up @@ -19,3 +19,13 @@ gdn2_h64,cudnn,gdn2,8,8192,64,64,128,1.752,14.667,206,74,0.000000,20,4.90,2.71
gdn2_h64,fla,gdn2,8,8192,64,64,128,11.361,44.740,32,24,0.000000,20,0.76,0.89
gdn2_h64,cudnn,gdn2,16,8192,64,64,128,3.066,28.539,235,76,0.000000,20,5.60,2.78
gdn2_h64,fla,gdn2,16,8192,64,64,128,22.279,84.983,32,25,0.000000,20,0.77,0.93
gdn2_hon,cudnn_state_on,gdn2,4,2048,64,64,128,0.375,1.762,120,77,0.000000,20,2.86,2.82
gdn2_hon,cudnn_state_on,gdn2,4,4096,64,64,128,0.724,3.483,125,78,0.000000,20,2.97,2.85
gdn2_hon,cudnn_state_on,gdn2,4,8192,64,64,128,1.426,6.937,126,78,0.000000,20,3.01,2.86
gdn2_hon,cudnn_state_on,gdn2,4,16384,64,64,128,2.833,13.838,127,78,0.000000,20,3.03,2.87
gdn2_hon,cudnn_state_on,gdn2,4,32768,64,64,128,5.643,27.648,128,78,0.000000,20,3.04,2.87
gdn2_hon,cudnn_state_on,gdn2,1,8192,64,64,128,0.684,3.115,66,43,0.000000,20,1.57,1.59
gdn2_hon,cudnn_state_on,gdn2,2,8192,64,64,128,0.708,3.452,127,78,0.000000,20,3.03,2.88
gdn2_hon,cudnn_state_on,gdn2,4,8192,64,64,128,1.428,6.937,126,78,0.000000,20,3.01,2.86
gdn2_hon,cudnn_state_on,gdn2,8,8192,64,64,128,2.850,12.342,127,88,0.000000,20,3.01,3.22
gdn2_hon,cudnn_state_on,gdn2,16,8192,64,64,128,5.265,24.217,137,89,0.000000,20,3.26,3.28
Loading
Sorry, something went wrong. Reload?
Sorry, we cannot display this file.
Sorry, this file is invalid so it cannot be displayed.
Loading
Sorry, something went wrong. Reload?
Sorry, we cannot display this file.
Sorry, this file is invalid so it cannot be displayed.
Loading
Sorry, something went wrong. Reload?
Sorry, we cannot display this file.
Sorry, this file is invalid so it cannot be displayed.
Loading
Sorry, something went wrong. Reload?
Sorry, we cannot display this file.
Sorry, this file is invalid so it cannot be displayed.
38 changes: 25 additions & 13 deletions docs/python_graph_and_execution_backends.md
Original file line number Diff line number Diff line change
Expand Up @@ -253,18 +253,27 @@ are close.
THD-only: token-packed `[total_T, heads, dim]` tensors with a required
`cu_seqlens` (a dense batch is `[0, T, 2T, ...]`). `gdn_bwd` takes the
forward inputs plus `dO` (and optionally `d_final_state`) and produces
`dQ/dK/dV/dG/dBeta` (+ `d_initial_state` iff `initial_state` is given);
`dQ/dK/dV/dG/dBeta` (+ `d_initial_state` iff `initial_state` is given,
+ `d_a_log`/`d_dt_bias` iff the node carries `safe_gate` — the gate
transform's parameter gradients, with `dG`/`dBeta` then in raw-logit
space under `safe_gate`/`use_beta_sigmoid`);
the cumulative gate and intra-chunk WY matrix are recomputed inside the
engine, so the graph contract carries no forward intermediates. Both are
python-engine-only ops: they have no cuDNN backend lowering, so routing
them to the backend entry raises `cudnnGraphNotSupportedError` at
lowering. The kernels live in `cudnn.linear_attention.cutile.kernels.gdn_chunk_cutile`;
lowering. The kernels live in `cudnn.linear_attention.cutile.kernels.gdn`;
the torch custom op `cudnn.linear_attention.ops.gated_delta_net` is a thin
adapter that builds and executes cached `gdn`/`gdn_bwd` graphs (the SDPA
op pattern), so it inherits whatever engine the planner selects. The
optional `use_qk_l2norm` attribute asks the engine to L2-normalize the q/k
rows in-kernel; `GdnFrostEngine` (the SM100/SM103 forward default) declines
such graphs, the cuTile engine serves them.
rows; `GdnFrostEngine` (the SM100/SM103 default, serving both `gdn` and
`gdn_bwd` on the FROST chunked kernels) serves it through a workspace
helper kernel (normalized q/k copies + saved inverse norms, with the
backward Jacobian projection applied in place after the head-group fold),
and likewise serves `safe_gate` (in-kernel raw-logit gate transform, with
`d_a_log`/`d_dt_bias` produced by a deterministic reduction helper) and
`use_beta_sigmoid`; the cuTile engine remains the fallback for non-128
head dims.
- `KdaFrostEngine` / `KdaCuTileEngine` do the same for the single-node
`kda` / `kda_bwd` ops (Kimi Delta Attention). KDA is GDN with a
per-key-channel decay: its `g` is the log-space vector gate
Expand All @@ -273,22 +282,25 @@ are close.
the forward default on SM100/SM103; the node's `use_qk_l2norm` attribute
(in-kernel L2-normalization of q/k — the KDA model's feature map) passes
through to the kernel (without the in-kernel norm the caller owns the q/k
conditioning). It declines `kda_bwd` (its backward
kernel is a stub), so gradients route to the cuTile engine
(`cudnn.linear_attention.cutile.kernels.kda_chunk_cutile`); the torch op
is `cudnn.linear_attention.ops.kimi_delta_attention`.
conditioning). It serves `kda_bwd` on the FROST backward kernel,
regenerating the per-chunk state checkpoints with a recompute pass when
the graph does not provide them; the cuTile engine
(`cudnn.linear_attention.cutile.kernels.kda`) is the
fallback slot. The torch op is
`cudnn.linear_attention.ops.kimi_delta_attention`.
- Gated DeltaNet v2 (`gdn2` / `gdn2_bwd`) has channel-wise gates — `g`/`beta`
`[total_T, HO, K]` plus a NEW per-value write gate `w` `[total_T, HO, V]`.
GDN-2 has **no** cuTile engine; `Gdn2FrostEngine`
(`cudnn.linear_attention.frost.gdn2_engine`, SM100/SM103) is its only
engine, passes the `use_qk_l2norm` attribute through to the kernel (like
`KdaFrostEngine`), and declines `gdn2_bwd` (stub backward kernel), so the
op (`cudnn.linear_attention.ops.gated_delta_net_v2`) is forward-only for
now.
`KdaFrostEngine`), and serves `gdn2_bwd` the same way (checkpoint
recompute when the series is absent); the op is
`cudnn.linear_attention.ops.gated_delta_net_v2`.
- The FROST engines are pure pass-through: `check_support` requires the
kernel-native dtypes (fp32 gates — io-dtype `beta`/`w` for GDN-2 — int32
`cu_seqlens`, fp32-or-bf16 state ports with matching initial/final dtypes,
fp32 state gradients) and execute hands the caller's buffers straight to
or int64 `cu_seqlens`, fp32-or-bf16 state ports with matching
initial/final dtypes, fp32 state gradients for GDN/KDA and io-dtype
`dBeta`/`dW` for GDN-2) and execute hands the caller's buffers straight to
the kernels, carving any scratch it needs out of the explicit workspace as
DLPack views. The cuTile engines follow the same buffer contract: outputs
are written in place (the caller's output buffers, required in the
Expand Down
73 changes: 57 additions & 16 deletions python/cudnn/_pygraph.py
Original file line number Diff line number Diff line change
Expand Up @@ -2621,7 +2621,7 @@ def _training_phase(node): # norm stats exist only in TRAINING forward phase
"gdn": dict(
node_type=NodeType.GDN,
inputs=("q", "k", "v", "g", "beta", "cu_seqlens", "initial_state", "a_log", "dt_bias"),
attrs=("scale", "output_final_state", "use_qk_l2norm", "checkpoint_every_n_tokens", "safe_gate", "batch_invariant"),
attrs=("scale", "output_final_state", "use_qk_l2norm", "checkpoint_every_n_tokens", "use_beta_sigmoid", "safe_gate", "batch_invariant"),
outputs=("O", "final_state", "state_checkpoints"),
maybe={
"final_state": lambda n: bool(n.params.get("output_final_state", False)),
Expand All @@ -2632,11 +2632,24 @@ def _training_phase(node): # norm stats exist only in TRAINING forward phase
),
"gdn_bwd": dict(
node_type=NodeType.GDN_BWD,
inputs=("q", "k", "v", "g", "beta", "cu_seqlens", "dO", "state_checkpoints", "initial_state", "d_final_state"),
attrs=("scale", "use_qk_l2norm", "batch_invariant"),
outputs=("dQ", "dK", "dV", "dG", "dBeta", "d_initial_state"),
maybe={"d_initial_state": lambda n: "initial_state" in n.inputs},
infer={"dQ": _like("q"), "dK": _like("k"), "dV": _like("v"), "dG": _like("g"), "dBeta": _like("beta"), "d_initial_state": _like("initial_state")},
inputs=("q", "k", "v", "g", "beta", "cu_seqlens", "dO", "state_checkpoints", "initial_state", "d_final_state", "a_log", "dt_bias"),
attrs=("scale", "use_qk_l2norm", "use_beta_sigmoid", "safe_gate", "batch_invariant"),
outputs=("dQ", "dK", "dV", "dG", "dBeta", "d_initial_state", "d_a_log", "d_dt_bias"),
maybe={
"d_initial_state": lambda n: "initial_state" in n.inputs,
"d_a_log": lambda n: bool(n.params.get("safe_gate", False)),
"d_dt_bias": lambda n: bool(n.params.get("safe_gate", False)),
},
infer={
"dQ": _like("q"),
"dK": _like("k"),
"dV": _like("v"),
"dG": _like("g"),
"dBeta": _like("beta"),
"d_initial_state": _like("initial_state"),
"d_a_log": _like("a_log"),
"d_dt_bias": _like("dt_bias"),
},
python_only=True,
),
"kda": dict(
Expand All @@ -2662,17 +2675,39 @@ def _training_phase(node): # norm stats exist only in TRAINING forward phase
),
"kda_bwd": dict(
node_type=NodeType.KDA_BWD,
inputs=("q", "k", "v", "g", "beta", "cu_seqlens", "dO", "state_checkpoints", "initial_state", "d_final_state"),
attrs=("scale", "use_qk_l2norm", "batch_invariant"),
outputs=("dQ", "dK", "dV", "dG", "dBeta", "d_initial_state"),
maybe={"d_initial_state": lambda n: "initial_state" in n.inputs},
infer={"dQ": _like("q"), "dK": _like("k"), "dV": _like("v"), "dG": _like("g"), "dBeta": _like("beta"), "d_initial_state": _like("initial_state")},
inputs=("q", "k", "v", "g", "beta", "cu_seqlens", "dO", "state_checkpoints", "initial_state", "d_final_state", "a_log", "dt_bias"),
attrs=("scale", "use_qk_l2norm", "use_beta_sigmoid", "safe_gate", "gate_lower_bound", "batch_invariant"),
outputs=("dQ", "dK", "dV", "dG", "dBeta", "d_initial_state", "d_a_log", "d_dt_bias"),
maybe={
"d_initial_state": lambda n: "initial_state" in n.inputs,
"d_a_log": lambda n: bool(n.params.get("safe_gate", False)),
"d_dt_bias": lambda n: bool(n.params.get("safe_gate", False)),
},
infer={
"dQ": _like("q"),
"dK": _like("k"),
"dV": _like("v"),
"dG": _like("g"),
"dBeta": _like("beta"),
"d_initial_state": _like("initial_state"),
"d_a_log": _like("a_log"),
"d_dt_bias": _like("dt_bias"),
},
python_only=True,
),
"gdn2": dict(
node_type=NodeType.GDN2,
inputs=("q", "k", "v", "g", "beta", "w", "cu_seqlens", "initial_state", "a_log", "dt_bias"),
attrs=("scale", "output_final_state", "use_qk_l2norm", "checkpoint_every_n_tokens", "safe_gate", "gate_lower_bound", "batch_invariant"),
attrs=(
"scale",
"output_final_state",
"use_qk_l2norm",
"checkpoint_every_n_tokens",
"use_beta_sigmoid",
"safe_gate",
"gate_lower_bound",
"batch_invariant",
),
outputs=("O", "final_state", "state_checkpoints"),
maybe={
"final_state": lambda n: bool(n.params.get("output_final_state", False)),
Expand All @@ -2683,10 +2718,14 @@ def _training_phase(node): # norm stats exist only in TRAINING forward phase
),
"gdn2_bwd": dict(
node_type=NodeType.GDN2_BWD,
inputs=("q", "k", "v", "g", "beta", "w", "cu_seqlens", "dO", "state_checkpoints", "initial_state", "d_final_state"),
attrs=("scale", "use_qk_l2norm", "batch_invariant"),
outputs=("dQ", "dK", "dV", "dG", "dBeta", "dW", "d_initial_state"),
maybe={"d_initial_state": lambda n: "initial_state" in n.inputs},
inputs=("q", "k", "v", "g", "beta", "w", "cu_seqlens", "dO", "state_checkpoints", "initial_state", "d_final_state", "a_log", "dt_bias"),
attrs=("scale", "use_qk_l2norm", "use_beta_sigmoid", "safe_gate", "gate_lower_bound", "batch_invariant"),
outputs=("dQ", "dK", "dV", "dG", "dBeta", "dW", "d_initial_state", "d_a_log", "d_dt_bias"),
maybe={
"d_initial_state": lambda n: "initial_state" in n.inputs,
"d_a_log": lambda n: bool(n.params.get("safe_gate", False)),
"d_dt_bias": lambda n: bool(n.params.get("safe_gate", False)),
},
infer={
"dQ": _like("q"),
"dK": _like("k"),
Expand All @@ -2695,6 +2734,8 @@ def _training_phase(node): # norm stats exist only in TRAINING forward phase
"dBeta": _like("beta"),
"dW": _like("w"),
"d_initial_state": _like("initial_state"),
"d_a_log": _like("a_log"),
"d_dt_bias": _like("dt_bias"),
},
python_only=True,
),
Expand Down
4 changes: 2 additions & 2 deletions python/cudnn/linear_attention/cutile/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,9 +4,9 @@
"""cudnn.linear_attention.cutile: cuTile GDN / KDA implementations.

``GdnCuTileEngine`` (a router ``BaseEngine``) executes single-node GDN /
GDN_BWD graphs on the chunked cuTile kernels in ``kernels/gdn_chunk_cutile``.
GDN_BWD graphs on the chunked cuTile kernels in ``kernels/gdn``.
``KdaCuTileEngine`` does the same for KDA / KDA_BWD graphs on
``kernels/kda_chunk_cutile``.
``kernels/kda``.
"""

from typing import Any
Expand Down
Loading