diff --git a/benchmark/linear_attention/results/gdn2/gb200/gdn2_20260814.csv b/benchmark/linear_attention/results/gdn2/gb200/gdn2_20260814.csv index f8bb549fe..e3a323218 100644 --- a/benchmark/linear_attention/results/gdn2/gb200/gdn2_20260814.csv +++ b/benchmark/linear_attention/results/gdn2/gb200/gdn2_20260814.csv @@ -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 diff --git a/benchmark/linear_attention/results/gdn2/gb200/gdn2_fixed_batch_bw.png b/benchmark/linear_attention/results/gdn2/gb200/gdn2_fixed_batch_bw.png index 07d543fd0..4adfd1387 100644 Binary files a/benchmark/linear_attention/results/gdn2/gb200/gdn2_fixed_batch_bw.png and b/benchmark/linear_attention/results/gdn2/gb200/gdn2_fixed_batch_bw.png differ diff --git a/benchmark/linear_attention/results/gdn2/gb200/gdn2_fixed_batch_flops.png b/benchmark/linear_attention/results/gdn2/gb200/gdn2_fixed_batch_flops.png index 4b3fdcfb4..3f13b8393 100644 Binary files a/benchmark/linear_attention/results/gdn2/gb200/gdn2_fixed_batch_flops.png and b/benchmark/linear_attention/results/gdn2/gb200/gdn2_fixed_batch_flops.png differ diff --git a/benchmark/linear_attention/results/gdn2/gb200/gdn2_fixed_seq_bw.png b/benchmark/linear_attention/results/gdn2/gb200/gdn2_fixed_seq_bw.png index bc1065be0..1bec90ae8 100644 Binary files a/benchmark/linear_attention/results/gdn2/gb200/gdn2_fixed_seq_bw.png and b/benchmark/linear_attention/results/gdn2/gb200/gdn2_fixed_seq_bw.png differ diff --git a/benchmark/linear_attention/results/gdn2/gb200/gdn2_fixed_seq_flops.png b/benchmark/linear_attention/results/gdn2/gb200/gdn2_fixed_seq_flops.png index f5e74978f..86d159c44 100644 Binary files a/benchmark/linear_attention/results/gdn2/gb200/gdn2_fixed_seq_flops.png and b/benchmark/linear_attention/results/gdn2/gb200/gdn2_fixed_seq_flops.png differ diff --git a/benchmark/linear_attention/results/gdn2/gb300/gdn2_20260814.csv b/benchmark/linear_attention/results/gdn2/gb300/gdn2_20260814.csv index 83a2af408..11ed31700 100644 --- a/benchmark/linear_attention/results/gdn2/gb300/gdn2_20260814.csv +++ b/benchmark/linear_attention/results/gdn2/gb300/gdn2_20260814.csv @@ -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 diff --git a/benchmark/linear_attention/results/gdn2/gb300/gdn2_fixed_batch_bw.png b/benchmark/linear_attention/results/gdn2/gb300/gdn2_fixed_batch_bw.png index fe5da5fce..952aa3897 100644 Binary files a/benchmark/linear_attention/results/gdn2/gb300/gdn2_fixed_batch_bw.png and b/benchmark/linear_attention/results/gdn2/gb300/gdn2_fixed_batch_bw.png differ diff --git a/benchmark/linear_attention/results/gdn2/gb300/gdn2_fixed_batch_flops.png b/benchmark/linear_attention/results/gdn2/gb300/gdn2_fixed_batch_flops.png index d8ddb0f5c..0d37fe54e 100644 Binary files a/benchmark/linear_attention/results/gdn2/gb300/gdn2_fixed_batch_flops.png and b/benchmark/linear_attention/results/gdn2/gb300/gdn2_fixed_batch_flops.png differ diff --git a/benchmark/linear_attention/results/gdn2/gb300/gdn2_fixed_seq_bw.png b/benchmark/linear_attention/results/gdn2/gb300/gdn2_fixed_seq_bw.png index c89176339..892285ac6 100644 Binary files a/benchmark/linear_attention/results/gdn2/gb300/gdn2_fixed_seq_bw.png and b/benchmark/linear_attention/results/gdn2/gb300/gdn2_fixed_seq_bw.png differ diff --git a/benchmark/linear_attention/results/gdn2/gb300/gdn2_fixed_seq_flops.png b/benchmark/linear_attention/results/gdn2/gb300/gdn2_fixed_seq_flops.png index 73822db2e..054ff4f92 100644 Binary files a/benchmark/linear_attention/results/gdn2/gb300/gdn2_fixed_seq_flops.png and b/benchmark/linear_attention/results/gdn2/gb300/gdn2_fixed_seq_flops.png differ diff --git a/docs/python_graph_and_execution_backends.md b/docs/python_graph_and_execution_backends.md index 9a595cfd1..1b7427b3f 100644 --- a/docs/python_graph_and_execution_backends.md +++ b/docs/python_graph_and_execution_backends.md @@ -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 @@ -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 diff --git a/python/cudnn/_pygraph.py b/python/cudnn/_pygraph.py index e327a475f..a84c3a8b8 100644 --- a/python/cudnn/_pygraph.py +++ b/python/cudnn/_pygraph.py @@ -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)), @@ -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( @@ -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)), @@ -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"), @@ -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, ), diff --git a/python/cudnn/linear_attention/cutile/__init__.py b/python/cudnn/linear_attention/cutile/__init__.py index 1d0d52060..d1ef1252b 100644 --- a/python/cudnn/linear_attention/cutile/__init__.py +++ b/python/cudnn/linear_attention/cutile/__init__.py @@ -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 diff --git a/python/cudnn/linear_attention/cutile/engine.py b/python/cudnn/linear_attention/cutile/engine.py new file mode 100644 index 000000000..324faba24 --- /dev/null +++ b/python/cudnn/linear_attention/cutile/engine.py @@ -0,0 +1,132 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Engine layer shared by the cuTile linear-attention families: the +check_support core both run, and the execute-time buffer gate.""" + +from __future__ import annotations + +import cudnn +from cudnn.frost import buffers + +from ..graph_analyzer import BUFFER_NAME_FROM_CUDNN + +CUTILE_ALIGN = {"cu_seqlens": 4, "a_log": 4, "dt_bias": 4, "beta": 4} + +CUTILE_MIN_CUDART = 13030 + +CUTILE_MAX_D_QK = 256 + + +def cutile_la_gate(engine: str, facts, op: str, dg_want) -> None: + """The cuTile LA engines' shared check_support core: the analyzer record, + the runtime environment, and the gates common to both kernel modules. + + ``dg_want`` is the family's dG output dtype -- fp32 for GDN, the gate's own + dtype for KDA -- and is the only gate that differs between them. Each + engine still probes its own kernel module for the cuda.tile runtime, then + adds its family's feature declines. Mirrors :func:`frost_la_gate`.""" + if facts is None or facts.op != op: + raise NotImplementedError(f"{engine} supports exactly one {op}/{op}_BWD node") + if facts.invalid: + raise NotImplementedError(f"{engine}: {facts.invalid}") + if buffers.current_sm() is None: + raise NotImplementedError(f"{engine} requires a CUDA device") + try: + from cuda.bindings import runtime + except ImportError as exc: + raise NotImplementedError(f"{engine} requires cuda.bindings: {exc}") from exc + err, cudart_version = runtime.cudaRuntimeGetVersion() + if int(err) != 0: + raise NotImplementedError(f"{engine}: cudaRuntimeGetVersion failed ({err})") + if cudart_version < CUTILE_MIN_CUDART: + raise NotImplementedError(f"{engine} requires CUDA 13.3+ (found {cudart_version})") + if facts.checkpoint_every_n_tokens > 0 or facts.wants_state_checkpoints: + raise NotImplementedError(f"{engine}: per-chunk state_checkpoints output is not supported") + + if not facts.uniform_io: + raise NotImplementedError(f"{engine}: q/k/v dtypes must match") + if facts.io_dtype not in (cudnn.data_type.HALF, cudnn.data_type.BFLOAT16, None): + raise NotImplementedError(f"{engine}: q/k/v must be fp16/bf16, got {facts.io_dtype}") + if not facts.thd_layout: + raise NotImplementedError(f"{engine}: q/k/v must be THD [total_T, heads, dim]") + if facts.h_k != facts.h_q: + raise NotImplementedError(f"{engine}: q and k head counts differ ({facts.h_q} vs {facts.h_k})") + if facts.h_q and facts.h_v % facts.h_q != 0: + raise NotImplementedError(f"{engine}: v heads ({facts.h_v}) must be a multiple of q heads ({facts.h_q}; GQA-style v broadcast is FROST-only)") + if facts.d_qk > CUTILE_MAX_D_QK: + raise NotImplementedError(f"{engine}: head dim K must be <= {CUTILE_MAX_D_QK}, got {facts.d_qk}") + if facts.cu_dtype not in (cudnn.data_type.INT32, None): + raise NotImplementedError(f"{engine}: cu_seqlens must be int32 (the device-side table builder reads it directly)") + + fp32 = cudnn.data_type.FLOAT + for port, got in ( + ("initial_state", facts.state_dtype), + ("final_state", facts.final_state_dtype), + ("d_final_state", facts.d_final_state_dtype), + ("d_initial_state", facts.d_initial_state_dtype), + ("a_log", facts.a_log_dtype), + ("dt_bias", facts.dt_bias_dtype), + ): + if got not in (fp32, None): + raise NotImplementedError(f"{engine}: '{port}' must be fp32 (callers convert), got {got}") + + io = facts.io_dtype + if not facts.is_bwd: + outputs = {"O": (facts.o_dtype, io), "final_state": (facts.final_state_dtype, fp32)} + else: + if facts.has_initial_state != facts.wants_d_initial_state: + raise NotImplementedError(f"{engine}: d_initial_state output must be requested iff initial_state is given") + outputs = { + "dQ": (facts.dq_dtype, io), + "dK": (facts.dk_dtype, io), + "dV": (facts.dv_dtype, io), + "dG": (facts.dg_dtype, dg_want), + "dBeta": (facts.dbeta_dtype, facts.beta_dtype), + "d_initial_state": (facts.d_initial_state_dtype, fp32), + } + for port, (got, want) in outputs.items(): + if got is not None and got != want: + raise NotImplementedError(f"{engine}: output {port!r} must be {want} (written in place), got {got}") + + +def expect_table(node, align=CUTILE_ALIGN) -> dict: + """Build-time ``{port: (dims, dtype_name, align_bytes)}`` for + :func:`check_layouts_compact`: bound buffers must match the node's frozen + geometry exactly (one graph per shape), and base pointers must satisfy the + kernel entry's alignment claim.""" + table = {} + for ports in (node.inputs, node.outputs): + for name, t in ports.items(): + if t is None: + continue + dims = tuple(int(d) for d in t.dim) if t.dim else None + table[name] = (dims, BUFFER_NAME_FROM_CUDNN.get(t.get_data_type()), align.get(name, 16)) + return table + + +def check_layouts_compact(plan_name: str, expect, names, views) -> None: + """Execute-time gate: every bound buffer must be CONTIGUOUS (the kernels + stage rank-merged views and whole-buffer zero fills) and must match the + node's build-time dims/dtype and base alignment per ``expect``. ``names`` + and ``views`` are the plan's bound port names and the index-aligned + ``variant_pack.operands`` result.""" + for name, b in zip(names, views): + shape = tuple(b.shape) + stride = getattr(b, "stride", None) + strides = tuple(stride()) if callable(stride) else getattr(b, "strides", None) + exp = expect.get(name) if expect else None + if exp is not None: + dims, dtype_name, align = exp + if dims is not None and shape != dims: + raise ValueError(f"{plan_name}: buffer for {name!r} must match the graph's build-time dims {dims}; got {shape}") + if dtype_name is not None and b.dtype != dtype_name: + raise ValueError(f"{plan_name}: buffer for {name!r} must be {dtype_name} (the node's declared dtype); got {b.dtype}") + if align: + ptr = b.data_ptr() + if ptr % align != 0: + raise ValueError(f"{plan_name}: buffer for {name!r} base pointer must be {align}-byte aligned; got 0x{ptr:x}") + if not buffers.is_contiguous(shape, strides): + raise ValueError( + f"{plan_name}: buffer for {name!r} must be contiguous (the cuTile backend stages rank-merged views); got shape {shape} strides {strides}" + ) diff --git a/python/cudnn/linear_attention/cutile/gdn_engine.py b/python/cudnn/linear_attention/cutile/gdn_engine.py index bf74b6da1..ce3013dbb 100644 --- a/python/cudnn/linear_attention/cutile/gdn_engine.py +++ b/python/cudnn/linear_attention/cutile/gdn_engine.py @@ -2,311 +2,214 @@ # SPDX-License-Identifier: Apache-2.0 """cuTile GDN engine: GDN / GDN_BWD nodes on the chunked cuTile kernels -(``kernels/gdn_chunk_cutile``).""" +(``kernels/gdn``).""" from typing import TYPE_CHECKING from cudnn import behavior_note -from cudnn.engines.base import BaseEngine, CompiledPlan, resolve_node_buffers +from cudnn.engines.base import BaseEngine, CompiledPlan, bind_ports from cudnn.graph_types import NodeType from cudnn.frost import buffers -from cudnn.frost.workspace import Workspace -from ..graph_analyzer import check_layouts_compact, analyze, expect_table, to_buffer_dtype - -# entry base-alignment expectations; ports not listed assume 16 -CUTILE_ALIGN = {"cu_seqlens": 4, "a_log": 4, "beta": 4} +from cudnn.frost.workspace import Workspace, WorkspaceLayout, carve_plan +from ..graph_analyzer import analyze, to_buffer_dtype +from .engine import check_layouts_compact, cutile_la_gate, expect_table if TYPE_CHECKING: from cudnn._pygraph import pygraph -def node_ws_layout(node): - """Static carve plan for one node's pipeline intermediates: name -> - (offset, dtype-name, shape). The chunk count is data-dependent (varlen), - so chunk-indexed entries are sized and SHAPED at the bound - ``cdiv(total, 64) + N`` — the device-built table's sentinel tail keeps - bound-gridded launches inert past the real count. Terminal pipeline - buffers (``o``/``final_state``; the backward's ``dq``/``dk`` finals, - ``wy_dv``, ``dg_cum``, ``db``, ``dstate0``) are NOT carved — execute plants - the caller's output buffers under those names.""" - from .kernels.gdn_chunk_cutile import BT_CHUNK, cdiv, next_power_of_2 - - q, v, cu = (node.inputs[p] for p in ("q", "v", "cu_seqlens")) - total, H, K = q.dim - HV, V = v.dim[1], v.dim[2] - N = cu.dim[0] - 1 - io = to_buffer_dtype(q.get_data_type()) - f32 = "float32" - NT_bound = cdiv(total, BT_CHUNK) + N - l2norm = bool(node.params.get("use_qk_l2norm", False)) - - size = 0 - table = {} - - def carve(name, dtype, shape): - nonlocal size - nbytes = buffers.DTYPE_ITEMSIZE[dtype] - for s in shape: - nbytes *= int(s) - table[name] = (size, dtype, tuple(int(s) for s in shape)) - size += (nbytes + 127) & ~127 # 128B-aligned sequential carve +class GdnCuTilePlan(CompiledPlan): + """Carve plan over the caller's workspace, driven from the normalized + variant pack: the carve layout, geometry, scale and kernel module are fixed + per node at build; between executes only the buffer addresses move.""" - carve("chunk_table", "int32", (NT_bound, 2)) - carve("chunk_count", "int32", (1,)) - carve("chunk_offsets", "int32", (N + 1,)) - carve("dummy", "int32", (4,)) # inert stub backing for absent optional kernel args - carve("g_cum", f32, (total, HV)) - carve("A", io, (total, HV, BT_CHUNK)) - carve("w", io, (total, HV, K)) - carve("u", io, (total, HV, V)) - carve("state_checkpoints", io, (NT_bound, HV, K, V)) - carve("v_new", io, (total, HV, V)) - if l2norm: - carve("q_norm", io, (total * H, K)) - carve("q_rstd", f32, (total * H,)) - carve("k_norm", io, (total * H, K)) - carve("k_rstd", f32, (total * H,)) - if node.node_type == NodeType.GDN_BWD: - carve("dv", io, (total, HV, V)) - carve("dstate", io, (NT_bound, HV, K, V)) - carve("dv2", io, (total, HV, V)) - NK = cdiv(K, min(max(next_power_of_2(K), 16), 64)) - carve("dg_nk", f32, (NK, total, HV)) - carve("dw", io, (total, HV, K)) - if HV != H or l2norm: - # dQ/dK are finals only without l2norm on an MHA config; every - # other combination keeps them (or their head-reduced pair) as - # pipeline intermediates - carve("dq", io, (total, HV, K)) - carve("dk", io, (total, HV, K)) - if HV != H: - carve("wy_dk_hred", io, (total, H, K)) - if l2norm: - carve("dq_hred", io, (total, H, K)) - carve("dk_hred", io, (total, H, K)) - carve("dg", f32, (total, HV)) - carve("wy_dk", io, (total, HV, K)) - carve("wy_dg", f32, (total, HV)) - return size, table + takes_variant_pack = True + plan_name = "GdnCuTileEngine" + def __init__(self, graph): + from .kernels import common + from .kernels import gdn as kernels + + (node,) = graph.nodes + self.kernels = kernels + self.common = common + self.is_bwd = node.node_type == NodeType.GDN_BWD + + q, v, cu = (node.inputs[p] for p in ("q", "v", "cu_seqlens")) + total, H, K = (int(d) for d in q.dim) + HV, V = int(v.dim[1]), int(v.dim[2]) + N = int(cu.dim[0]) - 1 + io = to_buffer_dtype(q.get_data_type()) + f32 = "float32" + isz = buffers.DTYPE_ITEMSIZE[io] + BT = kernels.BT_CHUNK + NT_bound = common.cdiv(total, BT) + N + l2norm = bool(node.params.get("use_qk_l2norm", False)) -class GdnCuTilePlan(CompiledPlan): - """Carve plan over the caller's workspace: the layout is static per node; - the buffer arrives with every execute.""" + layout = WorkspaceLayout() + regions = [ + ("chunk_table", layout.add(NT_bound * 2 * 4), "int32", (NT_bound, 2)), + ("chunk_count", layout.add(4), "int32", (1,)), + ("chunk_offsets", layout.add((N + 1) * 4), "int32", (N + 1,)), + ("dummy", layout.add(16), "int32", (4,)), + ("g_cum", layout.add(total * HV * 4), f32, (total, HV)), + ("A", layout.add(total * HV * BT * isz), io, (total, HV, BT)), + ("w", layout.add(total * HV * K * isz), io, (total, HV, K)), + ("u", layout.add(total * HV * V * isz), io, (total, HV, V)), + ("state_checkpoints", layout.add(NT_bound * HV * K * V * isz), io, (NT_bound, HV, K, V)), + ("v_new", layout.add(total * HV * V * isz), io, (total, HV, V)), + ] + if l2norm: + regions += [ + ("q_norm", layout.add(total * H * K * isz), io, (total * H, K)), + ("q_rstd", layout.add(total * H * 4), f32, (total * H,)), + ("k_norm", layout.add(total * H * K * isz), io, (total * H, K)), + ("k_rstd", layout.add(total * H * 4), f32, (total * H,)), + ] + if self.is_bwd: + NK = common.cdiv(K, min(max(common.next_power_of_2(K), 16), 64)) + regions += [ + ("dv", layout.add(total * HV * V * isz), io, (total, HV, V)), + ("dstate", layout.add(NT_bound * HV * K * V * isz), io, (NT_bound, HV, K, V)), + ("dv2", layout.add(total * HV * V * isz), io, (total, HV, V)), + ("dg_nk", layout.add(NK * total * HV * 4), f32, (NK, total, HV)), + ("dw", layout.add(total * HV * K * isz), io, (total, HV, K)), + ] + if HV != H or l2norm: + regions += [ + ("dq", layout.add(total * HV * K * isz), io, (total, HV, K)), + ("dk", layout.add(total * HV * K * isz), io, (total, HV, K)), + ] + if HV != H: + regions.append(("wy_dk_hred", layout.add(total * H * K * isz), io, (total, H, K))) + if l2norm: + regions += [ + ("dq_hred", layout.add(total * H * K * isz), io, (total, H, K)), + ("dk_hred", layout.add(total * H * K * isz), io, (total, H, K)), + ] + regions += [ + ("dg", layout.add(total * HV * 4), f32, (total, HV)), + ("wy_dk", layout.add(total * HV * K * isz), io, (total, HV, K)), + ("wy_dg", layout.add(total * HV * 4), f32, (total, HV)), + ] + + self.ws_bytes = layout.size + self.carve_names = [name for name, _off, _dtype, _shape in regions] + self.carve = carve_plan(self.plan_name, [(off, dtype, shape) for _name, off, dtype, shape in regions]) + self.expect = expect_table(node) + self.n_seqs = N + self.bound = NT_bound + self.bt_chunk = BT + self.scale = float(node.params.get("scale") or K**-0.5) + self.l2norm = l2norm + self.safe_gate = bool(node.params.get("safe_gate", False)) + self.want_state = "final_state" in node.outputs + + if self.is_bwd: + plant = [("dq_l2", "dQ"), ("dk_l2", "dK")] if l2norm else [("dq" if HV == H else "dq_hred", "dQ"), ("dk" if HV == H else "dk_hred", "dK")] + plant += [("wy_dv", "dV"), ("dg_cum", "dG"), ("db", "dBeta")] + if "initial_state" in node.inputs: + plant.append(("dstate0", "d_initial_state")) + else: + plant = [("o", "O")] + ([("final_state", "final_state")] if self.want_state else []) + self.plant = tuple(plant) - def __init__(self, graph): - self.layouts = [(node, *node_ws_layout(node)) for node in graph.nodes] - self.expects = {node: expect_table(node, CUTILE_ALIGN) for node in graph.nodes} - # nodes execute sequentially, each re-carving the same buffer - self.ws_bytes = max(nbytes for _, nbytes, _ in self.layouts) + self.ports = None + self.names = None + self.indices = None def get_workspace_size(self) -> int: return self.ws_bytes - def execute(self, graph, uid_to_data, ctx) -> None: - from .kernels.common import build_chunk_table, ensure_cuda_context - from .kernels.gdn_chunk_cutile import BT_CHUNK - - node_buffers = resolve_node_buffers(graph, uid_to_data) + def execute(self, graph, variant_pack, ctx) -> None: + if self.ports is None: + self.ports = bind_ports(graph, variant_pack) + (slots,) = self.ports.values() + self.names = list(slots.inputs) + list(slots.outputs) + self.indices = list(slots.inputs.values()) + list(slots.outputs.values()) + views = variant_pack.operands(self.indices) + check_layouts_compact(self.plan_name, self.expect, self.names, views) + nb = dict(zip(self.names, views)) stream = ctx.stream if ctx.stream is not None else 0 - ensure_cuda_context(stream) - ws = Workspace(ctx.workspace, self.ws_bytes, "GdnCuTileEngine") - for node, _nbytes, table in self.layouts: - nb = node_buffers[node] - check_layouts_compact("GdnCuTileEngine", self.expects[node], nb) - cu_seqlens = nb.inputs["cu_seqlens"] - N = node.inputs["cu_seqlens"].dim[0] - 1 - bufs = {name: ws.view(off, dt, shape) for name, (off, dt, shape) in table.items()} - bound = bufs["chunk_table"].shape[0] - build_chunk_table(bufs["chunk_table"], bufs["chunk_count"], bufs["chunk_offsets"], cu_seqlens, N, BT_CHUNK, bound, stream=stream) - if node.node_type == NodeType.GDN: - self.execute_fwd(node, nb, bufs, stream) - else: - self.execute_bwd(node, nb, bufs, stream) - - def execute_fwd(self, node, nb, bufs, stream) -> None: - from .kernels.gdn_chunk_cutile import chunk_gated_delta_rule_fwd, l2norm_fwd - - want_state = "final_state" in node.outputs - K = node.inputs["q"].dim[-1] - q, k, v = nb.inputs["q"], nb.inputs["k"], nb.inputs["v"] - g, beta = nb.inputs["g"], nb.inputs["beta"] - scale = node.params.get("scale") or K**-0.5 - if node.params.get("use_qk_l2norm", False): - q, _ = l2norm_fwd(q, out=bufs["q_norm"], rstd_out=bufs["q_rstd"], stream=stream) - k, _ = l2norm_fwd(k, out=bufs["k_norm"], rstd_out=bufs["k_rstd"], stream=stream) - # terminal pipeline buffers = the caller's output buffers - bufs["o"] = nb.outputs["O"] - if want_state: - bufs["final_state"] = nb.outputs["final_state"] - gate_kwargs = {} - if node.params.get("safe_gate", False): - gate_kwargs = dict(use_gate_in_kernel=True, A_log=nb.inputs["a_log"], dt_bias=nb.inputs["dt_bias"]) - chunk_gated_delta_rule_fwd( - q=q, - k=k, - v=v, - g=g, - beta=beta, - scale=scale, - initial_state=nb.inputs.get("initial_state"), - **gate_kwargs, - output_final_state=want_state, - cu_seqlens=nb.inputs["cu_seqlens"], - chunk_indices=bufs["chunk_table"], - bufs=bufs, + self.common.ensure_cuda_context(stream) + ws = Workspace.over(variant_pack, self.ws_bytes, self.plan_name) + region = dict(zip(self.carve_names, ws.carve(self.carve))) + self.common.build_chunk_table( + region["chunk_table"], + region["chunk_count"], + region["chunk_offsets"], + nb["cu_seqlens"], + self.n_seqs, + self.bt_chunk, + self.bound, stream=stream, ) - - def execute_bwd(self, node, nb, bufs, stream) -> None: - from .kernels.common import add_inplace, reshaped - from .kernels.gdn_chunk_cutile import ( - RCP_LN2, - BT_CHUNK, - chunk_gated_delta_rule_bwd, - chunk_gated_delta_rule_fwd_intra, - chunk_local_cumsum, - l2norm_bwd, - l2norm_fwd, + for name, port in self.plant: + region[name] = nb[port] + if self.is_bwd: + self.execute_bwd(nb, region, stream) + else: + self.execute_fwd(nb, region, stream) + + def execute_fwd(self, nb, region, stream) -> None: + gate = dict(use_gate_in_kernel=True, A_log=nb["a_log"], dt_bias=nb["dt_bias"]) if self.safe_gate else {} + self.kernels.chunk_gated_delta_rule( + nb["q"], + nb["k"], + nb["v"], + nb["g"], + nb["beta"], + scale=self.scale, + initial_state=nb.get("initial_state"), + output_final_state=self.want_state, + use_qk_l2norm_in_kernel=self.l2norm, + cu_seqlens=nb["cu_seqlens"], + chunk_indices=region["chunk_table"], + bufs=region, + stream=stream, + **gate, ) - H, K = node.inputs["q"].dim[1], node.inputs["q"].dim[-1] - HV = node.inputs["v"].dim[1] - q, k, v = nb.inputs["q"], nb.inputs["k"], nb.inputs["v"] - g, beta, do = nb.inputs["g"], nb.inputs["beta"], nb.inputs["dO"] - cu_seqlens = nb.inputs["cu_seqlens"] - initial_state = nb.inputs.get("initial_state") - dstate_in = nb.inputs.get("d_final_state") - chunk_indices = bufs["chunk_table"] - scale = node.params.get("scale") or K**-0.5 - l2norm = bool(node.params.get("use_qk_l2norm", False)) - if l2norm: - q, q_rstd = l2norm_fwd(q, out=bufs["q_norm"], rstd_out=bufs["q_rstd"], stream=stream) - k, k_rstd = l2norm_fwd(k, out=bufs["k_norm"], rstd_out=bufs["k_rstd"], stream=stream) - - # terminal pipeline buffers = the caller's output buffers - if not l2norm: - bufs["dq" if HV == H else "dq_hred"] = nb.outputs["dQ"] - bufs["dk" if HV == H else "dk_hred"] = nb.outputs["dK"] - bufs["wy_dv"] = nb.outputs["dV"] - bufs["dg_cum"] = nb.outputs["dG"] - bufs["db"] = nb.outputs["dBeta"] - if initial_state is not None: - bufs["dstate0"] = nb.outputs["d_initial_state"] - - # recompute the forward's cumulative gate and intra-chunk WY matrix - g_cum = chunk_local_cumsum(g, chunk_size=BT_CHUNK, scale=RCP_LN2, cu_seqlens=cu_seqlens, chunk_indices=chunk_indices, out=bufs["g_cum"], stream=stream) - _, _, A = chunk_gated_delta_rule_fwd_intra( - k=k, v=v, g=g_cum, beta=beta, cu_seqlens=cu_seqlens, chunk_indices=chunk_indices, bufs=bufs, compute_wu=False, stream=stream - ) - dq, dk, dk2, _, _, _, _, _, _ = chunk_gated_delta_rule_bwd( - q=q, - k=k, - v=v, - g=g_cum, - beta=beta, - A=A, - scale=scale, - initial_state=initial_state, - do=do, - dstate_in=dstate_in, - cu_seqlens=cu_seqlens, - chunk_indices=chunk_indices, - bufs=bufs, + def execute_bwd(self, nb, region, stream) -> None: + self.kernels.chunk_gated_delta_rule_grad( + nb["q"], + nb["k"], + nb["v"], + nb["g"], + nb["beta"], + nb["dO"], + dstate_in=nb.get("d_final_state"), + scale=self.scale, + initial_state=nb.get("initial_state"), + use_qk_l2norm_in_kernel=self.l2norm, + cu_seqlens=nb["cu_seqlens"], + chunk_indices=region["chunk_table"], + bufs=region, stream=stream, ) - if l2norm: - l2norm_bwd(q, q_rstd, dq, out=nb.outputs["dQ"], bufs=bufs, stream=stream) - l2norm_bwd(k, k_rstd, dk, dy2=dk2, out=nb.outputs["dK"], bufs=bufs, stream=stream) - else: - # dK/dK2 are the head-reduced finals for GVA, HV-head for MHA - n_dk = 1 - for s in dk.shape: - n_dk *= int(s) - add_inplace(reshaped(dk, (n_dk,)), reshaped(dk2, (n_dk,)), n_dk, stream=stream) class GdnCuTileEngine(BaseEngine): """cuTile chunked-kernel backend for single-node GDN graphs (THD layout).""" name = "gdn_cutile" - behavior_notes = (behavior_note.RUNTIME_COMPILATION,) # JIT-compiled + autotuned on first execute per shape + behavior_notes = (behavior_note.RUNTIME_COMPILATION,) def check_support(self, graph: "pygraph") -> None: import cudnn - if buffers.current_sm() is None: - raise NotImplementedError("GdnCuTileEngine requires a CUDA device") try: - from cuda.bindings import runtime - - err, cudart_version = runtime.cudaRuntimeGetVersion() - if int(err) != 0: - raise NotImplementedError(f"GdnCuTileEngine: cudaRuntimeGetVersion failed ({err})") - except ImportError as exc: - raise NotImplementedError(f"GdnCuTileEngine requires cuda.bindings: {exc}") from exc - if cudart_version < 13030: - raise NotImplementedError(f"GdnCuTileEngine requires CUDA 13.3+ (found {cudart_version})") - try: - from .kernels.gdn_chunk_cutile import chunk_gated_delta_rule # noqa: F401 — availability probe: ImportError = decline + from .kernels.gdn import chunk_gated_delta_rule # noqa: F401 — availability probe: ImportError = decline except ImportError as exc: raise NotImplementedError(f"GdnCuTileEngine requires the cuda.tile runtime: {exc}") from exc facts = graph._facts_for(analyze) - if facts is None or facts.op != "GDN": - raise NotImplementedError("GdnCuTileEngine supports exactly one GDN/GDN_BWD node") - if facts.invalid: - raise NotImplementedError(f"GdnCuTileEngine: {facts.invalid}") - if facts.checkpoint_every_n_tokens > 0 or facts.wants_state_checkpoints: - raise NotImplementedError("GdnCuTileEngine: per-chunk state_checkpoints output is not supported") + cutile_la_gate("GdnCuTileEngine", facts, "GDN", cudnn.data_type.FLOAT) if facts.is_bwd and facts.safe_gate: raise NotImplementedError("GdnCuTileEngine: safe_gate is forward-only") - f32 = cudnn.data_type.FLOAT - for port, got in ( - ("initial_state", facts.state_dtype), - ("final_state", facts.final_state_dtype), - ("d_final_state", facts.d_final_state_dtype), - ("d_initial_state", facts.d_initial_state_dtype), - ("a_log", facts.a_log_dtype), - ("dt_bias", facts.dt_bias_dtype), - ): - if got not in (f32, None): - raise NotImplementedError(f"GdnCuTileEngine: '{port}' must be fp32 (callers convert), got {got}") - if not facts.uniform_io: - raise NotImplementedError("GdnCuTileEngine: q/k/v dtypes must match") - if facts.io_dtype not in (cudnn.data_type.HALF, cudnn.data_type.BFLOAT16, None): - raise NotImplementedError(f"GdnCuTileEngine: q/k/v must be fp16/bf16, got {facts.io_dtype}") - if not facts.thd_layout: - raise NotImplementedError("GdnCuTileEngine: q/k/v must be THD [total_T, heads, dim]") - if facts.h_k != facts.h_q: - raise NotImplementedError(f"GdnCuTileEngine: q and k head counts differ ({facts.h_q} vs {facts.h_k})") - if facts.h_q and facts.h_v % facts.h_q != 0: - raise NotImplementedError( - f"GdnCuTileEngine: v heads ({facts.h_v}) must be a multiple of q heads ({facts.h_q}; GQA-style v broadcast is FROST-only)" - ) - if facts.d_qk > 256: - raise NotImplementedError(f"GdnCuTileEngine: head dim K must be <= 256, got {facts.d_qk}") - if facts.cu_dtype not in (cudnn.data_type.INT32, None): - raise NotImplementedError("GdnCuTileEngine: cu_seqlens must be int32 (the device-side table builder reads it directly)") - io = facts.io_dtype - f32 = cudnn.data_type.FLOAT - if not facts.is_bwd: - out_dtypes = {"O": (facts.o_dtype, io), "final_state": (facts.final_state_dtype, f32)} - else: - out_dtypes = { - "dQ": (facts.dq_dtype, io), - "dK": (facts.dk_dtype, io), - "dV": (facts.dv_dtype, io), - "dG": (facts.dg_dtype, f32), - "dBeta": (facts.dbeta_dtype, facts.beta_dtype), - "d_initial_state": (facts.d_initial_state_dtype, f32), - } - if facts.has_initial_state != facts.wants_d_initial_state: - raise NotImplementedError("GdnCuTileEngine: d_initial_state output must be requested iff initial_state is given") - for port, (got, want) in out_dtypes.items(): - if got is not None and got not in (want, None): - raise NotImplementedError(f"GdnCuTileEngine: output '{port}' must be {want} (written in place), got {got}") + if facts.use_beta_sigmoid: + raise NotImplementedError("GdnCuTileEngine: use_beta_sigmoid has no cuTile path (the FROST GDN engine serves it)") def build_plan(self, graph, plan, ctx=None) -> CompiledPlan: return GdnCuTilePlan(graph) diff --git a/python/cudnn/linear_attention/cutile/kda_engine.py b/python/cudnn/linear_attention/cutile/kda_engine.py index ca5a769f3..6e844b58a 100644 --- a/python/cudnn/linear_attention/cutile/kda_engine.py +++ b/python/cudnn/linear_attention/cutile/kda_engine.py @@ -2,220 +2,220 @@ # SPDX-License-Identifier: Apache-2.0 """cuTile KDA engine: KDA / KDA_BWD nodes on the chunked cuTile kernels -(``kernels/kda_chunk_cutile``).""" +(``kernels/kda``).""" from typing import TYPE_CHECKING from cudnn import behavior_note -from cudnn.engines.base import BaseEngine, CompiledPlan, resolve_node_buffers +from cudnn.engines.base import BaseEngine, CompiledPlan, bind_ports from cudnn.graph_types import NodeType from cudnn.frost import buffers -from cudnn.frost.workspace import Workspace -from ..graph_analyzer import check_layouts_compact, analyze, expect_table, to_buffer_dtype +from cudnn.frost.workspace import Workspace, WorkspaceLayout, carve_plan +from ..graph_analyzer import analyze, to_buffer_dtype +from .engine import check_layouts_compact, cutile_la_gate, expect_table -# entry base-alignment expectations; ports not listed assume 16 -CUTILE_ALIGN = {"cu_seqlens": 4, "a_log": 4, "beta": 4} +GATE_LOWER_BOUND_RANGE = (-5.0, 0.0) if TYPE_CHECKING: from cudnn._pygraph import pygraph -def node_ws_layout(node): - """Static carve plan for one node's ``chunk_kda`` pipeline intermediates: - name -> (offset, dtype-name, shape). The chunk count is data-dependent - (varlen), so chunk-indexed entries are sized and SHAPED at the bound - ``cdiv(total, 64) + N`` — the device-built table's sentinel tail keeps - bound-gridded launches inert past the real count. A KDA_BWD node carries - the union of the forward re-run's intermediates and the backward - temporaries, so both live disjointly in one carve. Terminal pipeline - buffers (the forward's ``o``/``final_state``; the backward's boundary casts / - l2norm outputs, ``dv2`` and ``dstate0``) are NOT carved — execute plants the - caller's output buffers under those names.""" - from .kernels.kda_chunk_cutile import BT_CHUNK, cdiv, next_power_of_2 - - q, v, g, cu = (node.inputs[p] for p in ("q", "v", "g", "cu_seqlens")) - total, H, K = q.dim - HV, V = v.dim[1], v.dim[2] - N = cu.dim[0] - 1 - io = to_buffer_dtype(q.get_data_type()) - f32 = "float32" - BC = 32 if K >= 64 else 16 # fwd_intra sub-chunk (see chunk_kda_fwd_intra) - NT_bound = cdiv(total, BT_CHUNK) + N - l2norm = bool(node.params.get("use_qk_l2norm", False)) - - size = 0 - table = {} - - def carve(name, dtype, shape): - nonlocal size - nbytes = buffers.DTYPE_ITEMSIZE[dtype] - for s in shape: - nbytes *= int(s) - table[name] = (size, dtype, tuple(int(s) for s in shape)) - size += (nbytes + 127) & ~127 # 128B-aligned sequential carve - - carve("chunk_table", "int32", (NT_bound, 2)) - carve("chunk_count", "int32", (1,)) - carve("chunk_offsets", "int32", (N + 1,)) - carve("dummy", "int32", (4,)) # inert stub backing for absent optional kernel args - - # chunk_kda forward (also re-run inside KDA_BWD) - carve("g_cum", f32, (total, HV, K)) - carve("Aqk", io, (total, HV, BT_CHUNK)) - carve("Akk", io, (total, HV, BT_CHUNK)) - carve("Akkd", f32, (total, HV, BC)) - carve("w", io, (total, HV, K)) - carve("u", io, (total, HV, V)) - carve("qg", io, (total, HV, K)) - carve("kg", io, (total, HV, K)) - carve("state_checkpoints", io, (NT_bound, HV, K, V)) - carve("v_new", io, (total, HV, V)) - if node.params.get("use_beta_sigmoid", False): - carve("beta_sig", f32, (total, HV)) - if node.node_type == NodeType.KDA_BWD: - carve("o", io, (total, HV, V)) # discarded output of the forward re-run - if l2norm: - carve("q_norm", io, (total * H, K)) - carve("q_rstd", f32, (total * H,)) - carve("k_norm", io, (total * H, K)) - carve("k_rstd", f32, (total * H,)) - if node.node_type == NodeType.KDA_BWD: - carve("dAqk", f32, (total, HV, BT_CHUNK)) - carve("dv_dAv", io, (total, HV, V)) - carve("dstate", io, (NT_bound, HV, K, V)) - carve("dv_dstate_u", io, (total, HV, V)) - carve("dq", f32, (total, HV, K)) - carve("dk", f32, (total, HV, K)) - carve("dg", f32, (total, HV, K)) - if to_buffer_dtype(node.inputs["beta"].get_data_type()) != f32: - carve("db", f32, (total, HV)) - carve("dAkk", f32, (total, HV, BT_CHUNK)) - carve("dq2", f32, (total, HV, K)) - carve("dk2", f32, (total, HV, K)) - if HV != H: - carve("dq_hred", f32, (total, H, K)) - carve("dk_hred", f32, (total, H, K)) - NK = cdiv(K, min(64, next_power_of_2(K))) # bwd_intra K-split (see chunk_kda_bwd_intra) - carve("db2", f32, (NK, total, HV)) - carve("dg2", f32, (total, HV, K)) - if to_buffer_dtype(g.get_data_type()) != f32: - carve("dg_cum", f32, (total, HV, K)) - return size, table - - class KdaCuTilePlan(CompiledPlan): - """Carve plan over the caller's workspace: the layout is static per node; - the buffer arrives with every execute.""" + """Carve plan over the caller's workspace, driven from the normalized + variant pack: the carve layout, geometry, scale and kernel module are fixed + per node at build; between executes only the buffer addresses move.""" + + takes_variant_pack = True + plan_name = "KdaCuTileEngine" def __init__(self, graph): - self.layouts = [(node, *node_ws_layout(node)) for node in graph.nodes] - self.expects = {node: expect_table(node, CUTILE_ALIGN) for node in graph.nodes} - # nodes execute sequentially, each re-carving the same buffer - self.ws_bytes = max(nbytes for _, nbytes, _ in self.layouts) + from .kernels import common + from .kernels import kda as kernels + + (node,) = graph.nodes + self.kernels = kernels + self.common = common + self.is_bwd = node.node_type == NodeType.KDA_BWD + + q, v, g, cu = (node.inputs[p] for p in ("q", "v", "g", "cu_seqlens")) + total, H, K = (int(d) for d in q.dim) + HV, V = int(v.dim[1]), int(v.dim[2]) + N = int(cu.dim[0]) - 1 + io = to_buffer_dtype(q.get_data_type()) + f32 = "float32" + isz = buffers.DTYPE_ITEMSIZE[io] + BT = kernels.BT_CHUNK + BC = 32 if K >= 64 else 16 + NT_bound = common.cdiv(total, BT) + N + l2norm = bool(node.params.get("use_qk_l2norm", False)) + + layout = WorkspaceLayout() + regions = [ + ("chunk_table", layout.add(NT_bound * 2 * 4), "int32", (NT_bound, 2)), + ("chunk_count", layout.add(4), "int32", (1,)), + ("chunk_offsets", layout.add((N + 1) * 4), "int32", (N + 1,)), + ("dummy", layout.add(16), "int32", (4,)), + ("g_cum", layout.add(total * HV * K * 4), f32, (total, HV, K)), + ("Aqk", layout.add(total * HV * BT * isz), io, (total, HV, BT)), + ("Akk", layout.add(total * HV * BT * isz), io, (total, HV, BT)), + ("Akkd", layout.add(total * HV * BC * 4), f32, (total, HV, BC)), + ("w", layout.add(total * HV * K * isz), io, (total, HV, K)), + ("u", layout.add(total * HV * V * isz), io, (total, HV, V)), + ("qg", layout.add(total * HV * K * isz), io, (total, HV, K)), + ("kg", layout.add(total * HV * K * isz), io, (total, HV, K)), + ("state_checkpoints", layout.add(NT_bound * HV * K * V * isz), io, (NT_bound, HV, K, V)), + ("v_new", layout.add(total * HV * V * isz), io, (total, HV, V)), + ] + if node.params.get("use_beta_sigmoid", False): + regions.append(("beta_sig", layout.add(total * HV * 4), f32, (total, HV))) + if self.is_bwd: + regions.append(("o", layout.add(total * HV * V * isz), io, (total, HV, V))) + if l2norm: + regions += [ + ("q_norm", layout.add(total * H * K * isz), io, (total * H, K)), + ("q_rstd", layout.add(total * H * 4), f32, (total * H,)), + ("k_norm", layout.add(total * H * K * isz), io, (total * H, K)), + ("k_rstd", layout.add(total * H * 4), f32, (total * H,)), + ] + if self.is_bwd: + regions += [ + ("dAqk", layout.add(total * HV * BT * 4), f32, (total, HV, BT)), + ("dv_dAv", layout.add(total * HV * V * isz), io, (total, HV, V)), + ("dstate", layout.add(NT_bound * HV * K * V * isz), io, (NT_bound, HV, K, V)), + ("dv_dstate_u", layout.add(total * HV * V * isz), io, (total, HV, V)), + ("dq", layout.add(total * HV * K * 4), f32, (total, HV, K)), + ("dk", layout.add(total * HV * K * 4), f32, (total, HV, K)), + ("dg", layout.add(total * HV * K * 4), f32, (total, HV, K)), + ] + if to_buffer_dtype(node.inputs["beta"].get_data_type()) != f32: + regions.append(("db", layout.add(total * HV * 4), f32, (total, HV))) + regions += [ + ("dAkk", layout.add(total * HV * BT * 4), f32, (total, HV, BT)), + ("dq2", layout.add(total * HV * K * 4), f32, (total, HV, K)), + ("dk2", layout.add(total * HV * K * 4), f32, (total, HV, K)), + ] + if HV != H: + regions += [ + ("dq_hred", layout.add(total * H * K * 4), f32, (total, H, K)), + ("dk_hred", layout.add(total * H * K * 4), f32, (total, H, K)), + ] + NK = common.cdiv(K, min(64, common.next_power_of_2(K))) + regions += [ + ("db2", layout.add(NK * total * HV * 4), f32, (NK, total, HV)), + ("dg2", layout.add(total * HV * K * 4), f32, (total, HV, K)), + ] + if to_buffer_dtype(g.get_data_type()) != f32: + regions.append(("dg_cum", layout.add(total * HV * K * 4), f32, (total, HV, K))) + + self.ws_bytes = layout.size + self.carve_names = [name for name, _off, _dtype, _shape in regions] + self.carve = carve_plan(self.plan_name, [(off, dtype, shape) for _name, off, dtype, shape in regions]) + self.expect = expect_table(node) + self.n_seqs = N + self.bound = NT_bound + self.bt_chunk = BT + + scale = node.params.get("scale") + self.scale = float(scale or K**-0.5) if self.is_bwd else scale + self.l2norm = l2norm + self.use_beta_sigmoid = bool(node.params.get("use_beta_sigmoid", False)) + self.safe_gate = bool(node.params.get("safe_gate", False)) + self.lower_bound = float(node.params.get("gate_lower_bound") or GATE_LOWER_BOUND_RANGE[0]) + self.want_state = "final_state" in node.outputs + + if self.is_bwd: + carved = set(self.carve_names) + plant = [ + ("dq_cast", "dQ"), + ("dk_cast", "dK"), + ("dv2", "dV"), + ("dg_cast" if "dg_cum" in carved else "dg_cum", "dG"), + ("db_cast" if "db" in carved else "db", "dBeta"), + ] + if l2norm: + plant += [("dq_l2", "dQ"), ("dk_l2", "dK")] + if "initial_state" in node.inputs: + plant.append(("dstate0", "d_initial_state")) + else: + plant = [("o", "O")] + ([("final_state", "final_state")] if self.want_state else []) + self.plant = tuple(plant) + + self.ports = None + self.names = None + self.indices = None def get_workspace_size(self) -> int: return self.ws_bytes - def execute(self, graph, uid_to_data, ctx) -> None: - from .kernels.common import build_chunk_table, ensure_cuda_context - from .kernels.kda_chunk_cutile import BT_CHUNK - - node_buffers = resolve_node_buffers(graph, uid_to_data) + def execute(self, graph, variant_pack, ctx) -> None: + if self.ports is None: + self.ports = bind_ports(graph, variant_pack) + (slots,) = self.ports.values() + self.names = list(slots.inputs) + list(slots.outputs) + self.indices = list(slots.inputs.values()) + list(slots.outputs.values()) + views = variant_pack.operands(self.indices) + check_layouts_compact(self.plan_name, self.expect, self.names, views) + nb = dict(zip(self.names, views)) stream = ctx.stream if ctx.stream is not None else 0 - ensure_cuda_context(stream) - ws = Workspace(ctx.workspace, self.ws_bytes, "KdaCuTileEngine") - for node, _nbytes, table in self.layouts: - nb = node_buffers[node] - check_layouts_compact("KdaCuTileEngine", self.expects[node], nb) - cu_seqlens = nb.inputs["cu_seqlens"] - N = node.inputs["cu_seqlens"].dim[0] - 1 - bufs = {name: ws.view(off, dt, shape) for name, (off, dt, shape) in table.items()} - bound = bufs["chunk_table"].shape[0] - build_chunk_table(bufs["chunk_table"], bufs["chunk_count"], bufs["chunk_offsets"], cu_seqlens, N, BT_CHUNK, bound, stream=stream) - if node.node_type == NodeType.KDA: - self.execute_fwd(node, nb, bufs, stream) - else: - self.execute_bwd(node, nb, bufs, stream) - - def execute_fwd(self, node, nb, bufs, stream) -> None: - from .kernels.kda_chunk_cutile import chunk_kda - - want_state = "final_state" in node.outputs - q, k, v = nb.inputs["q"], nb.inputs["k"], nb.inputs["v"] - g, beta = nb.inputs["g"], nb.inputs["beta"] - # terminal pipeline buffers = the caller's output buffers - bufs["o"] = nb.outputs["O"] - if want_state: - bufs["final_state"] = nb.outputs["final_state"] - raw_gate_kwargs = {} - if node.params.get("use_beta_sigmoid", False): - raw_gate_kwargs["use_beta_sigmoid_in_kernel"] = True - if node.params.get("safe_gate", False): - raw_gate_kwargs.update( - safe_gate=True, - use_gate_in_kernel=True, - lower_bound=float(node.params.get("gate_lower_bound") or -5.0), - A_log=nb.inputs["a_log"], - dt_bias=nb.inputs["dt_bias"], - ) - chunk_kda( - q, - k, - v, - g, - beta, - scale=node.params.get("scale"), - initial_state=nb.inputs.get("initial_state"), - output_final_state=want_state, - use_qk_l2norm_in_kernel=bool(node.params.get("use_qk_l2norm", False)), - cu_seqlens=nb.inputs["cu_seqlens"], - chunk_indices=bufs["chunk_table"], - bufs=bufs, + self.common.ensure_cuda_context(stream) + ws = Workspace.over(variant_pack, self.ws_bytes, self.plan_name) + region = dict(zip(self.carve_names, ws.carve(self.carve))) + self.common.build_chunk_table( + region["chunk_table"], + region["chunk_count"], + region["chunk_offsets"], + nb["cu_seqlens"], + self.n_seqs, + self.bt_chunk, + self.bound, stream=stream, - **raw_gate_kwargs, + ) + for name, port in self.plant: + region[name] = nb[port] + if self.is_bwd: + self.execute_bwd(nb, region, stream) + else: + self.execute_fwd(nb, region, stream) + + def execute_fwd(self, nb, region, stream) -> None: + gate = {} + if self.use_beta_sigmoid: + gate["use_beta_sigmoid_in_kernel"] = True + if self.safe_gate: + gate.update(safe_gate=True, use_gate_in_kernel=True, lower_bound=self.lower_bound, A_log=nb["a_log"], dt_bias=nb["dt_bias"]) + self.kernels.chunk_kda( + nb["q"], + nb["k"], + nb["v"], + nb["g"], + nb["beta"], + scale=self.scale, + initial_state=nb.get("initial_state"), + output_final_state=self.want_state, + use_qk_l2norm_in_kernel=self.l2norm, + cu_seqlens=nb["cu_seqlens"], + chunk_indices=region["chunk_table"], + bufs=region, + stream=stream, + **gate, ) - def execute_bwd(self, node, nb, bufs, stream) -> None: - from .kernels.common import reshaped - from .kernels.kda_chunk_cutile import chunk_kda_grad - - total, H, K = node.inputs["q"].dim - cu_seqlens = nb.inputs["cu_seqlens"] - initial_state = nb.inputs.get("initial_state") - dstate_in = nb.inputs.get("d_final_state") - do, q, k, v = nb.inputs["dO"], nb.inputs["q"], nb.inputs["k"], nb.inputs["v"] - g, beta = nb.inputs["g"], nb.inputs["beta"] - scale = node.params.get("scale") or K**-0.5 - - # terminal pipeline buffers = the caller's output buffers - dQ_out = nb.outputs["dQ"] - dK_out = nb.outputs["dK"] - if node.params.get("use_qk_l2norm", False): - bufs["dq_l2"] = reshaped(dQ_out, (total * H, K)) - bufs["dk_l2"] = reshaped(dK_out, (total * H, K)) - bufs["dq_cast"] = dQ_out - bufs["dk_cast"] = dK_out - bufs["dv2"] = nb.outputs["dV"] - bufs["dg_cast" if "dg_cum" in bufs else "dg_cum"] = nb.outputs["dG"] - bufs["db_cast" if "db" in bufs else "db"] = nb.outputs["dBeta"] - if initial_state is not None: - bufs["dstate0"] = nb.outputs["d_initial_state"] - - chunk_kda_grad( - q, - k, - v, - g, - beta, - do, - dstate_in=dstate_in, - scale=scale, - initial_state=initial_state, - use_qk_l2norm_in_kernel=bool(node.params.get("use_qk_l2norm", False)), - cu_seqlens=cu_seqlens, - chunk_indices=bufs["chunk_table"], - bufs=bufs, + def execute_bwd(self, nb, region, stream) -> None: + self.kernels.chunk_kda_grad( + nb["q"], + nb["k"], + nb["v"], + nb["g"], + nb["beta"], + nb["dO"], + dstate_in=nb.get("d_final_state"), + scale=self.scale, + initial_state=nb.get("initial_state"), + use_qk_l2norm_in_kernel=self.l2norm, + cu_seqlens=nb["cu_seqlens"], + chunk_indices=region["chunk_table"], + bufs=region, stream=stream, ) @@ -224,86 +224,22 @@ class KdaCuTileEngine(BaseEngine): """cuTile chunked-kernel backend for single-node KDA graphs (THD layout).""" name = "kda_cutile" - behavior_notes = (behavior_note.RUNTIME_COMPILATION,) # JIT-compiled + autotuned on first execute per shape + behavior_notes = (behavior_note.RUNTIME_COMPILATION,) def check_support(self, graph: "pygraph") -> None: - import cudnn - - if buffers.current_sm() is None: - raise NotImplementedError("KdaCuTileEngine requires a CUDA device") - try: - from cuda.bindings import runtime - - err, cudart_version = runtime.cudaRuntimeGetVersion() - if int(err) != 0: - raise NotImplementedError(f"KdaCuTileEngine: cudaRuntimeGetVersion failed ({err})") - except ImportError as exc: - raise NotImplementedError(f"KdaCuTileEngine requires cuda.bindings: {exc}") from exc - if cudart_version < 13030: - raise NotImplementedError(f"KdaCuTileEngine requires CUDA 13.3+ (found {cudart_version})") try: - from .kernels.kda_chunk_cutile import chunk_kda # noqa: F401 — availability probe: ImportError = decline + from .kernels.kda import chunk_kda # noqa: F401 — availability probe: ImportError = decline except ImportError as exc: raise NotImplementedError(f"KdaCuTileEngine requires the cuda.tile runtime: {exc}") from exc facts = graph._facts_for(analyze) - if facts is None or facts.op != "KDA": - raise NotImplementedError("KdaCuTileEngine supports exactly one KDA/KDA_BWD node") - if facts.invalid: - raise NotImplementedError(f"KdaCuTileEngine: {facts.invalid}") - if facts.checkpoint_every_n_tokens > 0 or facts.wants_state_checkpoints: - raise NotImplementedError("KdaCuTileEngine: per-chunk state_checkpoints output is not supported") + cutile_la_gate("KdaCuTileEngine", facts, "KDA", facts.g_dtype if facts is not None else None) if facts.is_bwd and (facts.safe_gate or facts.use_beta_sigmoid): raise NotImplementedError("KdaCuTileEngine: raw-logit gate modes (safe_gate / use_beta_sigmoid) are forward-only") - f32 = cudnn.data_type.FLOAT - for port, got in ( - ("initial_state", facts.state_dtype), - ("final_state", facts.final_state_dtype), - ("d_final_state", facts.d_final_state_dtype), - ("d_initial_state", facts.d_initial_state_dtype), - ("a_log", facts.a_log_dtype), - ("dt_bias", facts.dt_bias_dtype), - ): - if got not in (f32, None): - raise NotImplementedError(f"KdaCuTileEngine: '{port}' must be fp32 (callers convert), got {got}") - node = next(iter(graph.nodes), None) - glb = node.params.get("gate_lower_bound") if node is not None else None - if glb is not None and glb is not False and not (-5.0 <= float(glb) < 0): - raise NotImplementedError(f"KdaCuTileEngine: gate_lower_bound must be in [-5, 0) (chunk_kda log-gate floor), got {glb}") - if not facts.uniform_io: - raise NotImplementedError("KdaCuTileEngine: q/k/v dtypes must match") - if facts.io_dtype not in (cudnn.data_type.HALF, cudnn.data_type.BFLOAT16, None): - raise NotImplementedError(f"KdaCuTileEngine: q/k/v must be fp16/bf16, got {facts.io_dtype}") - if not facts.thd_layout: - raise NotImplementedError("KdaCuTileEngine: q/k/v must be THD [total_T, heads, dim]") - if facts.h_k != facts.h_q: - raise NotImplementedError(f"KdaCuTileEngine: q and k head counts differ ({facts.h_q} vs {facts.h_k})") - if facts.h_q and facts.h_v % facts.h_q != 0: - raise NotImplementedError( - f"KdaCuTileEngine: v heads ({facts.h_v}) must be a multiple of q heads ({facts.h_q}; GQA-style v broadcast is FROST-only)" - ) - if facts.d_qk > 256: - raise NotImplementedError(f"KdaCuTileEngine: head dim K must be <= 256, got {facts.d_qk}") - if facts.cu_dtype not in (cudnn.data_type.INT32, None): - raise NotImplementedError("KdaCuTileEngine: cu_seqlens must be int32 (the device-side table builder reads it directly)") - io = facts.io_dtype - f32 = cudnn.data_type.FLOAT - if not facts.is_bwd: - out_dtypes = {"O": (facts.o_dtype, io), "final_state": (facts.final_state_dtype, f32)} - else: - out_dtypes = { - "dQ": (facts.dq_dtype, io), - "dK": (facts.dk_dtype, io), - "dV": (facts.dv_dtype, io), - "dG": (facts.dg_dtype, facts.g_dtype), - "dBeta": (facts.dbeta_dtype, facts.beta_dtype), - "d_initial_state": (facts.d_initial_state_dtype, f32), - } - if facts.has_initial_state != facts.wants_d_initial_state: - raise NotImplementedError("KdaCuTileEngine: d_initial_state output must be requested iff initial_state is given") - for port, (got, want) in out_dtypes.items(): - if got is not None and got not in (want, None): - raise NotImplementedError(f"KdaCuTileEngine: output '{port}' must be {want} (written in place), got {got}") + low, high = GATE_LOWER_BOUND_RANGE + glb = facts.gate_lower_bound + if glb is not None and not (low <= glb < high): + raise NotImplementedError(f"KdaCuTileEngine: gate_lower_bound must be in [{low}, {high}) (chunk_kda log-gate floor), got {glb}") def build_plan(self, graph, plan, ctx=None) -> CompiledPlan: return KdaCuTilePlan(graph) diff --git a/python/cudnn/linear_attention/cutile/kernels/__init__.py b/python/cudnn/linear_attention/cutile/kernels/__init__.py index f360e903a..9ed1f35f8 100644 --- a/python/cudnn/linear_attention/cutile/kernels/__init__.py +++ b/python/cudnn/linear_attention/cutile/kernels/__init__.py @@ -17,15 +17,23 @@ """Kernel libraries backing the python execution engines. -``gdn_chunk_cutile``: the chunked Gated DeltaNet (GDN) cuTile kernels behind +``gdn``: the chunked Gated DeltaNet (GDN) cuTile kernels behind ``GdnCuTileEngine`` (forward, backward, and the standalone building blocks — ``chunk_gated_delta_rule*``, ``chunk_local_cumsum``). -``kda_chunk_cutile``: the chunked Kimi Delta -Attention (KDA) cuTile kernels behind ``KdaCuTileEngine`` (``chunk_kda*``, -``chunk_kda_fwd_intra``, ``chunk_kda_bwd``, ``chunk_local_cumsum``). +``kda``: the chunked Kimi Delta Attention (KDA) cuTile kernels behind +``KdaCuTileEngine`` (``chunk_kda*``, ``chunk_kda_fwd_intra``, +``chunk_kda_bwd``, ``chunk_local_cumsum``). + +``common``: the glue kernels both pipelines share. + +Both kernel modules carry the same twelve sections in the same order — +helpers, then one device-kernel group per pipeline stage, then the host +launchers mirroring those groups, then the pipeline drivers. A ``*_kernel`` +name is a ``@ct.kernel``; the launcher that grids and launches it drops the +suffix. All public wrappers take THD (token-packed) tensors — ``[total_T, H, D]`` values and ``[total_T, H]`` / ``[total_T, H, K]`` gates — with required -``cu_seqlens`` / ``chunk_indices``. +``cu_seqlens`` / ``chunk_indices``, and are called only by the engines. """ diff --git a/python/cudnn/linear_attention/cutile/kernels/common.py b/python/cudnn/linear_attention/cutile/kernels/common.py index d127ba4bb..3f24ebccf 100644 --- a/python/cudnn/linear_attention/cutile/kernels/common.py +++ b/python/cudnn/linear_attention/cutile/kernels/common.py @@ -15,89 +15,62 @@ # See the License for the specific language governing permissions and # limitations under the License. -"""Small cuTile glue kernels shared by the chunked GDN/KDA pipelines. - -The pipeline hosts stitch their main kernels together with a handful of -element-wise / reduction steps (gradient accumulation, leading-axis sums, -dtype-converting copies, chunk-table building). These entries perform those -steps as plain cuTile launches over DLPack/CAI device buffers with an -explicit stream handle. - -``add_inplace`` / ``cast_copy`` treat buffers as FLAT contiguous element -ranges (the callers own the shape bookkeeping) and take the element count -explicitly; ``sum_leading`` takes a 2-D ``[r, m]`` view. Tile-tail loads are -zero-padded and stores clip at the buffer extent. +"""Helpers and small kernels shared by the chunked GDN/KDA cuTile pipelines. + +Every buffer reaching this module comes from an engine (a variant-pack operand +or a workspace carve), so it is contiguous and needs no probing. """ +from types import SimpleNamespace + import cuda.tile as ct +from cuda.tile.tune import exhaustive_search + +from cudnn.frost.buffers import DTYPE_ITEMSIZE, DeviceView, current_device_id, dtype_name as dtname, memset_zero_async ConstInt = ct.Constant[int] TILE = 2048 -def cdiv(a: int, b: int) -> int: - return (a + b - 1) // b +# --- Host helpers --------------------------------------------------------------------------------- -def zero_fill(buf, *, stream) -> None: - """Stream-ordered zero of a whole contiguous buffer (any DLPack/CAI).""" - from cudnn.frost.buffers import DTYPE_ITEMSIZE, memset_zero_async, probe +def next_power_of_2(n: int) -> int: + return 1 << (n - 1).bit_length() - ptr, shape, _strides, dtype, _dev = probe(buf) - n = DTYPE_ITEMSIZE[dtype] - for s_ in shape: - n *= int(s_) - memset_zero_async(ptr, n, stream) +def cdiv(a: int, b: int) -> int: + return (a + b - 1) // b -def reshaped(buf, target_shape): - """A ``target_shape``-d DeviceView over the same pointer (contiguous by - the engine gate's contract; the kernels derive their index rank from the - array rank).""" - from cudnn.frost.buffers import DeviceView, probe - ptr, shape, _strides, dtype, dev = probe(buf) - return DeviceView(ptr, shape, dtype, dev).reshape(tuple(target_shape)) +def zero_fill(buf, *, stream) -> None: + """Stream-ordered zero of a whole contiguous buffer.""" + memset_zero_async(buf.data_ptr(), int(buf.nbytes), stream) def dummy(dtype_name: str, bufs): - """Inert typed view over the workspace's 16-byte ``dummy`` carve, for - ABSENT optional kernel args (always paired with a flag==0, never - dereferenced). Dtype-bound so the compiled signature stays stable; the - library allocates nothing.""" - from cudnn.frost.buffers import DTYPE_ITEMSIZE, DeviceView - + """Inert typed view over the workspace's 16-byte ``dummy`` carve, for an + absent optional kernel arg. Always paired with a flag==0, never read.""" d = bufs["dummy"] return DeviceView(d.data_ptr(), (16 // DTYPE_ITEMSIZE[dtype_name],), dtype_name, d.__dlpack_device__()[1]) def opt(t, bufs, dtype_name: str = "float32"): - """Resolve an optional tensor argument to a non-null cuTile launch arg: - the buffer if present (contiguous by the engine contract), else an inert - dummy (paired with a USE_*/HAS_* integer flag). cuTile never accepts None - in launch args, so this is the required dummy-tensor-plus-flag pattern.""" + """``t`` if present, else an inert dummy: cuTile takes no None launch arg.""" if t is None: return dummy(dtype_name, bufs) return t -def dev_id(buf) -> int: - """Device ordinal of a DLPack/CAI buffer.""" - from cudnn.frost.buffers import probe - - return probe(buf)[4] - - def ensure_cuda_context(stream=0) -> None: - """Make the calling thread's CUDA driver context current. - - cuTile launches and the autotuner's driver-API timing fail on threads - whose driver context stack is empty — e.g. autograd backward worker - threads, where cudaSetDevice alone binds nothing. Prefer the launch - stream's own context; else retain + set-current the current device's - primary context (retained only when no context is bound, so at most once - per thread). Best-effort: never fatal.""" + """Bind a driver context to the calling thread when none is bound. + + ``ct.launch`` and the autotuner read the calling thread's context stack, + and an autograd backward runs on a worker thread where ``cudaSetDevice`` + has only moved the runtime's thread-local slot. Prefer the launch stream's + context, else retain the device's primary one. Best-effort: a context this + cannot establish fails at the launch, with the launch's own diagnostics.""" try: from cuda.bindings import driver as drv @@ -109,18 +82,185 @@ def ensure_cuda_context(stream=0) -> None: if err == drv.CUresult.CUDA_SUCCESS: drv.cuCtxSetCurrent(sctx) return - from cuda.bindings import runtime as rt - - err_d, dev = rt.cudaGetDevice() - if int(err_d) != 0: + device = current_device_id() + if device is None: return - err, pctx = drv.cuDevicePrimaryCtxRetain(dev) + err, pctx = drv.cuDevicePrimaryCtxRetain(device) if err == drv.CUresult.CUDA_SUCCESS: drv.cuCtxSetCurrent(pctx) except Exception: # noqa: BLE001 pass +# --- Launch tuning -------------------------------------------------------------------------------- + + +# Only the @ct.kernel launch hints (occupancy x num_worker_warps) are explored; +# the grid, args and algorithm are unchanged, so tuning never moves numerics. +launch_hint_cache: dict = {} + + +def launch_hint_configs(occ_choices, nww_choices=(4, 8)): + """occupancy x num_worker_warps grid; deduped, default (occ=1,nww=4) first.""" + seen = set() + cfgs = [] + for occ in occ_choices: + for nww in nww_choices: + key = (occ, nww) + if key in seen: + continue + seen.add(key) + cfgs.append(SimpleNamespace(occupancy=occ, num_worker_warps=nww)) + return cfgs + + +def autotuned_launch(kernel, cache_key, grid, args, occ_choices=(1, 2, 3, 4), nww_choices=(4, 8), timeout=30, stream=None): + """Launch ``kernel`` with the best launch hints for ``cache_key``. + + Tune-once/cache/launch over launch hints only (grid, args and signature are + fixed). Falls back to the base kernel when tuning fails or times out. + Default config (occ=1, nww=4) is explored first so a no-improvement shape + keeps the base behaviour. ``cache_key`` is qualified + by the kernel's own name, so two kernels sharing a key shape stay apart. + """ + stream = 0 if stream is None else stream + cache_key = (getattr(getattr(kernel, "_pyfunc", None), "__name__", repr(kernel)), cache_key) + if cache_key not in launch_hint_cache: + tuned = None + try: + configs = launch_hint_configs(occ_choices, nww_choices) + with ct.compiler_timeout(timeout): + result = exhaustive_search( + configs, + stream, + lambda cfg: grid, + kernel, + lambda cfg: args, + lambda cfg: {"occupancy": cfg.occupancy, "num_worker_warps": cfg.num_worker_warps}, + ) + best = result.best.config + tuned = kernel.replace_hints(occupancy=best.occupancy, num_worker_warps=best.num_worker_warps) + except Exception: + tuned = None + launch_hint_cache[cache_key] = tuned + + tuned = launch_hint_cache[cache_key] + if tuned is None: + ct.launch(stream, grid, kernel, args) + else: + ct.launch(stream, grid, tuned, args) + + +# --- Device helpers ------------------------------------------------------------------------------- + + +def exp(x): + return ct.exp(ct.astype(x, ct.float32)) + + +def exp2(x): + return ct.exp2(ct.astype(x, ct.float32)) + + +def softplus(x): + # softplus: where(x <= 20, log1p(exp(x)), x) + return ct.where(x <= 20.0, ct.log(1.0 + ct.exp(x)), x) + + +def tf32(a): + """ct.mma/ct.matmul do not auto-cast fp32 operands to tf32; cast + explicitly (allow-tf32 matmul semantics).""" + return ct.astype(a, ct.tfloat32) if a.dtype == ct.float32 else a + + +def ct_min(a, b): + # scalar/tile min for runtime ints (builtin `min` is whitelisted; `hasattr` is not). + return min(a, b) + + +# --- Kernels -------------------------------------------------------------------------------------- + + +@ct.kernel +def l2norm_fwd_kernel1(x, y, rstd, eps, D, BD: ConstInt): + # D > 512 path: one row per program, row length D. + i_t = ct.bid(0) + cols = ct.arange(BD, dtype=ct.int32) + mask = cols < D + + b_x = ct.astype(ct.gather(x, (i_t, cols), mask=mask, check_bounds=False, padding_value=0.0), ct.float32) + b_rstd = ct.rsqrt(ct.sum(b_x * b_x) + eps) + b_y = b_x * b_rstd + ct.scatter(y, (i_t, cols), ct.astype(b_y, y.dtype), mask=mask, check_bounds=False) + ct.scatter(rstd, (i_t,), ct.astype(b_rstd, rstd.dtype)) + + +@ct.kernel +def l2norm_fwd_kernel(x, y, rstd, eps, T, D: ConstInt, BD: ConstInt, BT: ConstInt): + # D <= 512 path: BT rows per block, BD power-of-2 cols. Block-aligned -> + # ct.load with block index + ZERO padding. + i_t = ct.bid(0) + b_x = ct.astype(ct.load(x, index=(i_t, 0), shape=(BT, BD), padding_mode=ct.PaddingMode.ZERO), ct.float32) + b_rstd = ct.rsqrt(ct.sum(b_x * b_x, axis=1) + eps) + b_y = b_x * b_rstd[:, None] + ct.store(y, index=(i_t, 0), tile=ct.astype(b_y, y.dtype)) + ct.store(rstd, index=(i_t,), tile=ct.astype(b_rstd, rstd.dtype)) + + +@ct.kernel +def l2norm_bwd_kernel1(y, rstd, dy, dy2, dx, eps, D, BD: ConstInt, HAS_DY2: ConstInt): + i_t = ct.bid(0) + cols = ct.arange(BD, dtype=ct.int32) + mask = cols < D + + b_y = ct.astype(ct.gather(y, (i_t, cols), mask=mask, check_bounds=False, padding_value=0.0), ct.float32) + b_dy = ct.astype(ct.gather(dy, (i_t, cols), mask=mask, check_bounds=False, padding_value=0.0), ct.float32) + if HAS_DY2: + b_dy2 = ct.astype(ct.gather(dy2, (i_t, cols), mask=mask, check_bounds=False, padding_value=0.0), ct.float32) + # Preserve bf16 `dk.add_(dk2)` rounding before the fp32 normalization math. + b_dy = ct.astype(ct.astype(b_dy + b_dy2, dy.dtype), ct.float32) + b_rstd = ct.astype(ct.gather(rstd, (i_t,), check_bounds=False, padding_value=0.0), ct.float32).item() + + b_dx = b_dy * b_rstd - ct.sum(b_dy * b_y) * b_y * b_rstd + ct.scatter(dx, (i_t, cols), ct.astype(b_dx, dx.dtype), mask=mask, check_bounds=False) + + +@ct.kernel +def l2norm_bwd_kernel(y, rstd, dy, dy2, dx, eps, T, D: ConstInt, BD: ConstInt, BT: ConstInt, HAS_DY2: ConstInt): + i_t = ct.bid(0) + b_y = ct.astype(ct.load(y, index=(i_t, 0), shape=(BT, BD), padding_mode=ct.PaddingMode.ZERO), ct.float32) + b_rstd = ct.astype(ct.load(rstd, index=(i_t,), shape=(BT,), padding_mode=ct.PaddingMode.ZERO), ct.float32) + b_dy = ct.astype(ct.load(dy, index=(i_t, 0), shape=(BT, BD), padding_mode=ct.PaddingMode.ZERO), ct.float32) + if HAS_DY2: + b_dy2 = ct.astype(ct.load(dy2, index=(i_t, 0), shape=(BT, BD), padding_mode=ct.PaddingMode.ZERO), ct.float32) + b_dy = ct.astype(ct.astype(b_dy + b_dy2, dy.dtype), ct.float32) + b_dot = ct.sum(b_dy * b_y, axis=1) + b_dx = b_dy * b_rstd[:, None] - b_dot[:, None] * b_y * b_rstd[:, None] + ct.store(dx, index=(i_t, 0), tile=ct.astype(b_dx, dx.dtype)) + + +@ct.kernel +def fused_beta_sigmoid_fwd_kernel(x, y, scale, n_elements, BLOCK_SIZE: ConstInt): + pid = ct.bid(0) + offs = pid * BLOCK_SIZE + ct.arange(BLOCK_SIZE, dtype=ct.int32) + mask = offs < n_elements + b_x = ct.astype(ct.gather(x, offs, mask=mask, check_bounds=False, padding_value=0.0), ct.float32) + b_y = scale * (1.0 / (1.0 + ct.exp(-b_x))) + ct.scatter(y, offs, ct.astype(b_y, y.dtype), mask=mask, check_bounds=False) + + +@ct.kernel +def fused_beta_sigmoid_bwd_kernel(x, dy, dx, scale, n_elements, BLOCK_SIZE: ConstInt): + pid = ct.bid(0) + offs = pid * BLOCK_SIZE + ct.arange(BLOCK_SIZE, dtype=ct.int32) + mask = offs < n_elements + b_x = ct.astype(ct.gather(x, offs, mask=mask, check_bounds=False, padding_value=0.0), ct.float32) + b_dy = ct.astype(ct.gather(dy, offs, mask=mask, check_bounds=False, padding_value=0.0), ct.float32) + b_y = 1.0 / (1.0 + ct.exp(-b_x)) + b_dx = b_dy * scale * b_y * (1.0 - b_y) + ct.scatter(dx, offs, ct.astype(b_dx, dx.dtype), mask=mask, check_bounds=False) + + @ct.kernel def add_inplace_kernel(dst, src, TILE: ConstInt): pid = ct.bid(0) @@ -129,11 +269,6 @@ def add_inplace_kernel(dst, src, TILE: ConstInt): ct.store(dst, index=(pid,), tile=a + b) -def add_inplace(dst, src, numel: int, *, stream) -> None: - """``dst += src`` over ``numel`` flat elements (same dtype, contiguous).""" - ct.launch(stream, (cdiv(numel, TILE),), add_inplace_kernel, (dst, src, TILE)) - - @ct.kernel def cast_copy_kernel(dst, src, TILE: ConstInt): pid = ct.bid(0) @@ -141,12 +276,6 @@ def cast_copy_kernel(dst, src, TILE: ConstInt): ct.store(dst, index=(pid,), tile=ct.astype(t, dst.dtype)) -def cast_copy(dst, src, numel: int, *, stream) -> None: - """``dst[:] = src`` over ``numel`` flat elements, converting to ``dst``'s - dtype (a plain copy when the dtypes already match).""" - ct.launch(stream, (cdiv(numel, TILE),), cast_copy_kernel, (dst, src, TILE)) - - @ct.kernel def sum_leading_kernel(dst, src, R: ConstInt, ACC: ConstInt, TILE: ConstInt): pid = ct.bid(0) @@ -159,14 +288,6 @@ def sum_leading_kernel(dst, src, R: ConstInt, ACC: ConstInt, TILE: ConstInt): ct.store(dst, index=(pid,), tile=ct.astype(acc, dst.dtype)) -def sum_leading(dst, src, r: int, m: int, *, stream, accumulate: bool = False) -> None: - """Reduce ``src`` (a 2-D ``[r, m]`` row-major buffer) over its leading - axis into ``dst`` (flat ``[m]``), accumulating in fp32. ``r`` is a - compile-time constant (small fan-ins: split partials, head groups). - ``accumulate`` adds the reduction onto ``dst`` instead of overwriting.""" - ct.launch(stream, (cdiv(m, TILE),), sum_leading_kernel, (dst, src, r, int(accumulate), TILE)) - - @ct.kernel def build_chunk_table_kernel(cu_seqlens, table, count, offsets, N: ConstInt, CS: ConstInt, BOUND: ConstInt): run = 0 @@ -183,31 +304,16 @@ def build_chunk_table_kernel(cu_seqlens, table, count, offsets, N: ConstInt, CS: last = n run = run + nc ct.store(offsets, (n + 1,), run) - # sentinel tail: (last_nonempty_seq, BOUND) decodes to a token range - # starting at or past the packed end (BOUND * CS >= total), so a consumer - # launched at the bound grid loads zero-padding and its stores clip — no - # guard needed in the consuming kernels. The sentinel must reference a - # NONEMPTY sequence: a zero-length one turns seq-derived divisors to - # zero inside consumers (device trap). + # Sentinel tail: (last_nonempty_seq, BOUND) decodes to a token range at or + # past the packed end, so a consumer gridded at BOUND loads zero-padding + # and its stores clip. It must name a NONEMPTY sequence -- a zero-length + # one turns seq-derived divisors to zero inside consumers (device trap). for j in range(run, BOUND): ct.store(table, (j * 2,), last) ct.store(table, (j * 2 + 1,), BOUND) ct.store(count, (0,), run) -def build_chunk_table(table, count, offsets, cu_seqlens, n_seqs: int, chunk_size: int, bound: int, *, stream) -> None: - """Build the per-chunk ``(sequence, intra_chunk)`` index table ON DEVICE - from ``cu_seqlens`` — no host round-trip, so the launch stays async and - capture-safe. ``table`` is a flat int32 buffer of ``2 * bound`` entries - (``bound = cdiv(total, chunk_size) + n_seqs``, shape-derived); rows past - the real chunk count are filled with an inert sentinel whose decoded - token range lies at/past the packed end, so consumers may launch their - chunk grids at ``bound`` unchanged. ``count`` (one int32) receives the - real chunk count; ``offsets`` (int32 ``[n_seqs + 1]``) receives the - per-sequence chunk prefix (``prepare_chunk_offsets`` semantics).""" - ct.launch(stream, (1,), build_chunk_table_kernel, (cu_seqlens, reshaped(table, (2 * bound,)), count, offsets, n_seqs, chunk_size, bound)) - - @ct.kernel def head_group_sum_kernel(dst, src, G: ConstInt, BT: ConstInt, BK: ConstInt): t = ct.bid(0) @@ -219,9 +325,146 @@ def head_group_sum_kernel(dst, src, G: ConstInt, BT: ConstInt, BK: ConstInt): ct.store(dst, index=(t, h, k), tile=ct.astype(acc, dst.dtype)) +# --- Launchers ------------------------------------------------------------------------------------ + + +BETA_SIGMOID_BLOCK_SIZE = 2048 + +# l2norm is a memory-bound row reduction (each row loaded once, reduced, +# written once), so higher occupancy hides DRAM latency; do not force occ=1. +L2NORM_TUNE_OCC = (1, 2, 4, 8) + + +def l2norm_fwd(x, eps: float = 1e-6, out=None, rstd_out=None, stream=None): + stream = 0 if stream is None else stream + x_shape_og = x.shape + x = x.reshape((-1, x.shape[-1])) + y = out.reshape(tuple(x.shape)) + T, D = x.shape[0], x.shape[-1] + MAX_FUSED_SIZE = 65536 // x.element_size() + BD = min(MAX_FUSED_SIZE, next_power_of_2(D)) + rstd = rstd_out.reshape((T,)) + if D <= 512: + BT = 32 + grid = (cdiv(T, BT),) + autotuned_launch( + l2norm_fwd_kernel, + ("l2norm_fwd_kernel", D, BD, BT, str(x.dtype), current_device_id()), + grid, + (x, y, rstd, float(eps), T, D, BD, BT), + occ_choices=L2NORM_TUNE_OCC, + nww_choices=(4,), + stream=stream, + ) + else: + ct.launch(stream, (T,), l2norm_fwd_kernel1, (x, y, rstd, float(eps), D, BD)) + return y.view(x_shape_og), rstd.view(x_shape_og[:-1]) + + +def l2norm_bwd( + y, + rstd, + dy, + eps: float = 1e-6, + dy2=None, + out=None, + bufs=None, + stream=None, +): + stream = 0 if stream is None else stream + y_shape_og = y.shape + y = y.reshape(-1, dy.shape[-1]) + dy = dy.reshape(-1, dy.shape[-1]) + dy2_arg = dy2.reshape(-1, dy.shape[-1]) if dy2 is not None else dummy(dtname(dy), bufs) + dx = out.reshape(tuple(y.shape)) + T, D = y.shape[0], y.shape[-1] + MAX_FUSED_SIZE = 65536 // y.element_size() + BD = min(MAX_FUSED_SIZE, next_power_of_2(D)) + rstd_flat = rstd.reshape(-1) + if D <= 512: + BT = 32 + grid = (cdiv(T, BT),) + autotuned_launch( + l2norm_bwd_kernel, + ("l2norm_bwd_kernel", D, BD, BT, str(y.dtype), int(dy2 is not None), current_device_id()), + grid, + (y, rstd_flat, dy, dy2_arg, dx, float(eps), T, D, BD, BT, int(dy2 is not None)), + occ_choices=L2NORM_TUNE_OCC, + nww_choices=(4,), + stream=stream, + ) + else: + ct.launch( + stream, + (T,), + l2norm_bwd_kernel1, + (y, rstd_flat, dy, dy2_arg, dx, float(eps), D, BD, int(dy2 is not None)), + ) + return dx.view(y_shape_og) + + +def fused_beta_sigmoid_fwd(x, scale: float = 1.0, out=None, stream=None): + stream = 0 if stream is None else stream + y = out.reshape(tuple(x.shape)) + n = x.numel() + grid = (cdiv(n, BETA_SIGMOID_BLOCK_SIZE),) + ct.launch( + stream, + grid, + fused_beta_sigmoid_fwd_kernel, + (x.reshape((-1,)), y.reshape(-1), float(scale), n, BETA_SIGMOID_BLOCK_SIZE), + ) + return y + + +def fused_beta_sigmoid_bwd(x, dy, scale: float = 1.0, out=None, stream=None): + stream = 0 if stream is None else stream + dx = out.reshape(tuple(x.shape)) + n = x.numel() + grid = (cdiv(n, BETA_SIGMOID_BLOCK_SIZE),) + ct.launch( + stream, + grid, + fused_beta_sigmoid_bwd_kernel, + (x.reshape((-1,)), dy.reshape((-1,)), dx.reshape(-1), float(scale), n, BETA_SIGMOID_BLOCK_SIZE), + ) + return dx + + +def fused_beta_sigmoid(x, scale: float = 1.0, out=None, stream=None): + """Fused ``scale * sigmoid(x)`` (fp32, written to ``out``).""" + stream = 0 if stream is None else stream + return fused_beta_sigmoid_fwd(x, scale, out=out, stream=stream) + + +def add_inplace(dst, src, numel: int, *, stream) -> None: + """``dst += src`` over ``numel`` flat elements (same dtype, contiguous).""" + ct.launch(stream, (cdiv(numel, TILE),), add_inplace_kernel, (dst, src, TILE)) + + +def cast_copy(dst, src, numel: int, *, stream) -> None: + """``dst[:] = src`` over ``numel`` flat elements, converting to ``dst``'s dtype.""" + ct.launch(stream, (cdiv(numel, TILE),), cast_copy_kernel, (dst, src, TILE)) + + +def sum_leading(dst, src, r: int, m: int, *, stream, accumulate: bool = False) -> None: + """Reduce ``src`` ``[r, m]`` over its leading axis into ``dst`` ``[m]`` in + fp32. ``r`` is compile-time (split partials, head groups); ``accumulate`` + adds onto ``dst`` instead of overwriting.""" + ct.launch(stream, (cdiv(m, TILE),), sum_leading_kernel, (dst, src, r, int(accumulate), TILE)) + + +def build_chunk_table(table, count, offsets, cu_seqlens, n_seqs: int, chunk_size: int, bound: int, *, stream) -> None: + """Build the per-chunk ``(sequence, intra_chunk)`` index table ON DEVICE, so + the launch stays async and capture-safe. ``table`` holds ``2 * bound`` + int32s; rows past the real chunk count get the inert sentinel above, so + consumers may grid at ``bound``. ``count`` receives the real chunk count, + ``offsets`` the per-sequence chunk prefix.""" + ct.launch(stream, (1,), build_chunk_table_kernel, (cu_seqlens, table.reshape((2 * bound,)), count, offsets, n_seqs, chunk_size, bound)) + + def head_group_sum(dst, src, t: int, h: int, g: int, k: int, *, stream) -> None: - """Grouped-head reduction ``dst[t, h, :] = sum_g src[t, h*g + g', :]``: - ``src`` is a 3-D ``[t, h*g, k]`` buffer, ``dst`` 3-D ``[t, h, k]``; - fp32 accumulation, ``g`` consecutive heads per group (compile-time).""" + """``dst[t, h, :] = sum_g src[t, h*g + g', :]`` in fp32, ``g`` consecutive + heads per group (compile-time).""" BT, BK = 64, 128 # padded loads + clipped stores absorb ragged t/k ct.launch(stream, (cdiv(t, BT), h, cdiv(k, BK)), head_group_sum_kernel, (dst, src, g, BT, BK)) diff --git a/python/cudnn/linear_attention/cutile/kernels/gdn_chunk_cutile.py b/python/cudnn/linear_attention/cutile/kernels/gdn.py similarity index 86% rename from python/cudnn/linear_attention/cutile/kernels/gdn_chunk_cutile.py rename to python/cudnn/linear_attention/cutile/kernels/gdn.py index a6e2da2f0..8460024e7 100644 --- a/python/cudnn/linear_attention/cutile/kernels/gdn_chunk_cutile.py +++ b/python/cudnn/linear_attention/cutile/kernels/gdn.py @@ -16,16 +16,28 @@ # limitations under the License. -import logging -from types import SimpleNamespace - import cuda.tile as ct -from cuda.tile.tune import exhaustive_search - -from .common import add_inplace, dev_id, dummy, ensure_cuda_context, head_group_sum, opt, reshaped, sum_leading, zero_fill -from cudnn.frost.buffers import dtype_name as dtname -logger = logging.getLogger(__name__) +from .common import ( + add_inplace, + autotuned_launch, + cdiv, + ct_min, + dummy, + exp, + exp2, + fused_beta_sigmoid, + head_group_sum, + l2norm_bwd, + l2norm_fwd, + next_power_of_2, + opt, + softplus, + sum_leading, + tf32, + zero_fill, +) +from cudnn.frost.buffers import current_device_id, dtype_name as dtname ConstInt = ct.Constant[int] @@ -35,33 +47,7 @@ BT_CHUNK = 64 -# Host-side utilities - - -def cdiv(a: int, b: int) -> int: - return (a + b - 1) // b - - -def next_power_of_2(n: int) -> int: - return 1 << (n - 1).bit_length() - - -# Launch-hint autotuning (occupancy x num_worker_warps) for hot kernels. - -LAUNCH_HINT_CACHE: dict = {} - -# NOTE: the upstream occupancy=2 config is EXCLUDED: on tileiras 13.2 (which -# ignores num_worker_warps hints) a freshly compiled occ=2 -# chunk_bwd_kernel_dqkwg deadlocks on its first launch (deterministic; -# occ=1/4 are fine, and occ=2 is fine on the dhu/kkt kernels). Restore the -# SimpleNamespace(occupancy=2, num_worker_warps=4) entry once the runtime is -# on tileiras >= 13.3. -LAUNCH_HINT_CONFIGS = [ - SimpleNamespace(occupancy=4, num_worker_warps=4), - SimpleNamespace(occupancy=4, num_worker_warps=8), - SimpleNamespace(occupancy=1, num_worker_warps=8), - SimpleNamespace(occupancy=1, num_worker_warps=4), -] +# --- Host helpers --------------------------------------------------------------------------------- def device_attrs(): @@ -80,55 +66,15 @@ def device_attrs(): return 0, 0 -def tuned_launch(kernel, stream, grid, args, cache_key, configs=None): - """Launch ``kernel`` with autotuned (occupancy, num_worker_warps) hints. - - Tunes once per ``cache_key`` via exhaustive_search, caches the specialized - kernel, and re-launches it on every subsequent call. On any tuning failure - (e.g. compiler error) falls back to the un-hinted kernel. ``configs`` overrides - the default (occupancy, num_worker_warps) candidate list. - """ - kname = getattr(getattr(kernel, "_pyfunc", None), "__name__", repr(kernel)) - key = (kname, cache_key) - if key not in LAUNCH_HINT_CACHE: - tuned = None - ensure_cuda_context() - try: - with ct.compiler_timeout(20): - result = exhaustive_search( - list(configs if configs is not None else LAUNCH_HINT_CONFIGS), - stream, - lambda cfg: grid, - kernel, - lambda cfg: args, - lambda cfg: {"occupancy": cfg.occupancy, "num_worker_warps": cfg.num_worker_warps}, - ) - best = result.best.config - tuned = kernel.replace_hints(occupancy=best.occupancy, num_worker_warps=best.num_worker_warps) - except Exception as e: # noqa: BLE001 - logger.warning("launch-hint autotune failed for %s: %s; using default hints", kname, e) - tuned = kernel - LAUNCH_HINT_CACHE[key] = tuned - ct.launch(stream, grid, LAUNCH_HINT_CACHE[key], args) - - -def exp(x): - return ct.exp(ct.astype(x, ct.float32)) - - -def exp2(x): - return ct.exp2(ct.astype(x, ct.float32)) +# --- Launch tuning -------------------------------------------------------------------------------- -def softplus(x): - # TODO(cutile): softplus_nv inline-asm replaced with log1p formula. - return ct.where(x < 20.0, ct.log(1.0 + ct.exp(x)), x) +# occupancy=2 is EXCLUDED: on tileiras 13.2 a freshly compiled occ=2 +# chunk_bwd_kernel_dqkwg deadlocks on its first launch (occ=1/4 are fine). +TUNE_OCC = (1, 4) -def tf32(a): - """ct.mma/ct.matmul do not auto-cast fp32 operands to tf32; cast - explicitly (allow-tf32 matmul semantics).""" - return ct.astype(a, ct.tfloat32) if a.dtype == ct.float32 else a +# --- Device helpers ------------------------------------------------------------------------------- def safe_matmul(a, b): @@ -181,93 +127,7 @@ def scatter_flat_cb(raw, off, value, numel, *, mask=None): raw.store_offset(off, value, mask=m) -# l2norm kernels - - -@ct.kernel -def l2norm_fwd_kernel1(x, y, rstd, eps, D, BD: ConstInt): - # D > 512 path: one row per program, row length D. - i_t = ct.bid(0) - cols = ct.arange(BD, dtype=ct.int32) - mask = cols < D - - b_x = ct.astype(ct.gather(x, (i_t, cols), mask=mask, check_bounds=False, padding_value=0.0), ct.float32) - b_rstd = ct.rsqrt(ct.sum(b_x * b_x) + eps) - b_y = b_x * b_rstd - ct.scatter(y, (i_t, cols), ct.astype(b_y, y.dtype), mask=mask, check_bounds=False) - ct.scatter(rstd, (i_t,), ct.astype(b_rstd, rstd.dtype)) - - -@ct.kernel -def l2norm_bwd_kernel1(y, rstd, dy, dy2, dx, eps, D, BD: ConstInt, HAS_DY2: ConstInt): - i_t = ct.bid(0) - cols = ct.arange(BD, dtype=ct.int32) - mask = cols < D - - b_y = ct.astype(ct.gather(y, (i_t, cols), mask=mask, check_bounds=False, padding_value=0.0), ct.float32) - b_dy = ct.astype(ct.gather(dy, (i_t, cols), mask=mask, check_bounds=False, padding_value=0.0), ct.float32) - if HAS_DY2: - b_dy2 = ct.astype(ct.gather(dy2, (i_t, cols), mask=mask, check_bounds=False, padding_value=0.0), ct.float32) - # Preserve bf16 `dk.add_(dk2)` rounding before the fp32 normalization math. - b_dy = ct.astype(ct.astype(b_dy + b_dy2, dy.dtype), ct.float32) - b_rstd = ct.astype(ct.gather(rstd, (i_t,), check_bounds=False, padding_value=0.0), ct.float32).item() - - b_dx = b_dy * b_rstd - ct.sum(b_dy * b_y) * b_y * b_rstd - ct.scatter(dx, (i_t, cols), ct.astype(b_dx, dx.dtype), mask=mask, check_bounds=False) - - -@ct.kernel -def l2norm_fwd_kernel(x, y, rstd, eps, T, D: ConstInt, BD: ConstInt, BT: ConstInt): - # D <= 512 path: BT rows per block, BD power-of-2 cols. Block-aligned -> - # ct.load with block index + ZERO padding. - i_t = ct.bid(0) - b_x = ct.astype(ct.load(x, index=(i_t, 0), shape=(BT, BD), padding_mode=ct.PaddingMode.ZERO), ct.float32) - b_rstd = ct.rsqrt(ct.sum(b_x * b_x, axis=1) + eps) - b_y = b_x * b_rstd[:, None] - ct.store(y, index=(i_t, 0), tile=ct.astype(b_y, y.dtype)) - ct.store(rstd, index=(i_t,), tile=ct.astype(b_rstd, rstd.dtype)) - - -@ct.kernel -def l2norm_bwd_kernel(y, rstd, dy, dy2, dx, eps, T, D: ConstInt, BD: ConstInt, BT: ConstInt, HAS_DY2: ConstInt): - i_t = ct.bid(0) - b_y = ct.astype(ct.load(y, index=(i_t, 0), shape=(BT, BD), padding_mode=ct.PaddingMode.ZERO), ct.float32) - b_rstd = ct.astype(ct.load(rstd, index=(i_t,), shape=(BT,), padding_mode=ct.PaddingMode.ZERO), ct.float32) - b_dy = ct.astype(ct.load(dy, index=(i_t, 0), shape=(BT, BD), padding_mode=ct.PaddingMode.ZERO), ct.float32) - if HAS_DY2: - b_dy2 = ct.astype(ct.load(dy2, index=(i_t, 0), shape=(BT, BD), padding_mode=ct.PaddingMode.ZERO), ct.float32) - b_dy = ct.astype(ct.astype(b_dy + b_dy2, dy.dtype), ct.float32) - b_dot = ct.sum(b_dy * b_y, axis=1) - b_dx = b_dy * b_rstd[:, None] - b_dot[:, None] * b_y * b_rstd[:, None] - ct.store(dx, index=(i_t, 0), tile=ct.astype(b_dx, dx.dtype)) - - -# fused_beta_sigmoid - - -@ct.kernel -def fused_beta_sigmoid_fwd_kernel(x, y, scale, n_elements, BLOCK_SIZE: ConstInt): - pid = ct.bid(0) - offs = pid * BLOCK_SIZE + ct.arange(BLOCK_SIZE, dtype=ct.int32) - mask = offs < n_elements - b_x = ct.astype(ct.gather(x, offs, mask=mask, check_bounds=False, padding_value=0.0), ct.float32) - b_y = scale * (1.0 / (1.0 + ct.exp(-b_x))) - ct.scatter(y, offs, ct.astype(b_y, y.dtype), mask=mask, check_bounds=False) - - -@ct.kernel -def fused_beta_sigmoid_bwd_kernel(x, dy, dx, scale, n_elements, BLOCK_SIZE: ConstInt): - pid = ct.bid(0) - offs = pid * BLOCK_SIZE + ct.arange(BLOCK_SIZE, dtype=ct.int32) - mask = offs < n_elements - b_x = ct.astype(ct.gather(x, offs, mask=mask, check_bounds=False, padding_value=0.0), ct.float32) - b_dy = ct.astype(ct.gather(dy, offs, mask=mask, check_bounds=False, padding_value=0.0), ct.float32) - b_y = 1.0 / (1.0 + ct.exp(-b_x)) - b_dx = b_dy * scale * b_y * (1.0 - b_y) - ct.scatter(dx, offs, ct.astype(b_dx, dx.dtype), mask=mask, check_bounds=False) - - -# scalar chunk_local_cumsum (used for the gate g) +# --- Kernels: normalization and gates ------------------------------------------------------------- @ct.kernel @@ -306,9 +166,6 @@ def chunk_local_cumsum_scalar_kernel( ct.scatter(o, (t_idx, i_h), ct.astype(b_o, o.dtype), mask=mask, check_bounds=False) -# gdn_gate_chunk_cumsum / gdn_gate_bwd - - @ct.kernel def gdn_gate_chunk_cumsum_scalar_kernel( g, @@ -385,7 +242,9 @@ def gdn_gate_bwd_kernel(g, A_log, dt_bias, dyg, dg, dA, db, T, H: ConstInt, BT: ct.scatter(db, i_t * H + i_h, b_db) -# chunk_gated_delta_rule_fwd_kkt_solve_kernel (fused KK^T + solve, 4 sub-chunks) +# --- Kernels: WY representation ------------------------------------------------------------------- + + @ct.kernel def chunk_gated_delta_rule_fwd_kkt_solve_kernel( k, @@ -646,7 +505,6 @@ def chunk_gated_delta_rule_fwd_kkt_solve_kernel( b_Ai33 = b_Ai33 + ct.astype(m_I, ct.float32) # -- Step 4: block merge -- - # TODO(cutile): input_precision=ieee not representable; fp32 kept in fp32. b_Ai10 = -ct.matmul(ct.matmul(b_Ai11, b_A10), b_Ai00) b_Ai21 = -ct.matmul(ct.matmul(b_Ai22, b_A21), b_Ai11) b_Ai32 = -ct.matmul(ct.matmul(b_Ai33, b_A32), b_Ai22) @@ -1023,14 +881,7 @@ def prepare_wy_repr_bwd_kernel( scatter_flat(dg_flat, row_hv, ct.astype(b_dg, dg_flat.dtype), mask=m_t) -# chunk_delta_h (fwd_h state recurrence + bwd dhu) -# The recurrent hidden state (K x V, split into 64-wide K blocks b_h1..b_h4) is -# carried across chunks in a Python for-loop. - - -def ct_min(a, b): - # scalar/tile min for runtime ints (builtin `min` is whitelisted; `hasattr` is not). - return min(a, b) +# --- Kernels: state scan -------------------------------------------------------------------------- @ct.kernel @@ -1946,6 +1797,9 @@ def chunk_gated_delta_rule_bwd_kernel_dhu_blockdim64( ) +# --- Kernels: attention and gradients ------------------------------------------------------------- + + @ct.kernel def chunk_fwd_kernel_o( q, @@ -2296,120 +2150,9 @@ def chunk_bwd_kernel_dv_local( scatter_flat(dvf, dv_off, ct.astype(b_dv, dv.dtype), mask=do_mask) -# Host wrappers (compute grid + fixed tile sizes, then ct.launch) -# cuTile never accepts None; optional tensors are passed as an inert device stub + -# an integer flag (USE_*/HAS_*). - - -# explicit workspace: all device memory is the caller's — the engine plan (or a -# standalone harness) pre-carves every pipeline intermediate as a named view in -# ``bufs`` and passes outputs via ``out=``; nothing here allocates. - - -# l2norm -def l2norm_fwd(x, eps: float = 1e-6, out=None, rstd_out=None, stream=None): - stream = 0 if stream is None else stream - x_shape_og = x.shape - x = reshaped(x, (-1, x.shape[-1])) - y = reshaped(out, tuple(x.shape)) - T, D = x.shape[0], x.shape[-1] - MAX_FUSED_SIZE = 65536 // x.element_size() - BD = min(MAX_FUSED_SIZE, next_power_of_2(D)) - rstd = reshaped(rstd_out, (T,)) - if D <= 512: - BT = 32 - grid = (cdiv(T, BT),) - tuned_launch( - l2norm_fwd_kernel, - stream, - grid, - (x, y, rstd, float(eps), T, D, BD, BT), - cache_key=(D, BD, BT, str(x.dtype)), - ) - else: - ct.launch(stream, (T,), l2norm_fwd_kernel1, (x, y, rstd, float(eps), D, BD)) - return y.view(x_shape_og), rstd.view(x_shape_og[:-1]) - - -def l2norm_bwd( - y, - rstd, - dy, - eps: float = 1e-6, - dy2=None, - out=None, - bufs=None, - stream=None, -): - stream = 0 if stream is None else stream - y_shape_og = y.shape - y = y.reshape(-1, dy.shape[-1]) - dy = dy.reshape(-1, dy.shape[-1]) - dy2_arg = dy2.reshape(-1, dy.shape[-1]) if dy2 is not None else dummy(dtname(dy), bufs) - dx = reshaped(out, tuple(y.shape)) - T, D = y.shape[0], y.shape[-1] - MAX_FUSED_SIZE = 65536 // y.element_size() - BD = min(MAX_FUSED_SIZE, next_power_of_2(D)) - rstd_flat = rstd.reshape(-1) - if D <= 512: - BT = 32 - grid = (cdiv(T, BT),) - tuned_launch( - l2norm_bwd_kernel, - stream, - grid, - (y, rstd_flat, dy, dy2_arg, dx, float(eps), T, D, BD, BT, int(dy2 is not None)), - cache_key=(D, BD, BT, str(y.dtype), int(dy2 is not None)), - ) - else: - ct.launch( - stream, - (T,), - l2norm_bwd_kernel1, - (y, rstd_flat, dy, dy2_arg, dx, float(eps), D, BD, int(dy2 is not None)), - ) - return dx.view(y_shape_og) - - -# fused_beta_sigmoid -BETA_SIGMOID_BLOCK_SIZE = 2048 - - -def fused_beta_sigmoid_fwd(x, scale: float = 1.0, out=None, stream=None): - stream = 0 if stream is None else stream - y = reshaped(out, tuple(x.shape)) - n = x.numel() - grid = (cdiv(n, BETA_SIGMOID_BLOCK_SIZE),) - ct.launch( - stream, - grid, - fused_beta_sigmoid_fwd_kernel, - (reshaped(x, (-1,)), y.reshape(-1), float(scale), n, BETA_SIGMOID_BLOCK_SIZE), - ) - return y - - -def fused_beta_sigmoid_bwd(x, dy, scale: float = 1.0, out=None, stream=None): - stream = 0 if stream is None else stream - dx = reshaped(out, tuple(x.shape)) - n = x.numel() - grid = (cdiv(n, BETA_SIGMOID_BLOCK_SIZE),) - ct.launch( - stream, - grid, - fused_beta_sigmoid_bwd_kernel, - (reshaped(x, (-1,)), reshaped(dy, (-1,)), dx.reshape(-1), float(scale), n, BETA_SIGMOID_BLOCK_SIZE), - ) - return dx - - -def fused_beta_sigmoid(x, scale: float = 1.0, out=None, stream=None): - """Fused ``scale * sigmoid(x)`` (fp32, written to ``out``).""" - stream = 0 if stream is None else stream - return fused_beta_sigmoid_fwd(x, scale, out=out, stream=stream) +# --- Launchers: normalization and gates ----------------------------------------------------------- -# chunk_local_cumsum (scalar gate) / gdn_gate_chunk_cumsum / gdn_gate_bwd def chunk_local_cumsum_scalar( g, chunk_size, @@ -2424,7 +2167,7 @@ def chunk_local_cumsum_scalar( T, H = g.shape BT = chunk_size NT = len(chunk_indices) - g_out = reshaped(out, (T, H)) + g_out = out.reshape((T, H)) scale_val = float(scale) if scale is not None else 0.0 has_scale = int(scale is not None) cu_arg = cu_seqlens @@ -2487,8 +2230,8 @@ def gdn_gate_chunk_cumsum( T, H = g.shape BT = chunk_size NT = len(chunk_indices) - o = reshaped(out, (T, H)) - dt_arg = reshaped(opt(dt_bias, bufs, dtname(A_log)), (-1,)) + o = out.reshape((T, H)) + dt_arg = opt(dt_bias, bufs, dtname(A_log)).reshape((-1,)) scale_val = float(scale) if scale is not None else 0.0 cu_arg = cu_seqlens ci_arg = chunk_indices @@ -2498,7 +2241,7 @@ def gdn_gate_chunk_cumsum( gdn_gate_chunk_cumsum_scalar_kernel, ( g, - reshaped(A_log, (-1,)), + A_log.reshape((-1,)), dt_arg, o, scale_val, @@ -2523,20 +2266,20 @@ def gdn_gate_bwd(g, A_log, dt_bias, dyg, dg_out=None, dA_out=None, dbias_out=Non T = g.numel() // H BT = 32 NT = cdiv(T, BT) - dg = reshaped(dg_out, tuple(g.shape)) - dA_nt = reshaped(bufs["dA_gate"], (NT, H)) - db_nt = reshaped(bufs["db_gate"], (NT, H)) if dt_bias is not None else None - dt_arg = reshaped(opt(dt_bias, bufs, dtname(A_log)), (-1,)) - db_arg = db_nt.reshape(-1) if db_nt is not None else dummy("float32", bufs) + dg = dg_out.reshape(tuple(g.shape)) + dA_nt = bufs["dA_gate"].reshape((NT, H)) + db_nt = bufs["db_gate"].reshape((NT, H)) if dt_bias is not None else None + dt_arg = opt(dt_bias, bufs, dtname(A_log)).reshape((-1,)) + db_arg = opt(db_nt, bufs).reshape(-1) ct.launch( stream, (NT, H), gdn_gate_bwd_kernel, ( - reshaped(g, (-1,)), - reshaped(A_log, (-1,)), + g.reshape((-1,)), + A_log.reshape((-1,)), dt_arg, - reshaped(dyg, (-1,)), + dyg.reshape((-1,)), dg.reshape(-1), dA_nt.reshape(-1), db_arg, @@ -2546,13 +2289,15 @@ def gdn_gate_bwd(g, A_log, dt_bias, dyg, dg_out=None, dA_out=None, dbias_out=Non int(dt_bias is not None), ), ) - sum_leading(reshaped(dA_out, (H,)), dA_nt, NT, H, stream=stream) + sum_leading(dA_out.reshape((H,)), dA_nt, NT, H, stream=stream) if dt_bias is not None: - sum_leading(reshaped(dbias_out, (H,)), db_nt, NT, H, stream=stream) + sum_leading(dbias_out.reshape((H,)), db_nt, NT, H, stream=stream) return dg, dA_out, (dbias_out if dt_bias is not None else None) -# recompute_w_u_fwd / prepare_wy_repr_bwd +# --- Launchers: WY representation ----------------------------------------------------------------- + + def recompute_w_u_fwd(k, v, beta, A, g=None, cu_seqlens=None, chunk_indices=None, bufs=None, stream=None): stream = 0 if stream is None else stream T, H, K, V, HV = *k.shape, v.shape[-1], v.shape[1] @@ -2560,16 +2305,16 @@ def recompute_w_u_fwd(k, v, beta, A, g=None, cu_seqlens=None, chunk_indices=None BK = 64 BV = 64 NT = len(chunk_indices) - w = reshaped(bufs["w"], (T, HV, K)) - u = reshaped(bufs["u"], (T, HV, V)) - beta2 = reshaped(beta, (T, HV)) + w = bufs["w"].reshape((T, HV, K)) + u = bufs["u"].reshape((T, HV, V)) + beta2 = beta.reshape((T, HV)) A3 = A.reshape(T, HV, BT) g_arg = g.reshape(T, HV) if g is not None else dummy("float32", bufs) cu_arg = cu_seqlens ci_arg = chunk_indices - tuned_launch( + autotuned_launch( recompute_w_u_fwd_kernel, - stream, + (H, HV, K, V, BT, BK, BV, int(g is not None), str(k.dtype), current_device_id()), (NT, HV), ( k, @@ -2590,7 +2335,8 @@ def recompute_w_u_fwd(k, v, beta, A, g=None, cu_seqlens=None, chunk_indices=None BV, int(g is not None), ), - cache_key=(H, HV, K, V, BT, BK, BV, int(g is not None), str(k.dtype), dev_id(k)), + occ_choices=TUNE_OCC, + stream=stream, ) return w, u @@ -2603,17 +2349,17 @@ def prepare_wy_repr_bwd(k, v, beta, A, dw, du, g=None, cu_seqlens=None, chunk_in CONST_TILING = 64 BK = min(max(next_power_of_2(K), 16), CONST_TILING) BV = min(max(next_power_of_2(V), 16), CONST_TILING) - dk = reshaped(bufs["wy_dk"], (T, HV, K)) - dv = reshaped(bufs["wy_dv"], (T, HV, V)) - dg = reshaped(bufs["wy_dg"], (T, HV)) if g is not None else None - db = reshaped(bufs["db"], (T, HV)) + dk = bufs["wy_dk"].reshape((T, HV, K)) + dv = bufs["wy_dv"].reshape((T, HV, V)) + dg = bufs["wy_dg"].reshape((T, HV)) if g is not None else None + db = bufs["db"].reshape((T, HV)) g_arg = opt(g, bufs) - dg_arg = dg if dg is not None else dummy("float32", bufs) + dg_arg = opt(dg, bufs) cu_arg = cu_seqlens ci_arg = chunk_indices - tuned_launch( + autotuned_launch( prepare_wy_repr_bwd_kernel, - stream, + (H, HV, K, V, BT, BK, BV, int(g is not None), str(k.dtype), current_device_id()), (NT, HV), ( k, @@ -2640,16 +2386,61 @@ def prepare_wy_repr_bwd(k, v, beta, A, dw, du, g=None, cu_seqlens=None, chunk_in cdiv(K, BK), cdiv(V, BV), ), - cache_key=(H, HV, K, V, BT, BK, BV, int(g is not None), str(k.dtype), dev_id(k)), + occ_choices=TUNE_OCC, + stream=stream, ) if H != HV: - dk_r = reshaped(bufs["wy_dk_hred"], (T, H, K)) + dk_r = bufs["wy_dk_hred"].reshape((T, H, K)) head_group_sum(dk_r, dk, T, H, HV // H, K, stream=stream) dk = dk_r return dk, dv, db, dg -# chunk_gated_delta_rule_fwd_h / bwd_dhu +def chunk_gated_delta_rule_fwd_intra(k, v, g=None, beta=None, cu_seqlens=None, chunk_size=64, chunk_indices=None, bufs=None, compute_wu=True, stream=None): + stream = 0 if stream is None else stream + T, H, K, HV = *k.shape, beta.shape[1] + BT = chunk_size + NT = len(chunk_indices) + + A = bufs["A"].reshape((T, HV, BT)) + zero_fill(A, stream=stream) + BK = 64 + g_arg = opt(g, bufs) + cu_arg = cu_seqlens + ci_arg = chunk_indices + # Masked BC=16 4-sub-block kernel: masks the partial last chunk's tail. + BC = 16 + autotuned_launch( + chunk_gated_delta_rule_fwd_kkt_solve_kernel, + (H, HV, K, BT, BC, BK, int(g is not None), str(k.dtype), current_device_id()), + (NT, HV), + ( + k, + g_arg, + beta, + A, + cu_arg, + ci_arg, + H, + HV, + K, + BT, + BC, + BK, + int(g is not None), + ), + occ_choices=TUNE_OCC, + stream=stream, + ) + if not compute_wu: + return None, None, A + w, u = recompute_w_u_fwd(k=k, v=v, beta=beta, A=A, g=g, cu_seqlens=cu_seqlens, chunk_indices=chunk_indices, bufs=bufs, stream=stream) + return w, u, A + + +# --- Launchers: state scan ------------------------------------------------------------------------ + + def chunk_gated_delta_rule_fwd_h( k, w, @@ -2674,17 +2465,17 @@ def chunk_gated_delta_rule_fwd_h( chunk_offsets = bufs["chunk_offsets"] state_shape = (N, HV, V, K) if state_v_first else (N, HV, K, V) - h = reshaped(bufs["state_checkpoints"], (NT, HV) + state_shape[2:]) - final_state = reshaped(bufs["final_state"], state_shape) if output_final_state else None + h = bufs["state_checkpoints"].reshape((NT, HV) + state_shape[2:]) + final_state = bufs["final_state"].reshape(state_shape) if output_final_state else None if final_state is not None: zero_fill(final_state, stream=stream) - v_new = reshaped(bufs["v_new"], (T, HV, V)) if save_new_value else None + v_new = bufs["v_new"].reshape((T, HV, V)) if save_new_value else None - vnew_arg = v_new if v_new is not None else dummy(dtname(u), bufs) + vnew_arg = opt(v_new, bufs, dtname(u)) g_arg = opt(g, bufs) gk_arg = opt(gk, bufs) - h0_arg = initial_state if initial_state is not None else dummy("float32", bufs) - ht_arg = final_state if final_state is not None else dummy("float32", bufs) + h0_arg = opt(initial_state, bufs) + ht_arg = opt(final_state, bufs) cu_arg = cu_seqlens co_arg = chunk_offsets @@ -2707,22 +2498,9 @@ def chunk_gated_delta_rule_fwd_h( grid = (cdiv(V, BV), N * HV) # Multi-dim arrays passed as-is (per-dim indices + real strides). - tuned_launch( + autotuned_launch( chunk_gated_delta_rule_fwd_kernel_h_blockdim64, - stream, - grid, ( - k, - u, - w, - vnew_arg, - g_arg, - gk_arg, - h, - h0_arg, - ht_arg, - cu_arg, - co_arg, H, HV, K, @@ -2735,8 +2513,22 @@ def chunk_gated_delta_rule_fwd_h( int(output_final_state), int(save_new_value), int(state_v_first), + str(k.dtype), + current_device_id(), ), - cache_key=( + grid, + ( + k, + u, + w, + vnew_arg, + g_arg, + gk_arg, + h, + h0_arg, + ht_arg, + cu_arg, + co_arg, H, HV, K, @@ -2749,9 +2541,9 @@ def chunk_gated_delta_rule_fwd_h( int(output_final_state), int(save_new_value), int(state_v_first), - str(k.dtype), - dev_id(k), ), + occ_choices=TUNE_OCC, + stream=stream, ) return h, v_new, final_state @@ -2780,9 +2572,9 @@ def chunk_gated_delta_rule_bwd_dhu( N, NT = len(cu_seqlens) - 1, len(chunk_indices) chunk_offsets = bufs["chunk_offsets"] - dh = reshaped(bufs["dstate"], (NT, HV, V, K) if state_v_first else (NT, HV, K, V)) - dh0 = reshaped(bufs["dstate0"], tuple(h0.shape)) if h0 is not None else None - dv2 = reshaped(bufs["dv2"], (T, HV, V)) + dh = bufs["dstate"].reshape((NT, HV, V, K) if state_v_first else (NT, HV, K, V)) + dh0 = bufs["dstate0"].reshape(tuple(h0.shape)) if h0 is not None else None + dv2 = bufs["dv2"].reshape((T, HV, V)) # For K>128 the blockdim64 kernel's 4 (64xBV) fp32 dH accumulators overflow # tileiras allocation at BV=64; shrink BV to keep the live footprint in range. @@ -2791,13 +2583,25 @@ def chunk_gated_delta_rule_bwd_dhu( g_arg = opt(g, bufs) gk_arg = opt(gk, bufs) dht_arg = opt(dstate_in, bufs) - dh0_arg = dh0 if dh0 is not None else dummy("float32", bufs) + dh0_arg = opt(dh0, bufs) cu_arg = cu_seqlens co_arg = chunk_offsets - tuned_launch( + autotuned_launch( chunk_gated_delta_rule_bwd_kernel_dhu_blockdim64, - stream, + ( + H, + HV, + K, + V, + BT, + BV, + int(g is not None), + int(gk is not None), + int(state_v_first), + str(q.dtype), + current_device_id(), + ), grid, ( q, @@ -2827,24 +2631,15 @@ def chunk_gated_delta_rule_bwd_dhu( int(state_v_first), 1, ), - cache_key=( - H, - HV, - K, - V, - BT, - BV, - int(g is not None), - int(gk is not None), - int(state_v_first), - str(q.dtype), - dev_id(q), - ), + occ_choices=TUNE_OCC, + stream=stream, ) return dh, dh0, dv2 -# chunk_fwd_o / chunk_bwd_dqkwg / chunk_bwd_dv_local +# --- Launchers: attention and gradients ----------------------------------------------------------- + + def chunk_fwd_o( q, k, @@ -2866,7 +2661,7 @@ def chunk_fwd_o( NT = len(chunk_indices) if scale is None: scale = k.shape[-1] ** -0.5 - o = reshaped(bufs["o"], (T, HV, V)) + o = bufs["o"].reshape((T, HV, V)) BK = min(max(next_power_of_2(K), 16), 64) BV = 64 grid = (cdiv(V, BV), NT, HV) @@ -2879,21 +2674,9 @@ def chunk_fwd_o( h3 = h.reshape(NT * HV, V, K) else: h3 = h.reshape(NT * HV, K, V) - tuned_launch( + autotuned_launch( chunk_fwd_kernel_o, - stream, - grid, ( - q, - k, - v, - h3, - g_arg, - gg_arg, - o, - cu_arg, - ci_arg, - float(scale), H, HV, K, @@ -2904,8 +2687,21 @@ def chunk_fwd_o( int(g is not None), int(g_gamma is not None), int(state_v_first), + str(q.dtype), + current_device_id(), ), - cache_key=( + grid, + ( + q, + k, + v, + h3, + g_arg, + gg_arg, + o, + cu_arg, + ci_arg, + float(scale), H, HV, K, @@ -2916,9 +2712,9 @@ def chunk_fwd_o( int(g is not None), int(g_gamma is not None), int(state_v_first), - str(q.dtype), - dev_id(q), ), + occ_choices=TUNE_OCC, + stream=stream, ) return o @@ -2950,10 +2746,10 @@ def chunk_bwd_dqkwg( BK = min(max(next_power_of_2(K), 16), CONST_TILING) BV = min(max(next_power_of_2(V), 16), CONST_TILING) NK = cdiv(K, BK) - dq = reshaped(bufs["dq"], (T, HV, K)) - dk = reshaped(bufs["dk"], (T, HV, K)) - dg = reshaped(bufs["dg_nk"], (NK, T, HV)) if g is not None else None - dw = reshaped(bufs["dw"], (T, HV, K)) if w is not None else None + dq = bufs["dq"].reshape((T, HV, K)) + dk = bufs["dk"].reshape((T, HV, K)) + dg = bufs["dg_nk"].reshape((NK, T, HV)) if g is not None else None + dw = bufs["dw"].reshape((T, HV, K)) if w is not None else None grid = (NK, NT, HV) # h/dh flattened to (NT*HV,*,*) slabs for block-indexed TMA. if state_v_first: @@ -2964,15 +2760,29 @@ def chunk_bwd_dqkwg( dh3 = dh.reshape(NT * HV, K, V) g_arg = g.reshape(T, HV) if g is not None else dummy(dtname(q), bufs) gg_arg = opt(g_gamma, bufs) - dw3 = dw if dw is not None else dummy(dtname(k), bufs) + dw3 = opt(dw, bufs, dtname(k)) dv3 = dv.reshape(T, HV, V) if dv is not None else dummy(dtname(k), bufs) - dg3 = dg if dg is not None else dummy("float32", bufs) + dg3 = opt(dg, bufs) cu_arg = cu_seqlens ci_arg = chunk_indices use_dw = int(dw is not None and dv is not None) - tuned_launch( + autotuned_launch( chunk_bwd_kernel_dqkwg, - stream, + ( + H, + HV, + K, + V, + BT, + BK, + BV, + int(g is not None), + int(g_gamma is not None), + use_dw, + int(state_v_first), + str(q.dtype), + current_device_id(), + ), grid, ( q, @@ -3004,31 +2814,18 @@ def chunk_bwd_dqkwg( int(state_v_first), cdiv(V, BV), ), - cache_key=( - H, - HV, - K, - V, - BT, - BK, - BV, - int(g is not None), - int(g_gamma is not None), - use_dw, - int(state_v_first), - str(q.dtype), - dev_id(q), - ), + occ_choices=TUNE_OCC, + stream=stream, ) if H != HV: - dq_r = reshaped(bufs["dq_hred"], (T, H, K)) - dk_r = reshaped(bufs["dk_hred"], (T, H, K)) + dq_r = bufs["dq_hred"].reshape((T, H, K)) + dk_r = bufs["dk_hred"].reshape((T, H, K)) head_group_sum(dq_r, dq, T, H, HV // H, K, stream=stream) head_group_sum(dk_r, dk, T, H, HV // H, K, stream=stream) dq, dk = dq_r, dk_r if dg is not None: - dg_r = reshaped(bufs["dg"], (T, HV)) - sum_leading(reshaped(dg_r, (T * HV,)), reshaped(dg, (NK, T * HV)), NK, T * HV, stream=stream) + dg_r = bufs["dg"].reshape((T, HV)) + sum_leading(dg_r.reshape((T * HV,)), dg.reshape((NK, T * HV)), NK, T * HV, stream=stream) dg = dg_r return dq, dk, dw, dg @@ -3041,7 +2838,7 @@ def chunk_bwd_dv_local(q, k, do, g=None, g_gamma=None, A=None, scale=None, cu_se BK = min(max(next_power_of_2(K), 16), CONST_TILING) BV = min(max(next_power_of_2(V), 16), CONST_TILING) NT = len(chunk_indices) - dv = reshaped(bufs["dv"], (T, HV, V)) + dv = bufs["dv"].reshape((T, HV, V)) grid = (NT, HV) g_arg = opt(g, bufs) gg_arg = opt(g_gamma, bufs) @@ -3049,33 +2846,9 @@ def chunk_bwd_dv_local(q, k, do, g=None, g_gamma=None, A=None, scale=None, cu_se cu_arg = cu_seqlens ci_arg = chunk_indices scale_val = float(scale) if scale is not None else 0.0 - tuned_launch( + autotuned_launch( chunk_bwd_kernel_dv_local, - stream, - grid, ( - q, - k, - g_arg, - gg_arg, - A_arg, - do, - dv, - cu_arg, - ci_arg, - scale_val, - H, - HV, - K, - V, - BT, - BK, - BV, - int(g is not None), - int(g_gamma is not None), - int(A is not None), - ), - cache_key=( H, HV, K, @@ -3087,55 +2860,38 @@ def chunk_bwd_dv_local(q, k, do, g=None, g_gamma=None, A=None, scale=None, cu_se int(g_gamma is not None), int(A is not None), str(q.dtype), - dev_id(q), + current_device_id(), ), - ) - return dv - - -# fused intra (KK^T + solve_tril + recompute_w_u) -def chunk_gated_delta_rule_fwd_intra(k, v, g=None, beta=None, cu_seqlens=None, chunk_size=64, chunk_indices=None, bufs=None, compute_wu=True, stream=None): - stream = 0 if stream is None else stream - T, H, K, HV = *k.shape, beta.shape[1] - BT = chunk_size - NT = len(chunk_indices) - - A = reshaped(bufs["A"], (T, HV, BT)) - zero_fill(A, stream=stream) - BK = 64 - g_arg = opt(g, bufs) - cu_arg = cu_seqlens - ci_arg = chunk_indices - # Masked BC=16 4-sub-block kernel: masks the partial last chunk's tail. - BC = 16 - tuned_launch( - chunk_gated_delta_rule_fwd_kkt_solve_kernel, - stream, - (NT, HV), + grid, ( + q, k, g_arg, - beta, - A, + gg_arg, + A_arg, + do, + dv, cu_arg, ci_arg, + scale_val, H, HV, K, + V, BT, - BC, BK, + BV, int(g is not None), + int(g_gamma is not None), + int(A is not None), ), - cache_key=(H, HV, K, BT, BC, BK, int(g is not None), str(k.dtype), dev_id(k)), + occ_choices=TUNE_OCC, + stream=stream, ) - if not compute_wu: - return None, None, A - w, u = recompute_w_u_fwd(k=k, v=v, beta=beta, A=A, g=g, cu_seqlens=cu_seqlens, chunk_indices=chunk_indices, bufs=bufs, stream=stream) - return w, u, A + return dv -# forward / backward drivers + entry point +# --- Pipelines ------------------------------------------------------------------------------------ def chunk_gated_delta_rule_fwd( @@ -3271,7 +3027,7 @@ def chunk_gated_delta_rule_bwd( m_dg = 1 for s_ in dg.shape: m_dg *= int(s_) - add_inplace(reshaped(dg, (m_dg,)), reshaped(dg2, (m_dg,)), m_dg, stream=stream) + add_inplace(dg.reshape((m_dg,)), dg2.reshape((m_dg,)), m_dg, stream=stream) dg = chunk_local_cumsum(dg, chunk_size=64, reverse=True, cu_seqlens=cu_seqlens, chunk_indices=chunk_indices, out=bufs["dg_cum"], stream=stream) dA_log, ddt_bias = None, None if use_gate_in_kernel: @@ -3299,63 +3055,33 @@ def chunk_gated_delta_rule_grad( dstate_in=None, scale=None, initial_state=None, + use_qk_l2norm_in_kernel=False, state_v_first=False, cu_seqlens=None, - cu_seqlens_cpu=None, chunk_indices=None, - use_qk_l2norm_in_kernel=False, - use_gate_in_kernel=False, - A_log=None, - dt_bias=None, - use_beta_sigmoid_in_kernel=False, - allow_neg_eigval=False, bufs=None, stream=None, ): - r"""GDN backward as a plain pipeline over explicit THD arguments. + r"""GDN backward over THD (token-packed) inputs. ``q``/``k``/``v``/``do`` are ``[total_T, H, D]``, ``g``/``beta`` are - ``[total_T, H]``; ``cu_seqlens`` and ``chunk_indices`` are required. - Recomputes the forward's prep (L2-normalized Q/K + rstd, cumulative gate, - intra-chunk WY matrix) from the ORIGINAL inputs, then runs the backward - kernels. - - Returns ``(dq, dk, dv, dg, dbeta, dh0, dA_log, ddt_bias)`` in THD layout. - """ + ``[total_T, HV]``; ``cu_seqlens``, ``chunk_indices`` and the pre-carved + ``bufs`` views are required. Recomputes the forward's prep (L2-normalized + q/k + rstd, cumulative gate, intra-chunk WY matrix) from the inputs, then + runs the backward kernels. Gradients land in the caller's planted carves; + the returned handles are those same buffers.""" stream = 0 if stream is None else stream - if cu_seqlens is None: - raise ValueError("cu_seqlens is required (THD layout)") if scale is None: scale = k.shape[-1] ** -0.5 - q_rstd, k_rstd = None, None - q_in, k_in = q, k + q_in, k_in, q_rstd, k_rstd = q, k, None, None if use_qk_l2norm_in_kernel: q_in, q_rstd = l2norm_fwd(q, out=bufs["q_norm"], rstd_out=bufs["q_rstd"], stream=stream) k_in, k_rstd = l2norm_fwd(k, out=bufs["k_norm"], rstd_out=bufs["k_rstd"], stream=stream) - beta_raw = beta - if use_beta_sigmoid_in_kernel: - beta = fused_beta_sigmoid(beta_raw, scale=2.0 if allow_neg_eigval else 1.0, out=bufs["beta_sig"], stream=stream) - if cu_seqlens is not None and chunk_indices is None: - raise ValueError("varlen (cu_seqlens) requires chunk_indices — callers build the (seq, intra) table") - g_cum, _o, A, _fs, initial_state, g_input = chunk_gated_delta_rule_fwd( - q=q_in, - k=k_in, - v=v, - g=g, - beta=beta, - scale=scale, - initial_state=initial_state, - output_final_state=False, - cu_seqlens=cu_seqlens, - chunk_indices=chunk_indices, - state_v_first=state_v_first, - use_gate_in_kernel=use_gate_in_kernel, - A_log=A_log, - dt_bias=dt_bias, - bufs=bufs, - stream=stream, + g_cum = chunk_local_cumsum(g, chunk_size=BT_CHUNK, scale=RCP_LN2, cu_seqlens=cu_seqlens, chunk_indices=chunk_indices, out=bufs["g_cum"], stream=stream) + _w, _u, A = chunk_gated_delta_rule_fwd_intra( + k=k_in, v=v, g=g_cum, beta=beta, cu_seqlens=cu_seqlens, chunk_indices=chunk_indices, bufs=bufs, compute_wu=False, stream=stream ) - dq, dk, dk2, dv, db, dg, dh0, dA_log, ddt_bias = chunk_gated_delta_rule_bwd( + dq, dk, dk2, dv, db, dg, dh0, _dA_log, _ddt_bias = chunk_gated_delta_rule_bwd( q=q_in, k=k_in, v=v, @@ -3366,27 +3092,20 @@ def chunk_gated_delta_rule_grad( initial_state=initial_state, do=do, dstate_in=dstate_in, + state_v_first=state_v_first, cu_seqlens=cu_seqlens, chunk_indices=chunk_indices, - state_v_first=state_v_first, - use_gate_in_kernel=use_gate_in_kernel, - g_input=g_input, - A_log=A_log, - dt_bias=dt_bias, bufs=bufs, stream=stream, ) if use_qk_l2norm_in_kernel: - dq = l2norm_bwd(q_in, q_rstd, dq, out=bufs["dq_final"], bufs=bufs, stream=stream) - dk = l2norm_bwd(k_in, k_rstd, dk, dy2=dk2, out=bufs["dk_final"], bufs=bufs, stream=stream) + dq = l2norm_bwd(q_in, q_rstd, dq, out=bufs["dq_l2"], bufs=bufs, stream=stream) + dk = l2norm_bwd(k_in, k_rstd, dk, dy2=dk2, out=bufs["dk_l2"], bufs=bufs, stream=stream) else: - n_dk = 1 - for s_ in dk.shape: - n_dk *= int(s_) - add_inplace(reshaped(dk, (n_dk,)), reshaped(dk2, (n_dk,)), n_dk, stream=stream) - if use_beta_sigmoid_in_kernel: - db = fused_beta_sigmoid_bwd(beta_raw, db, scale=2.0 if allow_neg_eigval else 1.0, out=db, stream=stream) - return dq, dk, dv, dg, db, dh0, dA_log, ddt_bias + # dk/dk2 are the head-reduced finals for GVA, HV-head for MHA + n_dk = k.shape[0] * k.shape[1] * k.shape[2] + add_inplace(dk.reshape(-1), dk2.reshape(-1), n_dk, stream=stream) + return dq, dk, dv, db, dg, dh0 def chunk_gated_delta_rule( diff --git a/python/cudnn/linear_attention/cutile/kernels/kda_chunk_cutile.py b/python/cudnn/linear_attention/cutile/kernels/kda.py similarity index 80% rename from python/cudnn/linear_attention/cutile/kernels/kda_chunk_cutile.py rename to python/cudnn/linear_attention/cutile/kernels/kda.py index bd0189fc0..0e01ed969 100644 --- a/python/cudnn/linear_attention/cutile/kernels/kda_chunk_cutile.py +++ b/python/cudnn/linear_attention/cutile/kernels/kda.py @@ -17,14 +17,32 @@ import logging -import os from types import SimpleNamespace import cuda.tile as ct from cuda.tile.tune import exhaustive_search -from .common import dev_id, dummy, head_group_sum, opt, reshaped, sum_leading, zero_fill -from cudnn.frost.buffers import dtype_name as dtname +from .common import ( + autotuned_launch, + cdiv, + ct_min, + dummy, + exp, + exp2, + fused_beta_sigmoid, + fused_beta_sigmoid_bwd, + head_group_sum, + l2norm_bwd, + l2norm_fwd, + launch_hint_cache, + next_power_of_2, + opt, + softplus, + sum_leading, + tf32, + zero_fill, +) +from cudnn.frost.buffers import current_device_id, dtype_name as dtname logger = logging.getLogger(__name__) @@ -36,68 +54,24 @@ BT_CHUNK = 64 -# --------------------------------------------------------------------------- -# Launch-hint autotune. Only the @ct.kernel launch hints -# (occupancy x num_worker_warps) are explored via kernel.replace_hints(...); -# the grid / args / algorithm are kept UNCHANGED. The first config in every -# grid equals the kernel default (occupancy=1, num_worker_warps=4) so any shape -# that does not improve keeps the original behaviour (no regression). -# --------------------------------------------------------------------------- -DISABLE_TUNE = os.environ.get("DISABLE_TUNE", "") not in ("", "0") -launch_hint_cache: dict = {} - - -def launch_hint_configs(occ_choices, nww_choices=(4, 8)): - """occupancy x num_worker_warps grid; deduped, default (occ=1,nww=4) first.""" - seen = set() - cfgs = [] - for occ in occ_choices: - for nww in nww_choices: - key = (occ, nww) - if key in seen: - continue - seen.add(key) - cfgs.append(SimpleNamespace(occupancy=occ, num_worker_warps=nww)) - return cfgs - - -def autotuned_launch(kernel, cache_key, grid, args, occ_choices=(1, 2, 3, 4), nww_choices=(4, 8), timeout=30, stream=None): - """Launch ``kernel`` with the best launch hints for ``cache_key``. - - Tune-once/cache/launch over launch hints only (grid, args and signature are - fixed). Falls back to the base kernel when DISABLE_TUNE is set or tuning - fails / times out. Default config (occ=1, nww=4) is explored first so a - no-improvement shape keeps the base behaviour. - """ - stream = 0 if stream is None else stream - if DISABLE_TUNE: - ct.launch(stream, grid, kernel, args) - return +# --- Host helpers --------------------------------------------------------------------------------- - if cache_key not in launch_hint_cache: - tuned = None - try: - configs = launch_hint_configs(occ_choices, nww_choices) - with ct.compiler_timeout(timeout): - result = exhaustive_search( - configs, - stream, - lambda cfg: grid, - kernel, - lambda cfg: args, - lambda cfg: {"occupancy": cfg.occupancy, "num_worker_warps": cfg.num_worker_warps}, - ) - best = result.best.config - tuned = kernel.replace_hints(occupancy=best.occupancy, num_worker_warps=best.num_worker_warps) - except Exception: - tuned = None - launch_hint_cache[cache_key] = tuned - tuned = launch_hint_cache[cache_key] - if tuned is None: - ct.launch(stream, grid, kernel, args) - else: - ct.launch(stream, grid, tuned, args) +def cast(bufs, name, src, ref, stream=None): + """Dtype cast at the gradient boundary, through the ``bufs[name]`` carve.""" + if str(src.dtype).split(".")[-1] == str(ref.dtype).split(".")[-1]: + return src + from .common import cast_copy + + dst = bufs[name] + n = 1 + for s_ in src.shape: + n *= int(s_) + cast_copy(dst.reshape((n,)), src.reshape((n,)), n, stream=0 if stream is None else stream) + return dst + + +# --- Launch tuning -------------------------------------------------------------------------------- def autotuned_launch_bv(kernel, cache_key, bv_choices, grid_fn, args_fn, timeout=40, stream=None): @@ -109,7 +83,7 @@ def autotuned_launch_bv(kernel, cache_key, bv_choices, grid_fn, args_fn, timeout replace_hints; it is swept as a real config dimension. ``grid_fn(bv)`` / ``args_fn(bv)`` rebuild the BV-dependent grid and argument tuple. The winning BV is cached and the BV-specialized kernel re-launched on every subsequent - call. ``bv_choices[0]`` is the safe fallback on DISABLE_TUNE / tuning failure. + call. ``bv_choices[0]`` is the safe fallback when tuning fails. Why the sweep is BV-ONLY: the state scan is register-bound (255 reg/thread -> ~1 block/SM) with a serial per-CTA NT inter-chunk loop, so the only @@ -120,11 +94,6 @@ def autotuned_launch_bv(kernel, cache_key, bv_choices, grid_fn, args_fn, timeout this search (the larger sweep mis-ranks BV). """ stream = 0 if stream is None else stream - if DISABLE_TUNE: - bv = bv_choices[0] - ct.launch(stream, grid_fn(bv), kernel, args_fn(bv)) - return - if cache_key not in launch_hint_cache: chosen = None try: @@ -148,72 +117,7 @@ def autotuned_launch_bv(kernel, cache_key, bv_choices, grid_fn, args_fn, timeout ct.launch(stream, grid_fn(bv), kernel, args_fn(bv)) -# =========================================================================== -# Host-side utilities -# =========================================================================== - - -def cdiv(a: int, b: int) -> int: - return (a + b - 1) // b - - -def next_power_of_2(n: int) -> int: - return 1 << (n - 1).bit_length() - - -def i32_flat(t): - """Required index buffer as a flat 1-D view. Kernels index these - via flat loads (e.g. chunk_indices[i_t*2]); chunk_indices is (NT, 2) and - its row-major flat layout is [seg0, intra0, seg1, intra1, ...].""" - from .common import reshaped - - n = 1 - for s_ in t.shape: - n *= int(s_) - return reshaped(t, (n,)) - - -# explicit workspace: all device memory is the caller's — the engine plan (or a -# standalone harness) pre-carves every pipeline intermediate as a named view in -# ``bufs`` and passes outputs via ``out=``; nothing here allocates. - - -def cast(bufs, name, src, ref, stream=None): - """Dtype cast at the gradient boundary, through the ``bufs[name]`` carve.""" - if str(src.dtype).split(".")[-1] == str(ref.dtype).split(".")[-1]: - return src - from .common import cast_copy, reshaped - - dst = bufs[name] - n = 1 - for s_ in src.shape: - n *= int(s_) - cast_copy(reshaped(dst, (n,)), reshaped(src, (n,)), n, stream=0 if stream is None else stream) - return dst - - -# =========================================================================== -# Device helpers (plain `def` — NOT @ct.kernel) -# =========================================================================== - - -def exp(x): - return ct.exp(ct.astype(x, ct.float32)) - - -def exp2(x): - return ct.exp2(ct.astype(x, ct.float32)) - - -def softplus(x): - # softplus: where(x <= 20, log1p(exp(x)), x) - return ct.where(x <= 20.0, ct.log(1.0 + ct.exp(x)), x) - - -def tf32(a): - """ct.mma/ct.matmul do not auto-cast fp32 operands to tf32; cast - explicitly (allow-tf32 matmul semantics).""" - return ct.astype(a, ct.tfloat32) if a.dtype == ct.float32 else a +# --- Device helpers ------------------------------------------------------------------------------- def mma_operands(a, b): @@ -262,107 +166,29 @@ def reg_matmul(a, b): return ct.sum(aa[:, :, None] * bb[None, :, :], axis=1) -def ct_min(a, b): - # scalar/tile min for runtime ints. Avoid `hasattr` -- cutile traces this - # as device code and `hasattr` is not a supported builtin; the builtin - # `min` IS whitelisted and handles both scalar ints and tile values. - return min(a, b) - - -# =========================================================================== -# l2norm kernels -# =========================================================================== -# GATHER RATIONALE (l2norm): each row is loaded once, reduced once, written -# once (single-pass streaming, no cooperative reuse) -> ct.gather/ct.scatter -# avoids the smem staging tax of ct.load. Column mask handles BD-power-of-2 -# padding; partial-row handling via masks. - - -@ct.kernel -def l2norm_fwd_kernel1(x, y, rstd, eps, D, BD: ConstInt): - # D > 512 path: one row per program, row length D. - i_t = ct.bid(0) - cols = ct.arange(BD, dtype=ct.int32) - mask = cols < D - - b_x = ct.astype(ct.gather(x, (i_t, cols), mask=mask, check_bounds=False, padding_value=0.0), ct.float32) - b_rstd = ct.rsqrt(ct.sum(b_x * b_x) + eps) - b_y = b_x * b_rstd - ct.scatter(y, (i_t, cols), ct.astype(b_y, y.dtype), mask=mask, check_bounds=False) - ct.scatter(rstd, (i_t,), ct.astype(b_rstd, rstd.dtype)) - - -@ct.kernel -def l2norm_bwd_kernel1(y, rstd, dy, dx, eps, D, BD: ConstInt): - i_t = ct.bid(0) - cols = ct.arange(BD, dtype=ct.int32) - mask = cols < D - - b_y = ct.astype(ct.gather(y, (i_t, cols), mask=mask, check_bounds=False, padding_value=0.0), ct.float32) - b_dy = ct.astype(ct.gather(dy, (i_t, cols), mask=mask, check_bounds=False, padding_value=0.0), ct.float32) - b_rstd = ct.astype(ct.gather(rstd, (i_t,), check_bounds=False, padding_value=0.0), ct.float32).item() - - b_dx = b_dy * b_rstd - ct.sum(b_dy * b_y) * b_y * b_rstd - ct.scatter(dx, (i_t, cols), ct.astype(b_dx, dx.dtype), mask=mask, check_bounds=False) - - -@ct.kernel -def l2norm_fwd_kernel(x, y, rstd, eps, T, D: ConstInt, BD: ConstInt, BT: ConstInt): - # D <= 512 path: BT rows per block, BD power-of-2 cols. Block-aligned -> - # ct.load with block index + ZERO padding (boundary-checked block load). - i_t = ct.bid(0) - b_x = ct.astype(ct.load(x, index=(i_t, 0), shape=(BT, BD), padding_mode=ct.PaddingMode.ZERO), ct.float32) - b_rstd = ct.rsqrt(ct.sum(b_x * b_x, axis=1) + eps) - b_y = b_x * b_rstd[:, None] - ct.store(y, index=(i_t, 0), tile=ct.astype(b_y, y.dtype)) - ct.store(rstd, index=(i_t,), tile=ct.astype(b_rstd, rstd.dtype)) - - -@ct.kernel -def l2norm_bwd_kernel(y, rstd, dy, dx, eps, T, D: ConstInt, BD: ConstInt, BT: ConstInt): - i_t = ct.bid(0) - b_y = ct.astype(ct.load(y, index=(i_t, 0), shape=(BT, BD), padding_mode=ct.PaddingMode.ZERO), ct.float32) - b_rstd = ct.astype(ct.load(rstd, index=(i_t,), shape=(BT,), padding_mode=ct.PaddingMode.ZERO), ct.float32) - b_dy = ct.astype(ct.load(dy, index=(i_t, 0), shape=(BT, BD), padding_mode=ct.PaddingMode.ZERO), ct.float32) - b_dot = ct.sum(b_dy * b_y, axis=1) - b_dx = b_dy * b_rstd[:, None] - b_dot[:, None] * b_y * b_rstd[:, None] - ct.store(dx, index=(i_t, 0), tile=ct.astype(b_dx, dx.dtype)) - - -# =========================================================================== -# fused_beta_sigmoid -# =========================================================================== -# GATHER RATIONALE: pure 1-D element-wise streaming, each element touched once. - - -@ct.kernel -def fused_beta_sigmoid_fwd_kernel(x, y, scale, n_elements, BLOCK_SIZE: ConstInt): - pid = ct.bid(0) - offs = pid * BLOCK_SIZE + ct.arange(BLOCK_SIZE, dtype=ct.int32) - mask = offs < n_elements - b_x = ct.astype(ct.gather(x, offs, mask=mask, check_bounds=False, padding_value=0.0), ct.float32) - b_y = scale * (1.0 / (1.0 + ct.exp(-b_x))) - ct.scatter(y, offs, ct.astype(b_y, y.dtype), mask=mask, check_bounds=False) +def load_bc_bk(arr, base, row0, col0, stride_row, BC: ConstInt, BK: ConstInt, T_eff, K: ConstInt): + # Boundary-checked (BC,BK) block load at (row0,col0) of the (T,K) view + # (row stride stride_row), cast to fp32, on the flattened 1-D view. + o_r = ct.arange(BC, dtype=ct.int32) + o_c = ct.arange(BK, dtype=ct.int32) + rows = row0 + o_r + cols = col0 + o_c + off = base + rows[:, None] * stride_row + cols[None, :] + mask = (rows < T_eff)[:, None] & (cols < K)[None, :] + return ct.astype(ct.gather(arr, off, mask=mask, check_bounds=False, padding_value=0.0), ct.float32) -@ct.kernel -def fused_beta_sigmoid_bwd_kernel(x, dy, dx, scale, n_elements, BLOCK_SIZE: ConstInt): - pid = ct.bid(0) - offs = pid * BLOCK_SIZE + ct.arange(BLOCK_SIZE, dtype=ct.int32) - mask = offs < n_elements - b_x = ct.astype(ct.gather(x, offs, mask=mask, check_bounds=False, padding_value=0.0), ct.float32) - b_dy = ct.astype(ct.gather(dy, offs, mask=mask, check_bounds=False, padding_value=0.0), ct.float32) - b_y = 1.0 / (1.0 + ct.exp(-b_x)) - b_dx = b_dy * scale * b_y * (1.0 - b_y) - ct.scatter(dx, offs, ct.astype(b_dx, dx.dtype), mask=mask, check_bounds=False) +def store_bc_bc(arr, base, row0, col0, stride_row, blk, BC: ConstInt, BT_or_BC: ConstInt, T_eff): + o_r = ct.arange(BC, dtype=ct.int32) + o_c = ct.arange(BC, dtype=ct.int32) + rows = row0 + o_r + cols = col0 + o_c + off = base + rows[:, None] * stride_row + cols[None, :] + mask = (rows < T_eff)[:, None] & (cols < BT_or_BC)[None, :] + ct.scatter(arr, off, ct.astype(blk, arr.dtype), mask=mask, check_bounds=False) -# =========================================================================== -# vector chunk_local_cumsum (used for the vector gate G) -# =========================================================================== -# GATHER RATIONALE: head-interleaved [T, H, S] sub-view; loaded once, -# cumsum'd along T, written back. s/o passed flattened (1-D) by the wrapper. -# The (T, S) sub-view has stride (H*S, 1). +# --- Kernels: normalization and gates ------------------------------------------------------------- @ct.kernel @@ -410,13 +236,6 @@ def chunk_local_cumsum_vector_kernel( ct.scatter(o, off, ct.astype(b_o, o.dtype), mask=m, check_bounds=False) -# =========================================================================== -# kda_gate_chunk_cumsum_vector / kda_gate_bwd -# =========================================================================== -# Vector gate: G is [B, T, H, S]; A_log is [H]; dt_bias is [H*S]. -# GATHER RATIONALE: head-interleaved [B, T, H, S] sub-view streamed once. - - @ct.kernel def kda_gate_chunk_cumsum_vector_kernel( s, @@ -541,25 +360,23 @@ def kda_gate_bwd_kernel( ct.scatter(db, (i_t * H + i_h) * D + o_d, b_db, mask=m_d, check_bounds=False) -# =========================================================================== -# chunk_gla_fwd_kernel_o -- output projection O = QG @ H + A @ V -# =========================================================================== -# GATHER RATIONALE: Q/G/V are reused across the i_k loop (K split), but with -# head-interleaved strides + a transposed STATE_V_FIRST H-view that TMA cannot -# express here, so the non-TMA element-offset gather path is used. +# --- Kernels: WY representation ------------------------------------------------------------------- @ct.kernel -def chunk_gla_fwd_kernel_o( +def recompute_w_u_fwd_kda_kernel( q, + k, + qg, + kg, v, - g, - h, - o, + beta, + w, + u, A, + gk, cu_seqlens, chunk_indices, - scale, H: ConstInt, HV: ConstInt, K: ConstInt, @@ -567,1224 +384,1167 @@ def chunk_gla_fwd_kernel_o( BT: ConstInt, BK: ConstInt, BV: ConstInt, - STATE_V_FIRST: ConstInt, + STORE_U: ConstInt, + STORE_QG: ConstInt, + STORE_KG: ConstInt, ): - i_v = ct.bid(0) - i_t = ct.bid(1) - i_hv = ct.bid(2) + # Arrays arrive pre-flattened to 1-D from the host (cuTile cannot reshape a + # rank-4 dynamic array in-kernel). Use flat element offsets row*stride+col. + k_flat = k + q_flat = q + qg_flat = qg + kg_flat = kg + w_flat = w + gk_flat = gk + v_flat = v + u_flat = u + beta_flat = beta + A_flat = A + + i_t = ct.bid(0) + i_hv = ct.bid(1) i_h = i_hv // (HV // H) - # grid dim-1 is the GLOBAL chunk index; H is laid out per global chunk - # (H-state kernel writes slot chunk_offsets[i_n] + local). Capture it - # before i_t is reassigned to the per-sequence (local) chunk index. - i_tg = i_t i_n = ct.load(chunk_indices, (i_t * 2,), shape=()).item() - i_t = ct.load(chunk_indices, (i_t * 2 + 1,), shape=()).item() + i_t_loc = ct.load(chunk_indices, (i_t * 2 + 1,), shape=()).item() bos = ct.load(cu_seqlens, (i_n,), shape=()).item() eos = ct.load(cu_seqlens, (i_n + 1,), shape=()).item() - T = eos - bos - NT = ct.cdiv(T, BT) - o_bt = ct.arange(BT, dtype=ct.int32) - o_bk = ct.arange(BK, dtype=ct.int32) - o_bv = ct.arange(BV, dtype=ct.int32) - m_s = o_bt[:, None] >= o_bt[None, :] + Tloc = eos - bos - q_base = (bos * H + i_h) * K - g_base = (bos * HV + i_hv) * K - v_base = (bos * HV + i_hv) * V - o_base = (bos * HV + i_hv) * V - h_base = (i_tg * HV + i_hv) * K * V - A_base = (bos * HV + i_hv) * BT + t_off = i_t_loc * BT + ct.arange(BT, dtype=ct.int32) + m_t = t_off < Tloc - t_row = i_t * BT + o_bt - m_t = t_row < T + b_idx = (bos + t_off) * HV + i_hv + b_b = ct.astype(ct.gather(beta_flat, b_idx, mask=m_t, check_bounds=False, padding_value=0.0), ct.float32) - b_o = ct.zeros((BT, BV), dtype=ct.float32) - num_k = (K + BK - 1) // BK - for i_k in range(num_k): - k_col = i_k * BK + o_bk - m_k = k_col < K - v_col = i_v * BV + o_bv - m_v = v_col < V + a_rows = ((bos + t_off) * HV + i_hv)[:, None] + a_cols = ct.arange(BT, dtype=ct.int32)[None, :] + a_off = ct.broadcast_to(a_rows, (BT, BT)) * BT + ct.broadcast_to(a_cols, (BT, BT)) + b_A = ct.gather( + A_flat, + a_off, + mask=ct.broadcast_to(m_t[:, None], (BT, BT)), + check_bounds=False, + padding_value=0.0, + ) - # b_h: STATE_V_FIRST -> view (V,K) block (BV,BK) at (i_v*BV, i_k*BK); - # else -> view (K,V) block (BK,BV) at (i_k*BK, i_v*BV). - if STATE_V_FIRST: - h_off = h_base + v_col[:, None] * K + k_col[None, :] - h_mask = m_v[:, None] & m_k[None, :] - b_h = ct.gather(h, h_off, mask=h_mask, check_bounds=True, padding_value=0.0) - else: - h_off = h_base + k_col[:, None] * V + v_col[None, :] - h_mask = m_k[:, None] & m_v[None, :] - b_h = ct.gather(h, h_off, mask=h_mask, check_bounds=True, padding_value=0.0) + if STORE_U: + v_off = ct.arange(BV, dtype=ct.int32) + for i_v in range(ct.cdiv(V, BV)): + vcols = (i_v * BV + v_off)[None, :] + vrows = ((bos + t_off) * HV + i_hv)[:, None] + m_v = m_t[:, None] & ((i_v * BV + v_off) < V)[None, :] + v_offset = ct.broadcast_to(vrows, (BT, BV)) * V + ct.broadcast_to(vcols, (BT, BV)) + b_v = ct.gather( + v_flat, + v_offset, + mask=m_v, + check_bounds=False, + padding_value=0.0, + ) + b_vb = ct.astype(ct.astype(b_v, ct.float32) * b_b[:, None], b_v.dtype) + b_u = safe_matmul(b_A, b_vb) + ct.scatter( + u_flat, + v_offset, + ct.astype(b_u, u_flat.dtype), + mask=m_v, + check_bounds=False, + ) - q_off = q_base + t_row[:, None] * (H * K) + k_col[None, :] - g_off = g_base + t_row[:, None] * (HV * K) + k_col[None, :] - qg_mask = m_t[:, None] & m_k[None, :] - b_q = ct.gather(q, q_off, mask=qg_mask, check_bounds=True, padding_value=0.0) - b_g = ct.astype(ct.gather(g, g_off, mask=qg_mask, check_bounds=True, padding_value=0.0), ct.float32) - b_qg = ct.astype(b_q * exp2(b_g), b_q.dtype) - - if STATE_V_FIRST: - b_o = safe_mma(b_qg, ct.astype(ct.transpose(b_h), b_qg.dtype), b_o) - else: - b_o = safe_mma(b_qg, ct.astype(b_h, b_qg.dtype), b_o) - - b_o = b_o * scale - - v_col = i_v * BV + o_bv - m_v = v_col < V - v_off = v_base + t_row[:, None] * (HV * V) + v_col[None, :] - v_mask = m_t[:, None] & m_v[None, :] - b_v = ct.gather(v, v_off, mask=v_mask, check_bounds=True, padding_value=0.0) - - A_col = o_bt - A_off = A_base + t_row[:, None] * (HV * BT) + A_col[None, :] - A_mask = m_t[:, None] & (A_col[None, :] < BT) - b_A = ct.gather(A, A_off, mask=A_mask, check_bounds=True, padding_value=0.0) - b_A = ct.astype(ct.where(m_s, b_A, ct.zeros((BT, BT), dtype=b_A.dtype)), b_v.dtype) - b_o = safe_mma(b_A, b_v, b_o) - - o_off = o_base + t_row[:, None] * (HV * V) + v_col[None, :] - ct.scatter(o, o_off, ct.astype(b_o, o.dtype), mask=v_mask, check_bounds=True) + last_idx = ct_min(i_t_loc * BT + BT, Tloc) - 1 + k_off = ct.arange(BK, dtype=ct.int32) + for i_k in range(ct.cdiv(K, BK)): + kcols = (i_k * BK + k_off)[None, :] + m_kcol = ((i_k * BK + k_off) < K)[None, :] + m_k = m_t[:, None] & m_kcol + krows = ((bos + t_off) * H + i_h)[:, None] + gkrows = ((bos + t_off) * HV + i_hv)[:, None] + bk_col = ct.broadcast_to(kcols, (BT, BK)) + k_offset = ct.broadcast_to(krows, (BT, BK)) * K + bk_col + gk_offset = ct.broadcast_to(gkrows, (BT, BK)) * K + bk_col + b_k = ct.gather( + k_flat, + k_offset, + mask=m_k, + check_bounds=False, + padding_value=0.0, + ) + b_gk = ct.astype( + ct.gather( + gk_flat, + gk_offset, + mask=m_k, + check_bounds=False, + padding_value=0.0, + ), + ct.float32, + ) + b_egk = exp2(b_gk) + b_kb = ct.astype(b_k, ct.float32) * b_b[:, None] * b_egk + if STORE_QG: + qrows = ((bos + t_off) * H + i_h)[:, None] + q_offset = ct.broadcast_to(qrows, (BT, BK)) * K + bk_col + b_q = ct.gather( + q_flat, + q_offset, + mask=m_k, + check_bounds=False, + padding_value=0.0, + ) + b_qg = ct.astype(b_q, ct.float32) * b_egk + qgrows = ((bos + t_off) * HV + i_hv)[:, None] + qg_offset = ct.broadcast_to(qgrows, (BT, BK)) * K + bk_col + ct.scatter( + qg_flat, + qg_offset, + ct.astype(b_qg, qg_flat.dtype), + mask=m_k, + check_bounds=False, + ) + if STORE_KG: + gn_rows = ct.broadcast_to(((bos + last_idx) * HV + i_hv), (1, BK)) + gn_off = gn_rows * K + kcols + b_gn = ct.astype( + ct.gather( + gk_flat, + gn_off, + mask=m_kcol, + check_bounds=False, + padding_value=0.0, + ), + ct.float32, + ) + decay = ct.where(m_t[:, None], exp2(b_gn - b_gk), ct.zeros((BT, BK), dtype=ct.float32)) + b_kg = ct.astype(b_k, ct.float32) * decay + kgrows = ((bos + t_off) * HV + i_hv)[:, None] + kg_offset = ct.broadcast_to(kgrows, (BT, BK)) * K + bk_col + ct.scatter( + kg_flat, + kg_offset, + ct.astype(b_kg, kg_flat.dtype), + mask=m_k, + check_bounds=False, + ) -# =========================================================================== -# chunk_delta_h (state-H fwd + dhu bwd) -# =========================================================================== -# For KDA the gate carried is the per-key vector gate -# `gk` (USE_GK=1, USE_G=0); the per-time scalar gate `g` path is inert here. + b_w = safe_matmul(b_A, ct.astype(b_kb, b_k.dtype)) + wrows = ((bos + t_off) * HV + i_hv)[:, None] + w_offset = ct.broadcast_to(wrows, (BT, BK)) * K + bk_col + ct.scatter( + w_flat, + w_offset, + ct.astype(b_w, w_flat.dtype), + mask=m_k, + check_bounds=False, + ) @ct.kernel -def chunk_gated_delta_rule_fwd_kernel_h_blockdim64( +def chunk_kda_fwd_kernel_intra_token_parallel( + q, k, - v, - w, - v_new, g, - gk, - h, - h0, - ht, + beta, + Aqk, + Akk, + scale, cu_seqlens, - chunk_offsets, + N, + T, H: ConstInt, HV: ConstInt, K: ConstInt, - V: ConstInt, BT: ConstInt, - BV: ConstInt, - BK: ConstInt, # next_pow2(K) -- full-width state/K tile (general for any K) - USE_G: ConstInt, - USE_GK: ConstInt, - USE_INITIAL_STATE: ConstInt, - STORE_FINAL_STATE: ConstInt, - SAVE_NEW_VALUE: ConstInt, - STATE_V_FIRST: ConstInt, + BC: ConstInt, + BH: ConstInt, + BK: ConstInt, ): - # The KV state is carried as a SINGLE full-width tile ((BV, BK) when - # STATE_V_FIRST else (BK, BV), BK=next_pow2(K)), so the kernel is general - # for ANY K (no K<=256 cap). Every K-axis load zero-pads cols [K:BK] and - # every K-axis store masks them off, so the tail is 0 on load, contributes - # 0 to every MMA, and stays 0 in the state; a padded Gk tail loads as 0 so - # exp2(0)=1 leaves those zero cols unchanged. - i_v = ct.bid(0) - i_nh = ct.bid(1) - i_n = i_nh // HV - i_h = i_nh % HV + # BK = next_power_of_2(K) is passed by host + # (cuTile tile shapes must be compile-time constants). + i_tg = ct.bid(0) + i_hg = ct.bid(1) + bos = (i_tg // T) * T + i_t = i_tg % T + T_eff = T + left = 0 + right = N + # Unrolled binary search to find i_n s.t. cu[i_n] <= i_tg < cu[i_n+1] + for _ in range(20): + if left < right: + mid = (left + right) // 2 + cmid = ct.load(cu_seqlens, (mid + 1,), shape=()).item() + if i_tg < cmid: + right = mid + else: + left = mid + 1 + i_n = left bos = ct.load(cu_seqlens, (i_n,), shape=()).item() eos = ct.load(cu_seqlens, (i_n + 1,), shape=()).item() - T = eos - bos - NT = ct.cdiv(T, BT) - boh = ct.load(chunk_offsets, (i_n,), shape=()).item() - if STATE_V_FIRST: - b_h = ct.zeros((BV, BK), dtype=ct.float32) - else: - b_h = ct.zeros((BK, BV), dtype=ct.float32) - - h_base = (boh * HV + i_h) * K * V - v_base = (bos * HV + i_h) * V - k_base = (bos * H + i_h // (HV // H)) * K - w_base = (bos * HV + i_h) * K - vnew_base = (bos * HV + i_h) * V - h0_base = i_nh * K * V - ht_base = i_nh * K * V + T_eff = eos - bos + i_t = i_tg - bos - o_bk = ct.arange(BK, dtype=ct.int32) - o_bt = ct.arange(BT, dtype=ct.int32) - o_bv = ct.arange(BV, dtype=ct.int32) + if i_t >= T_eff: + return - # K-tile / V-tile validity masks. The kernel carries a single BK-wide K tile - # (BK=next_pow2(K)) and BV-wide V tiles; when K or V is not a multiple of the - # tile width the extra lanes alias the neighbouring head's H slot. The matmul - # state rows are zeroed via K-masked K/Gk loads, but the raw H/H0/Ht - # gather/scatter are also masked here -> no cross-head corruption when - # K % BK != 0 (i.e. K < BK) or V % BV != 0. - mkh = o_bk < K - mvh = (i_v * BV + o_bv) < V + i_c = i_t // BT + i_s = (i_t % BT) // BC + i_tc = i_c * BT + i_ts = i_tc + i_s * BC - if USE_INITIAL_STATE: - if STATE_V_FIRST: - row = (i_v * BV + o_bv)[:, None] - b_h = b_h + ct.astype( - ct.gather( - h0, - h0_base + row * K + o_bk[None, :], - mask=mvh[:, None] & mkh[None, :], - check_bounds=True, - padding_value=0.0, - ), - ct.float32, - ) - else: - col = (i_v * BV + o_bv)[None, :] - b_h = b_h + ct.astype( - ct.gather( - h0, - h0_base + o_bk[:, None] * V + col, - mask=mkh[:, None] & mvh[None, :], - check_bounds=True, - padding_value=0.0, - ), - ct.float32, - ) + G = HV // H - for i_t in range(NT): - h_chunk = h_base + i_t * HV * K * V + Aqk_base = bos * HV * BT + Akk_base = bos * HV * BC + beta_base = bos * HV - if STATE_V_FIRST: - row = (i_v * BV + o_bv)[:, None] - ct.scatter( - h, - h_chunk + row * K + o_bk[None, :], - ct.astype(b_h, h.dtype), - mask=mvh[:, None] & mkh[None, :], - check_bounds=True, - ) - else: - col = (i_v * BV + o_bv)[None, :] - ct.scatter( - h, - h_chunk + o_bk[:, None] * V + col, - ct.astype(b_h, h.dtype), - mask=mkh[:, None] & mvh[None, :], - check_bounds=True, - ) + # cuTile gather/scatter require the index tuple rank to match the array rank + # (no raw pointer arithmetic). Arrays arrive pre-flattened from the + # host: Q/K -> (B*T*H, K), G -> (B*T*HV, K), Beta/Aqk/Akk -> 1-D. + o_hv = i_hg * BH + ct.arange(BH, dtype=ct.int32) + o_h = o_hv // G + o_k = ct.arange(BK, dtype=ct.int32) + m_hv = o_hv < HV + m_k = o_k < K + m_hk = m_hv[:, None] & m_k[None, :] - w_row = (i_t * BT + o_bt)[:, None] - wmask_r = (i_t * BT + o_bt) < T - # Full-width K: padded cols [K:BK] gather OOB (masked to 0) and row-mask - # zeros the partial-chunk tail, so both contribute 0 to the MMA. - bw = ct.gather(w, w_base + w_row * (HV * K) + o_bk[None, :], mask=mkh[None, :], check_bounds=True, padding_value=0.0) - bw = ct.where(wmask_r[:, None], bw, ct.zeros((BT, BK), dtype=bw.dtype)) - b_v = ct.zeros((BT, BV), dtype=ct.float32) - bmat = ct.astype(ct.transpose(b_h), bw.dtype) if STATE_V_FIRST else ct.astype(b_h, bw.dtype) - b_v = safe_mma(bw, bmat, b_v) + col = ct.broadcast_to(o_k[None, :], (BH, BK)) - v_col = (i_v * BV + o_bv)[None, :] - vmask_c = (i_v * BV + o_bv) < V - v_off = v_base + (i_t * BT + o_bt)[:, None] * (HV * V) + v_col - v_full = (wmask_r[:, None]) & (vmask_c[None, :]) - b_v_load = ct.gather(v, v_off, check_bounds=True, padding_value=0.0) - b_v_load = ct.where(v_full, b_v_load, ct.zeros((BT, BV), dtype=b_v_load.dtype)) - b_v = ct.astype(b_v_load, ct.float32) - b_v + # Q/K: row = (bos + token) * H + head; col = key + qk_row = ct.broadcast_to(((bos + i_t) * H + o_h)[:, None], (BH, BK)) + b_q = ct.astype(ct.gather(q, (qk_row, col), mask=m_hk, check_bounds=False, padding_value=0.0), ct.float32) + b_k = ct.astype(ct.gather(k, (qk_row, col), mask=m_hk, check_bounds=False, padding_value=0.0), ct.float32) - if SAVE_NEW_VALUE: - vn_off = vnew_base + (i_t * BT + o_bt)[:, None] * (HV * V) + v_col - vn_oob = ct.full((BT, BV), v.shape[0], dtype=ct.int32) - vn_off = ct.where(v_full, vn_off, vn_oob) - ct.scatter(v_new, vn_off, ct.astype(b_v, v_new.dtype), check_bounds=True) + # G: row = (bos + token) * HV + head; Beta: idx = (bos + token) * HV + head + g_row = ct.broadcast_to(((bos + i_t) * HV + o_hv)[:, None], (BH, BK)) + b_g = ct.astype(ct.gather(g, (g_row, col), mask=m_hk, check_bounds=False, padding_value=0.0), ct.float32) + b_beta = ct.astype(ct.gather(beta, beta_base + i_t * HV + o_hv, mask=m_hv, check_bounds=False, padding_value=0.0), ct.float32) + b_k = b_k * b_beta[:, None] - last_idx = ct_min((i_t + 1) * BT, T) - 1 - - if USE_G: - m_t = (i_t * BT + o_bt) < T - b_g_last = ct.astype(ct.gather(g, (bos * HV + last_idx * HV + i_h,), check_bounds=True, padding_value=0.0), ct.float32) - g_off = (bos * HV + i_h) + (i_t * BT + o_bt) * HV - b_g = ct.astype(ct.gather(g, g_off, mask=m_t, check_bounds=True, padding_value=0.0), ct.float32) - decay = ct.where(m_t, exp2(b_g_last - b_g), ct.zeros((BT,), dtype=ct.float32)) - b_v = b_v * decay[:, None] - b_g_last = exp2(b_g_last) - b_h = b_h * b_g_last - - if USE_GK: - # Padded tail [K:BK] loads as 0 -> exp2(0)=1 leaves zero state cols unchanged. - gk_base = (bos + last_idx) * HV * K + i_h * K - b_gk = ct.astype(ct.gather(gk, gk_base + o_bk, mask=o_bk < K, check_bounds=True, padding_value=0.0), ct.float32) - if STATE_V_FIRST: - b_h = b_h * exp2(b_gk)[None, :] - else: - b_h = b_h * exp2(b_gk)[:, None] - - b_v = ct.astype(b_v, k.dtype) + # Counted loop over the (static) BC-wide sub-chunk window. A runtime-bounded + # `for j in range(i_ts, j_hi)` lowers to per-iteration branches that tileiras + # cannot unroll/predicate; a counted `for jj in range(BC)` with a runtime + # guard derives `j` from `jj` and stays fully unrolled. + for jj in range(BC): + j = i_ts + jj + if j < i_t + 1 and j < T_eff and j < i_ts + BC: + kj_row = ct.broadcast_to(((bos + j) * H + o_h)[:, None], (BH, BK)) + gj_row = ct.broadcast_to(((bos + j) * HV + o_hv)[:, None], (BH, BK)) + b_kj = ct.astype(ct.gather(k, (kj_row, col), mask=m_hk, check_bounds=False, padding_value=0.0), ct.float32) + b_gj = ct.astype(ct.gather(g, (gj_row, col), mask=m_hk, check_bounds=False, padding_value=0.0), ct.float32) - k_row = o_bk[:, None] - k_time = (i_t * BT + o_bt)[None, :] - kmask = (o_bk[:, None] < K) & ((i_t * BT + o_bt)[None, :] < T) - bk = ct.gather(k, k_base + k_row * 1 + k_time * (H * K), check_bounds=True, padding_value=0.0) - bk = ct.where(kmask, bk, ct.zeros((BK, BT), dtype=bk.dtype)) - prod = safe_mma(bk, b_v, ct.zeros((BK, BV), dtype=ct.float32)) - if STATE_V_FIRST: - b_h = b_h + ct.transpose(prod) - else: - b_h = b_h + prod + b_kgj = ct.where(m_k[None, :], b_kj * exp2(b_g - b_gj), ct.zeros((BH, BK), dtype=ct.float32)) + b_Aqk = ct.sum(b_q * b_kgj, axis=1) * scale + b_Akk = ct.sum(b_k * b_kgj, axis=1) * (1.0 if j < i_t else 0.0) - if STORE_FINAL_STATE: - if STATE_V_FIRST: - row = (i_v * BV + o_bv)[:, None] ct.scatter( - ht, - ht_base + row * K + o_bk[None, :], - ct.astype(b_h, ht.dtype), - mask=mvh[:, None] & mkh[None, :], - check_bounds=True, + Aqk, + Aqk_base + i_t * HV * BT + o_hv * BT + (j % BT), + ct.astype(b_Aqk, Aqk.dtype), + mask=m_hv, + check_bounds=False, ) - else: - col = (i_v * BV + o_bv)[None, :] ct.scatter( - ht, - ht_base + o_bk[:, None] * V + col, - ct.astype(b_h, ht.dtype), - mask=mkh[:, None] & mvh[None, :], - check_bounds=True, + Akk, + Akk_base + i_t * HV * BC + o_hv * BC + (j - i_ts), + ct.astype(b_Akk, Akk.dtype), + mask=m_hv, + check_bounds=False, ) @ct.kernel -def chunk_gated_delta_rule_bwd_kernel_dhu_blockdim64( +def chunk_kda_fwd_kernel_inter_diag_compute_solve( q, k, - w, g, - gk, - dstate_in, - dh0, - do, - dh, - dv, - dv2, - cu_seqlens, - chunk_offsets, + beta, + Aqk, + Akk, scale, + cu_seqlens, + chunk_indices, H: ConstInt, HV: ConstInt, K: ConstInt, - V: ConstInt, BT: ConstInt, - BV: ConstInt, - BK: ConstInt, # next_pow2(K) -- full-width state/K tile (general for any K) - USE_G: ConstInt, - USE_GK: ConstInt, - USE_INITIAL_STATE: ConstInt, - USE_FINAL_STATE_GRADIENT: ConstInt, - STATE_V_FIRST: ConstInt, + BC: ConstInt, + BK: ConstInt, ): - # dH state is a SINGLE full-width tile ((BV, BK) when STATE_V_FIRST else - # (BK, BV), BK=next_pow2(K)), so the kernel is general for ANY K (no K<=256 - # cap). Every flat gather/scatter over the K axis uses oK=arange(BK) with a - # `(oK < K)` mask, so rows/cols [K:BK] are zero on load, contribute 0 to - # every MMA, and are masked out on store; a Gk padding tail loads as 0 so - # exp2(0)=1 leaves those zero state cols unchanged. - i_v = ct.bid(0) - i_nh = ct.bid(1) - i_n = i_nh // HV - i_h = i_nh % HV + i_t = ct.bid(0) + i_i = ct.bid(1) + i_hv = ct.bid(2) + i_h = i_hv // (HV // H) + + i_n = ct.load(chunk_indices, (i_t * 2,), shape=()).item() + i_t = ct.load(chunk_indices, (i_t * 2 + 1,), shape=()).item() bos = ct.load(cu_seqlens, (i_n,), shape=()).item() eos = ct.load(cu_seqlens, (i_n + 1,), shape=()).item() - T = eos - bos - NT = ct.cdiv(T, BT) - boh = ct.load(chunk_offsets, (i_n,), shape=()).item() - if STATE_V_FIRST: - b_dh = ct.zeros((BV, BK), dtype=ct.float32) - else: - b_dh = ct.zeros((BK, BV), dtype=ct.float32) + T_eff = eos - bos - q_base = (bos * H + i_h // (HV // H)) * K - k_base = (bos * H + i_h // (HV // H)) * K - w_base = (bos * HV + i_h) * K - do_base = (bos * HV + i_h) * V - dv_base = (bos * HV + i_h) * V - dv2_base = (bos * HV + i_h) * V - dh_base = (boh * HV + i_h) * K * V - gk_base = (bos * HV + i_h) * K - dh0_base = i_nh * K * V - dht_base = i_nh * K * V + i_ti = i_t * BT + i_i * BC + if i_ti >= T_eff: + return + o_bc = ct.arange(BC, dtype=ct.int32) o_bk = ct.arange(BK, dtype=ct.int32) - o_bt = ct.arange(BT, dtype=ct.int32) - o_bv = ct.arange(BV, dtype=ct.int32) + o_c = i_ti + o_bc + m_c = o_c < T_eff + m_k = o_bk < K - # V-boundary mask for the state-gradient (dH/dH0) scatters: a partial - # trailing V tile (V not a multiple of BV) must not write rows/cols >= V, - # else the store spills into the neighbouring chunk/head's state-gradient - # (boundary-checked store semantics). - m_v = (i_v * BV + o_bv) < V + q_base = (bos * H + i_h) * K + k_base = (bos * H + i_h) * K + g_base = (bos * HV + i_hv) * K + beta_base = bos * HV + i_hv + Aqk_base = (bos * HV + i_hv) * BT + Akk_base = (bos * HV + i_hv) * BC - # --- Load final state gradient dHt -> b_dh (single full-width tile; [K:BK] loads 0) --- - if USE_FINAL_STATE_GRADIENT: - if STATE_V_FIRST: - row = (i_v * BV + o_bv)[:, None] - b_dh = b_dh + ct.gather( - dstate_in, - dht_base + row * K + o_bk[None, :], - mask=(o_bk < K)[None, :], - check_bounds=True, - padding_value=0.0, - ) - else: - col = (i_v * BV + o_bv)[None, :] - b_dh = b_dh + ct.gather( - dstate_in, - dht_base + o_bk[:, None] * V + col, - mask=(o_bk < K)[:, None], - check_bounds=True, - padding_value=0.0, - ) + qk_rows = o_c[:, None] + qk_cols = o_bk[None, :] + m_qk = m_c[:, None] & m_k[None, :] + q_off = q_base + qk_rows * (H * K) + qk_cols + k_off = k_base + qk_rows * (H * K) + qk_cols + g_off = g_base + qk_rows * (HV * K) + qk_cols + b_q = ct.gather(q, q_off, mask=m_qk, check_bounds=False, padding_value=0.0) + b_k = ct.gather(k, k_off, mask=m_qk, check_bounds=False, padding_value=0.0) + b_g = ct.gather(g, g_off, mask=m_qk, check_bounds=False, padding_value=0.0) + b_beta = ct.gather(beta, beta_base + o_c * HV, mask=m_c, check_bounds=False, padding_value=0.0) - # cuTile range() requires a positive step; iterate forward and reverse the - # index to preserve the backward (NT-1 .. 0) chunk traversal. - for _i_t in range(NT): - i_t = NT - 1 - _i_t - dh_chunk = dh_base + i_t * HV * K * V + gn_row = i_ti + ct_min(BC // 2, T_eff - i_ti - 1) + b_gn = ct.gather(g, g_base + gn_row * (HV * K) + o_bk, mask=m_k, check_bounds=False, padding_value=0.0) + b_gn = ct.astype(b_gn, ct.float32)[None, :] - # Store current b_dh to dH[i_t] (single full-width tile; [K:BK] masked out) - if STATE_V_FIRST: - row = (i_v * BV + o_bv)[:, None] - ct.scatter( - dh, - dh_chunk + row * K + o_bk[None, :], - ct.astype(b_dh, dh.dtype), - mask=m_v[:, None] & (o_bk < K)[None, :], - check_bounds=False, - ) - else: - col = (i_v * BV + o_bv)[None, :] - ct.scatter( - dh, - dh_chunk + o_bk[:, None] * V + col, - ct.astype(b_dh, dh.dtype), - mask=(o_bk < K)[:, None] & m_v[None, :], - check_bounds=False, - ) + b_gm = ct.astype(b_g, ct.float32) - b_gn + b_gq = ct.where(m_c[:, None], exp2(b_gm), ct.zeros((BC, BK), dtype=ct.float32)) + b_gk = ct.where(m_c[:, None], exp2(-b_gm), ct.zeros((BC, BK), dtype=ct.float32)) - last_idx = ct_min((i_t + 1) * BT, T) - 1 + if K < 256: + b_kgt = ct.transpose(ct.astype(ct.astype(b_k, ct.float32) * b_gk, b_k.dtype)) + b_qg = ct.astype(ct.astype(b_q, ct.float32) * b_gq, b_q.dtype) + b_kg = ct.astype(ct.astype(b_k, ct.float32) * b_gq, b_k.dtype) + else: + b_kgt = ct.transpose(ct.astype(b_k, ct.float32) * b_gk) + b_qg = ct.astype(b_q, ct.float32) * b_gq + b_kg = ct.astype(b_k, ct.float32) * b_gq - m_t = (i_t * BT + o_bt) < T - if USE_G: - bg_last = ct.astype(ct.gather(g, ((bos + last_idx) * HV + i_h,), check_bounds=True, padding_value=0.0), ct.float32) - g_off = (bos * HV + i_h) + (i_t * BT + o_bt) * HV - b_g = ct.astype(ct.gather(g, g_off, mask=m_t, check_bounds=True, padding_value=0.0), ct.float32) - bg_last_exp = exp2(bg_last) - b_g_exp = exp2(b_g) - else: - bg_last_exp = ct.astype(0.0, ct.float32) - b_g_exp = ct.zeros((BT,), dtype=ct.float32) + b_Aqk = safe_matmul(b_qg, b_kgt) * scale + b_Akk = safe_matmul(b_kg, b_kgt) * ct.astype(b_beta, ct.float32)[:, None] - # dO, dV, dV2 tiles - v_col = (i_v * BV + o_bv)[None, :] - vmask_c = (i_v * BV + o_bv) < V - v_full = (m_t[:, None]) & (vmask_c[None, :]) - do_off = do_base + (i_t * BT + o_bt)[:, None] * (HV * V) + v_col - b_do = ct.gather(do, do_off, check_bounds=True, padding_value=0.0) - b_do = ct.where(v_full, b_do, ct.zeros((BT, BV), dtype=b_do.dtype)) + o_i = o_bc + m_Aqk = o_i[:, None] >= o_i[None, :] + m_Akk = o_i[:, None] > o_i[None, :] + m_I = o_i[:, None] == o_i[None, :] - # b_dv = b_k @ b_dh (single full-width K MMA). b_k tile is (BT, BK); [K:BK] - # is masked to 0 so it contributes nothing. - bk = ct.gather( - k, - k_base + (i_t * BT + o_bt)[:, None] * (H * K) + o_bk[None, :], - mask=(o_bk < K)[None, :], - check_bounds=True, - padding_value=0.0, - ) - bk = ct.where(m_t[:, None], bk, ct.zeros((BT, BK), dtype=bk.dtype)) - if USE_GK: - # Gk offset: base (bos*HV+i_h)*K ; then + last_idx*HV*K + o_k. - # Padded tail [K:BK] loads as 0 -> exp2(0)=1 leaves those zero cols unchanged. - gkl = (bos * HV + i_h) * K + last_idx * HV * K - b_gk_last = ct.astype(ct.gather(gk, gkl + o_bk, mask=o_bk < K, check_bounds=True, padding_value=0.0), ct.float32) - bmat = ct.astype(ct.transpose(b_dh), bk.dtype) if STATE_V_FIRST else ct.astype(b_dh, bk.dtype) - b_dv = safe_mma(bk, bmat, ct.zeros((BT, BV), dtype=ct.float32)) - - if USE_G: - decay = ct.where(m_t, exp2(bg_last - b_g), ct.zeros((BT,), dtype=ct.float32)) - b_dv = b_dv * decay[:, None] - - # b_dv += dV ; store to dV2 - dv_off = dv_base + (i_t * BT + o_bt)[:, None] * (HV * V) + v_col - b_dv_load = ct.gather(dv, dv_off, check_bounds=True, padding_value=0.0) - b_dv_load = ct.where(v_full, b_dv_load, ct.zeros((BT, BV), dtype=b_dv_load.dtype)) - b_dv = b_dv + ct.astype(b_dv_load, ct.float32) - dv2_off = dv2_base + (i_t * BT + o_bt)[:, None] * (HV * V) + v_col - dv2_oob = ct.full((BT, BV), dv2.shape[0], dtype=ct.int32) - dv2_off = ct.where(v_full, dv2_off, dv2_oob) - ct.scatter(dv2, dv2_off, ct.astype(b_dv, dv2.dtype), check_bounds=True) - - # b_dh += trans(b_q@b_do*scale - b_w@b_dv) (b_q,b_w are (BK,BT) transposed; - # rows [K:BK] are masked to 0 so contribute nothing to the update) - b_dv_c = ct.astype(b_dv, do.dtype) - time = (i_t * BT + o_bt)[None, :] - tmask = (i_t * BT + o_bt)[None, :] < T - kr = o_bk[:, None] - wmask = (o_bk[:, None] < K) & tmask - b_w = ct.gather(w, w_base + kr * 1 + time * (HV * K), check_bounds=True, padding_value=0.0) - b_w = ct.where(wmask, b_w, ct.zeros((BK, BT), dtype=b_w.dtype)) - b_q = ct.gather(q, q_base + kr * 1 + time * (H * K), check_bounds=True, padding_value=0.0) - b_q = ct.where(wmask, b_q, ct.zeros((BK, BT), dtype=b_q.dtype)) - if USE_G: - b_dh = b_dh * bg_last_exp - b_q = b_q * b_g_exp[None, :] - if USE_GK: - if STATE_V_FIRST: - b_dh = b_dh * exp2(b_gk_last)[None, :] - else: - b_dh = b_dh * exp2(b_gk_last[:, None]) - term = safe_matmul(b_q, b_do) * scale - safe_matmul(b_w, b_dv_c) - if STATE_V_FIRST: - b_dh = b_dh + ct.transpose(term) - else: - b_dh = b_dh + term + b_Aqk = ct.where(m_Aqk, b_Aqk, ct.zeros((BC, BC), dtype=ct.float32)) + b_Akk = ct.where(m_Akk, b_Akk, ct.zeros((BC, BC), dtype=ct.float32)) - if USE_INITIAL_STATE: - if STATE_V_FIRST: - row = (i_v * BV + o_bv)[:, None] - ct.scatter( - dh0, - dh0_base + row * K + o_bk[None, :], - ct.astype(b_dh, dh0.dtype), - mask=m_v[:, None] & (o_bk < K)[None, :], - check_bounds=False, - ) - else: - col = (i_v * BV + o_bv)[None, :] - ct.scatter( - dh0, - dh0_base + o_bk[:, None] * V + col, - ct.astype(b_dh, dh0.dtype), - mask=(o_bk < K)[:, None] & m_v[None, :], - check_bounds=False, - ) + # store Aqk (Akk for this kernel writes to the fp32 diagonal buffer) + aqk_rows = o_c[:, None] + aqk_cols = (i_i * BC + o_bc)[None, :] + aqk_off = Aqk_base + aqk_rows * (HV * BT) + aqk_cols + m_aqk = m_c[:, None] & ((i_i * BC + o_bc) < BT)[None, :] + ct.scatter(Aqk, aqk_off, ct.astype(b_Aqk, Aqk.dtype), mask=m_aqk, check_bounds=False) + # diagonal Akk -> inverse via Neumann series by squaring, written to Akk(diag buf). + # + # Akk is strictly-lower-triangular and nilpotent (Akk^BC = 0); with + # N := -Akk, (I + Akk)^-1 = (I - N)^-1 = sum_{k=0}^{BC-1} N^k + # = prod_{j=0}^{log2(BC)-1} (I + N^(2^j)) -- log2(BC) block matmuls, no + # serial row dependency. Rows beyond T_eff have N = 0 (b_gq/b_gk were + # zeroed by m_c) and converge to the identity. + # + # Precision: the squarings run at tf32 via safe_matmul (fp32 operands + # would fall back to SIMT even at M=32). + b_N = -b_Akk # N := -Akk, so (I + Akk)^-1 = (I - N)^-1 = sum_k N^k + b_Ai = ct.astype(m_I, ct.float32) + b_N # (I + N) + # Squaring stages: `range(2, BC)` traces as a compile-time-bounded loop with + # concrete Python `i` during unroll (BC is a ConstInt). We do one squaring + # each time `i` is a power of two (i = 2, 4, 8, 16 for BC=32), giving exactly + # ceil(log2(BC)) - 1 stages (factors N^2, N^4, N^8, N^16). All the branch + # decisions are host-side (Python int `i`), so nothing lowers to a device + # branch -- the body is fully unrolled into log2(BC) block matmuls. + for i in range(2, BC): + if (i & (i - 1)) == 0: # i is a power of two (host-time test) + b_N = safe_matmul(b_N, b_N) # N -> N^2 -> N^4 ... + b_Ai = b_Ai + safe_matmul(b_Ai, b_N) # (prod so far) @ (I + N^(2^j)) -# =========================================================================== -# wy_fast (recompute_w_u_fwd) -# =========================================================================== -# 2D structured gather/scatter on reshaped flat views; head-interleaved -# strides folded into the (row, col) tuple indices. Gk is the per-key gate. + akk_rows = o_c[:, None] + akk_cols = o_bc[None, :] + akk_off = Akk_base + akk_rows * (HV * BC) + akk_cols + m_akk = m_c[:, None] & (o_bc < BC)[None, :] + ct.scatter(Akk, akk_off, ct.astype(b_Ai, Akk.dtype), mask=m_akk, check_bounds=False) @ct.kernel -def recompute_w_u_fwd_kda_kernel( +def chunk_kda_fwd_kernel_inter_solve_fused( q, k, - qg, - kg, - v, + g, beta, - w, - u, - A, - gk, + Aqk, + Akkd, + Akk, + scale, cu_seqlens, chunk_indices, H: ConstInt, HV: ConstInt, K: ConstInt, - V: ConstInt, BT: ConstInt, + BC: ConstInt, + NC: ConstInt, BK: ConstInt, - BV: ConstInt, - STORE_U: ConstInt, - STORE_QG: ConstInt, - STORE_KG: ConstInt, + USE_SAFE_GATE: ConstInt, ): - # Arrays arrive pre-flattened to 1-D from the host (cuTile cannot reshape a - # rank-4 dynamic array in-kernel). Use flat element offsets row*stride+col. - k_flat = k - q_flat = q - qg_flat = qg - kg_flat = kg - w_flat = w - gk_flat = gk - v_flat = v - u_flat = u - beta_flat = beta - A_flat = A - i_t = ct.bid(0) i_hv = ct.bid(1) i_h = i_hv // (HV // H) i_n = ct.load(chunk_indices, (i_t * 2,), shape=()).item() - i_t_loc = ct.load(chunk_indices, (i_t * 2 + 1,), shape=()).item() + i_t = ct.load(chunk_indices, (i_t * 2 + 1,), shape=()).item() bos = ct.load(cu_seqlens, (i_n,), shape=()).item() eos = ct.load(cu_seqlens, (i_n + 1,), shape=()).item() - Tloc = eos - bos + T_eff = eos - bos - t_off = i_t_loc * BT + ct.arange(BT, dtype=ct.int32) - m_t = t_off < Tloc + if i_t * BT >= T_eff: + return - b_idx = (bos + t_off) * HV + i_hv - b_b = ct.astype(ct.gather(beta_flat, b_idx, mask=m_t, check_bounds=False, padding_value=0.0), ct.float32) + i_tc0 = i_t * BT + i_tc1 = i_t * BT + BC + i_tc2 = i_t * BT + 2 * BC + i_tc3 = i_t * BT + 3 * BC - a_rows = ((bos + t_off) * HV + i_hv)[:, None] - a_cols = ct.arange(BT, dtype=ct.int32)[None, :] - a_off = ct.broadcast_to(a_rows, (BT, BT)) * BT + ct.broadcast_to(a_cols, (BT, BT)) - b_A = ct.gather( - A_flat, - a_off, - mask=ct.broadcast_to(m_t[:, None], (BT, BT)), - check_bounds=False, - padding_value=0.0, - ) + q_base = (bos * H + i_h) * K + k_base = (bos * H + i_h) * K + g_base = (bos * HV + i_hv) * K + Aqk_base = (bos * HV + i_hv) * BT + Akk_base = (bos * HV + i_hv) * BT + Akkd_base = (bos * HV + i_hv) * BC + beta_base = bos * HV + i_hv - if STORE_U: - v_off = ct.arange(BV, dtype=ct.int32) - for i_v in range(ct.cdiv(V, BV)): - vcols = (i_v * BV + v_off)[None, :] - vrows = ((bos + t_off) * HV + i_hv)[:, None] - m_v = m_t[:, None] & ((i_v * BV + v_off) < V)[None, :] - v_offset = ct.broadcast_to(vrows, (BT, BV)) * V + ct.broadcast_to(vcols, (BT, BV)) - b_v = ct.gather( - v_flat, - v_offset, - mask=m_v, - check_bounds=False, - padding_value=0.0, - ) - b_vb = ct.astype(ct.astype(b_v, ct.float32) * b_b[:, None], b_v.dtype) - b_u = safe_matmul(b_A, b_vb) - ct.scatter( - u_flat, - v_offset, - ct.astype(b_u, u_flat.dtype), - mask=m_v, - check_bounds=False, + o_i = ct.arange(BC, dtype=ct.int32) + o_k = ct.arange(BK, dtype=ct.int32) + m_tc1 = (i_tc1 + o_i) < T_eff + m_tc2 = (i_tc2 + o_i) < T_eff + m_tc3 = (i_tc3 + o_i) < T_eff + + z = ct.zeros((BC, BC), dtype=ct.float32) + b_Aqk10 = z + b_Akk10 = z + b_Aqk20 = z + b_Akk20 = z + b_Aqk21 = z + b_Akk21 = z + b_Aqk30 = z + b_Akk30 = z + b_Aqk31 = z + b_Akk31 = z + b_Aqk32 = z + b_Akk32 = z + + # ---- off-diagonal blocks ----------------------------------------------------- + num_k = (K + BK - 1) // BK + for i_k in range(num_k): + kk = i_k * BK + o_k + m_k = kk < K + b_k0 = load_bc_bk(k, k_base, i_tc0, i_k * BK, H * K, BC, BK, T_eff, K) + b_g0 = load_bc_bk(g, g_base, i_tc0, i_k * BK, HV * K, BC, BK, T_eff, K) + + if i_tc1 < T_eff: + b_q1 = load_bc_bk(q, q_base, i_tc1, i_k * BK, H * K, BC, BK, T_eff, K) + b_k1 = load_bc_bk(k, k_base, i_tc1, i_k * BK, H * K, BC, BK, T_eff, K) + b_g1 = load_bc_bk(g, g_base, i_tc1, i_k * BK, HV * K, BC, BK, T_eff, K) + b_gn1 = ct.astype( + ct.gather(g, g_base + i_tc1 * (HV * K) + kk, mask=m_k, check_bounds=False, padding_value=0.0), + ct.float32, ) + b_gqn = ct.where(m_tc1[:, None], exp2(b_g1 - b_gn1[None, :]), ct.zeros((BC, BK), dtype=ct.float32)) + b_kgt = ct.transpose(b_k0 * exp2(b_gn1[None, :] - b_g0)) + b_Aqk10 = safe_mma(b_q1 * b_gqn, b_kgt, b_Aqk10) + b_Akk10 = safe_mma(b_k1 * b_gqn, b_kgt, b_Akk10) - last_idx = ct_min(i_t_loc * BT + BT, Tloc) - 1 - k_off = ct.arange(BK, dtype=ct.int32) - for i_k in range(ct.cdiv(K, BK)): - kcols = (i_k * BK + k_off)[None, :] - m_kcol = ((i_k * BK + k_off) < K)[None, :] - m_k = m_t[:, None] & m_kcol - krows = ((bos + t_off) * H + i_h)[:, None] - gkrows = ((bos + t_off) * HV + i_hv)[:, None] - bk_col = ct.broadcast_to(kcols, (BT, BK)) - k_offset = ct.broadcast_to(krows, (BT, BK)) * K + bk_col - gk_offset = ct.broadcast_to(gkrows, (BT, BK)) * K + bk_col - b_k = ct.gather( - k_flat, - k_offset, - mask=m_k, - check_bounds=False, - padding_value=0.0, + if NC >= 3 and i_tc2 < T_eff: + b_q2 = load_bc_bk(q, q_base, i_tc2, i_k * BK, H * K, BC, BK, T_eff, K) + b_k2 = load_bc_bk(k, k_base, i_tc2, i_k * BK, H * K, BC, BK, T_eff, K) + b_g2 = load_bc_bk(g, g_base, i_tc2, i_k * BK, HV * K, BC, BK, T_eff, K) + b_gn2 = ct.astype( + ct.gather(g, g_base + i_tc2 * (HV * K) + kk, mask=m_k, check_bounds=False, padding_value=0.0), + ct.float32, + ) + b_gqn2 = ct.where(m_tc2[:, None], exp2(b_g2 - b_gn2[None, :]), ct.zeros((BC, BK), dtype=ct.float32)) + b_qg2 = b_q2 * b_gqn2 + b_kg2 = b_k2 * b_gqn2 + b_kgt = ct.transpose(b_k0 * exp2(b_gn2[None, :] - b_g0)) + b_Aqk20 = safe_mma(b_qg2, b_kgt, b_Aqk20) + b_Akk20 = safe_mma(b_kg2, b_kgt, b_Akk20) + b_kgt = ct.transpose(b_k1 * exp2(b_gn2[None, :] - b_g1)) + b_Aqk21 = safe_mma(b_qg2, b_kgt, b_Aqk21) + b_Akk21 = safe_mma(b_kg2, b_kgt, b_Akk21) + + if NC >= 4 and i_tc3 < T_eff: + b_q3 = load_bc_bk(q, q_base, i_tc3, i_k * BK, H * K, BC, BK, T_eff, K) + b_k3 = load_bc_bk(k, k_base, i_tc3, i_k * BK, H * K, BC, BK, T_eff, K) + b_g3 = load_bc_bk(g, g_base, i_tc3, i_k * BK, HV * K, BC, BK, T_eff, K) + b_gn3 = ct.astype( + ct.gather(g, g_base + i_tc3 * (HV * K) + kk, mask=m_k, check_bounds=False, padding_value=0.0), + ct.float32, + ) + b_gqn3 = ct.where(m_tc3[:, None], exp2(b_g3 - b_gn3[None, :]), ct.zeros((BC, BK), dtype=ct.float32)) + b_qg3 = b_q3 * b_gqn3 + b_kg3 = b_k3 * b_gqn3 + b_kgt = ct.transpose(b_k0 * exp2(b_gn3[None, :] - b_g0)) + b_Aqk30 = safe_mma(b_qg3, b_kgt, b_Aqk30) + b_Akk30 = safe_mma(b_kg3, b_kgt, b_Akk30) + b_kgt = ct.transpose(b_k1 * exp2(b_gn3[None, :] - b_g1)) + b_Aqk31 = safe_mma(b_qg3, b_kgt, b_Aqk31) + b_Akk31 = safe_mma(b_kg3, b_kgt, b_Akk31) + b_kgt = ct.transpose(b_k2 * exp2(b_gn3[None, :] - b_g2)) + b_Aqk32 = safe_mma(b_qg3, b_kgt, b_Aqk32) + b_Akk32 = safe_mma(b_kg3, b_kgt, b_Akk32) + + # ---- save off-diagonal Aqk blocks, scale Akk by Beta ------------------------- + if i_tc1 < T_eff: + store_bc_bc(Aqk, Aqk_base, i_tc1, 0, HV * BT, b_Aqk10 * scale, BC, BT, T_eff) + b_b1 = ct.astype( + ct.gather(beta, beta_base + (i_tc1 + o_i) * HV, mask=m_tc1, check_bounds=False, padding_value=0.0), + ct.float32, ) - b_gk = ct.astype( - ct.gather( - gk_flat, - gk_offset, - mask=m_k, - check_bounds=False, - padding_value=0.0, - ), + b_Akk10 = b_Akk10 * b_b1[:, None] + if NC >= 3 and i_tc2 < T_eff: + store_bc_bc(Aqk, Aqk_base, i_tc2, 0, HV * BT, b_Aqk20 * scale, BC, BT, T_eff) + store_bc_bc(Aqk, Aqk_base, i_tc2, BC, HV * BT, b_Aqk21 * scale, BC, BT, T_eff) + b_b2 = ct.astype( + ct.gather(beta, beta_base + (i_tc2 + o_i) * HV, mask=m_tc2, check_bounds=False, padding_value=0.0), ct.float32, ) - b_egk = exp2(b_gk) - b_kb = ct.astype(b_k, ct.float32) * b_b[:, None] * b_egk + b_Akk20 = b_Akk20 * b_b2[:, None] + b_Akk21 = b_Akk21 * b_b2[:, None] + if NC >= 4 and i_tc3 < T_eff: + store_bc_bc(Aqk, Aqk_base, i_tc3, 0, HV * BT, b_Aqk30 * scale, BC, BT, T_eff) + store_bc_bc(Aqk, Aqk_base, i_tc3, BC, HV * BT, b_Aqk31 * scale, BC, BT, T_eff) + store_bc_bc(Aqk, Aqk_base, i_tc3, 2 * BC, HV * BT, b_Aqk32 * scale, BC, BT, T_eff) + b_b3 = ct.astype( + ct.gather(beta, beta_base + (i_tc3 + o_i) * HV, mask=m_tc3, check_bounds=False, padding_value=0.0), + ct.float32, + ) + b_Akk30 = b_Akk30 * b_b3[:, None] + b_Akk31 = b_Akk31 * b_b3[:, None] + b_Akk32 = b_Akk32 * b_b3[:, None] - if STORE_QG: - qrows = ((bos + t_off) * H + i_h)[:, None] - q_offset = ct.broadcast_to(qrows, (BT, BK)) * K + bk_col - b_q = ct.gather( - q_flat, - q_offset, - mask=m_k, - check_bounds=False, - padding_value=0.0, - ) - b_qg = ct.astype(b_q, ct.float32) * b_egk - qgrows = ((bos + t_off) * HV + i_hv)[:, None] - qg_offset = ct.broadcast_to(qgrows, (BT, BK)) * K + bk_col - ct.scatter( - qg_flat, - qg_offset, - ct.astype(b_qg, qg_flat.dtype), - mask=m_k, - check_bounds=False, - ) - if STORE_KG: - gn_rows = ct.broadcast_to(((bos + last_idx) * HV + i_hv), (1, BK)) - gn_off = gn_rows * K + kcols - b_gn = ct.astype( - ct.gather( - gk_flat, - gn_off, - mask=m_kcol, - check_bounds=False, - padding_value=0.0, - ), - ct.float32, - ) - decay = ct.where(m_t[:, None], exp2(b_gn - b_gk), ct.zeros((BT, BK), dtype=ct.float32)) - b_kg = ct.astype(b_k, ct.float32) * decay - kgrows = ((bos + t_off) * HV + i_hv)[:, None] - kg_offset = ct.broadcast_to(kgrows, (BT, BK)) * K + bk_col - ct.scatter( - kg_flat, - kg_offset, - ct.astype(b_kg, kg_flat.dtype), - mask=m_k, - check_bounds=False, - ) + # ---- load diagonal inverse blocks from Akkd (fp32) --------------------------- + b_Ai00 = load_bc_bk(Akkd, Akkd_base, i_tc0, 0, HV * BC, BC, BC, T_eff, BC) + b_Ai11 = load_bc_bk(Akkd, Akkd_base, i_tc1, 0, HV * BC, BC, BC, T_eff, BC) + b_Ai22 = load_bc_bk(Akkd, Akkd_base, i_tc2, 0, HV * BC, BC, BC, T_eff, BC) if NC >= 3 else z + b_Ai33 = load_bc_bk(Akkd, Akkd_base, i_tc3, 0, HV * BC, BC, BC, T_eff, BC) if NC >= 4 else z + + # ---- forward substitution on diagonals (only when gate not pre-solved) ------- + if not USE_SAFE_GATE: + m_A = o_i[:, None] > o_i[None, :] + m_I = o_i[:, None] == o_i[None, :] + b_Ai00 = -ct.where(m_A, b_Ai00, z) + b_Ai11 = -ct.where(m_A, b_Ai11, z) + if NC >= 3: + b_Ai22 = -ct.where(m_A, b_Ai22, z) + if NC >= 4: + b_Ai33 = -ct.where(m_A, b_Ai33, z) + + # Counted loops over the static [BC] forward-substitution window with a + # runtime guard (i < T_eff - i_tc0). A runtime upper bound (ct_min(..., + # T_eff - i_tc0)) lowers to per-iteration branches tileiras can't + # unroll/predicate; the counted form stays fully unrolled. + for i in range(2, BC): + if i < T_eff - i_tc0: + b_a00 = -ct.astype( + ct.gather(Akkd, Akkd_base + (i_tc0 + i) * (HV * BC) + o_i, check_bounds=False, padding_value=0.0), + ct.float32, + ) + b_a00 = ct.where(o_i < i, b_a00, ct.zeros((BC,), dtype=ct.float32)) + b_a00 = b_a00 + ct.sum(b_a00[:, None] * b_Ai00, axis=0) + b_Ai00 = ct.where((o_i == i)[:, None], b_a00, b_Ai00) + for i in range(BC + 2, 2 * BC): + if i < T_eff - i_tc0: + b_a11 = -ct.astype( + ct.gather(Akkd, Akkd_base + (i_tc0 + i) * (HV * BC) + o_i, check_bounds=False, padding_value=0.0), + ct.float32, + ) + b_a11 = ct.where(o_i < i - BC, b_a11, ct.zeros((BC,), dtype=ct.float32)) + b_a11 = b_a11 + ct.sum(b_a11[:, None] * b_Ai11, axis=0) + b_Ai11 = ct.where((o_i == i - BC)[:, None], b_a11, b_Ai11) + if NC >= 3: + for i in range(2 * BC + 2, 3 * BC): + if i < T_eff - i_tc0: + b_a22 = -ct.astype( + ct.gather(Akkd, Akkd_base + (i_tc0 + i) * (HV * BC) + o_i, check_bounds=False, padding_value=0.0), + ct.float32, + ) + b_a22 = ct.where(o_i < i - 2 * BC, b_a22, ct.zeros((BC,), dtype=ct.float32)) + b_a22 = b_a22 + ct.sum(b_a22[:, None] * b_Ai22, axis=0) + b_Ai22 = ct.where((o_i == i - 2 * BC)[:, None], b_a22, b_Ai22) + if NC >= 4: + for i in range(3 * BC + 2, 4 * BC): + if i < T_eff - i_tc0: + b_a33 = -ct.astype( + ct.gather(Akkd, Akkd_base + (i_tc0 + i) * (HV * BC) + o_i, check_bounds=False, padding_value=0.0), + ct.float32, + ) + b_a33 = ct.where(o_i < i - 3 * BC, b_a33, ct.zeros((BC,), dtype=ct.float32)) + b_a33 = b_a33 + ct.sum(b_a33[:, None] * b_Ai33, axis=0) + b_Ai33 = ct.where((o_i == i - 3 * BC)[:, None], b_a33, b_Ai33) - b_w = safe_matmul(b_A, ct.astype(b_kb, b_k.dtype)) - wrows = ((bos + t_off) * HV + i_hv)[:, None] - w_offset = ct.broadcast_to(wrows, (BT, BK)) * K + bk_col - ct.scatter( - w_flat, - w_offset, - ct.astype(b_w, w_flat.dtype), - mask=m_k, - check_bounds=False, - ) + b_Ai00 = b_Ai00 + ct.astype(m_I, ct.float32) + b_Ai11 = b_Ai11 + ct.astype(m_I, ct.float32) + if NC >= 3: + b_Ai22 = b_Ai22 + ct.astype(m_I, ct.float32) + if NC >= 4: + b_Ai33 = b_Ai33 + ct.astype(m_I, ct.float32) + + # ---- merged inverse using off-diagonals (tf32) ------------------------------- + b_Ai10 = -safe_matmul(safe_matmul(b_Ai11, b_Akk10), b_Ai00) + b_Ai20 = z + b_Ai21 = z + b_Ai30 = z + b_Ai31 = z + b_Ai32 = z + if NC >= 3: + b_Ai21 = -safe_matmul(safe_matmul(b_Ai22, b_Akk21), b_Ai11) + b_Ai20 = -safe_matmul(b_Ai22, safe_matmul(b_Akk20, b_Ai00) + safe_matmul(b_Akk21, b_Ai10)) + if NC >= 4: + b_Ai32 = -safe_matmul(safe_matmul(b_Ai33, b_Akk32), b_Ai22) + b_Ai31 = -safe_matmul(b_Ai33, safe_matmul(b_Akk31, b_Ai11) + safe_matmul(b_Akk32, b_Ai21)) + b_Ai30 = -safe_matmul(b_Ai33, safe_matmul(b_Akk30, b_Ai00) + safe_matmul(b_Akk31, b_Ai10) + safe_matmul(b_Akk32, b_Ai20)) + + # ---- store full Akk_inv to Akk ----------------------------------------------- + store_bc_bc(Akk, Akk_base, i_tc0, 0, HV * BT, b_Ai00, BC, BT, T_eff) + store_bc_bc(Akk, Akk_base, i_tc1, 0, HV * BT, b_Ai10, BC, BT, T_eff) + store_bc_bc(Akk, Akk_base, i_tc1, BC, HV * BT, b_Ai11, BC, BT, T_eff) + if NC >= 3: + store_bc_bc(Akk, Akk_base, i_tc2, 0, HV * BT, b_Ai20, BC, BT, T_eff) + store_bc_bc(Akk, Akk_base, i_tc2, BC, HV * BT, b_Ai21, BC, BT, T_eff) + store_bc_bc(Akk, Akk_base, i_tc2, 2 * BC, HV * BT, b_Ai22, BC, BT, T_eff) + if NC >= 4: + store_bc_bc(Akk, Akk_base, i_tc3, 0, HV * BT, b_Ai30, BC, BT, T_eff) + store_bc_bc(Akk, Akk_base, i_tc3, BC, HV * BT, b_Ai31, BC, BT, T_eff) + store_bc_bc(Akk, Akk_base, i_tc3, 2 * BC, HV * BT, b_Ai32, BC, BT, T_eff) + store_bc_bc(Akk, Akk_base, i_tc3, 3 * BC, HV * BT, b_Ai33, BC, BT, T_eff) -# =========================================================================== -# chunk_kda_fwd_kernel_intra_token_parallel -# =========================================================================== -# Token-parallel: one block per (token, head-group). BH heads at a time, each -# head contributes a scalar Aqk/Akk per partner token j. GATHER RATIONALE: -# per-token head-interleaved scalar reads/writes -> structured gather/scatter. +# --- Kernels: state scan -------------------------------------------------------------------------- @ct.kernel -def chunk_kda_fwd_kernel_intra_token_parallel( - q, +def chunk_gated_delta_rule_fwd_kernel_h_blockdim64( k, + v, + w, + v_new, g, - beta, - Aqk, - Akk, - scale, + gk, + h, + h0, + ht, cu_seqlens, - N, - T, + chunk_offsets, H: ConstInt, HV: ConstInt, K: ConstInt, + V: ConstInt, BT: ConstInt, - BC: ConstInt, - BH: ConstInt, - BK: ConstInt, + BV: ConstInt, + BK: ConstInt, # next_pow2(K) -- full-width state/K tile (general for any K) + USE_G: ConstInt, + USE_GK: ConstInt, + USE_INITIAL_STATE: ConstInt, + STORE_FINAL_STATE: ConstInt, + SAVE_NEW_VALUE: ConstInt, + STATE_V_FIRST: ConstInt, ): - # BK = next_power_of_2(K) is passed by host - # (cuTile tile shapes must be compile-time constants). - i_tg = ct.bid(0) - i_hg = ct.bid(1) + # The KV state is carried as a SINGLE full-width tile ((BV, BK) when + # STATE_V_FIRST else (BK, BV), BK=next_pow2(K)), so the kernel is general + # for ANY K (no K<=256 cap). Every K-axis load zero-pads cols [K:BK] and + # every K-axis store masks them off, so the tail is 0 on load, contributes + # 0 to every MMA, and stays 0 in the state; a padded Gk tail loads as 0 so + # exp2(0)=1 leaves those zero cols unchanged. + i_v = ct.bid(0) + i_nh = ct.bid(1) + i_n = i_nh // HV + i_h = i_nh % HV - bos = (i_tg // T) * T - i_t = i_tg % T - T_eff = T - left = 0 - right = N - # Unrolled binary search to find i_n s.t. cu[i_n] <= i_tg < cu[i_n+1] - for _ in range(20): - if left < right: - mid = (left + right) // 2 - cmid = ct.load(cu_seqlens, (mid + 1,), shape=()).item() - if i_tg < cmid: - right = mid - else: - left = mid + 1 - i_n = left bos = ct.load(cu_seqlens, (i_n,), shape=()).item() eos = ct.load(cu_seqlens, (i_n + 1,), shape=()).item() - T_eff = eos - bos - i_t = i_tg - bos + T = eos - bos + NT = ct.cdiv(T, BT) + boh = ct.load(chunk_offsets, (i_n,), shape=()).item() + if STATE_V_FIRST: + b_h = ct.zeros((BV, BK), dtype=ct.float32) + else: + b_h = ct.zeros((BK, BV), dtype=ct.float32) - if i_t >= T_eff: - return + h_base = (boh * HV + i_h) * K * V + v_base = (bos * HV + i_h) * V + k_base = (bos * H + i_h // (HV // H)) * K + w_base = (bos * HV + i_h) * K + vnew_base = (bos * HV + i_h) * V + h0_base = i_nh * K * V + ht_base = i_nh * K * V - i_c = i_t // BT - i_s = (i_t % BT) // BC - i_tc = i_c * BT - i_ts = i_tc + i_s * BC + o_bk = ct.arange(BK, dtype=ct.int32) + o_bt = ct.arange(BT, dtype=ct.int32) + o_bv = ct.arange(BV, dtype=ct.int32) - G = HV // H + # K-tile / V-tile validity masks. The kernel carries a single BK-wide K tile + # (BK=next_pow2(K)) and BV-wide V tiles; when K or V is not a multiple of the + # tile width the extra lanes alias the neighbouring head's H slot. The matmul + # state rows are zeroed via K-masked K/Gk loads, but the raw H/H0/Ht + # gather/scatter are also masked here -> no cross-head corruption when + # K % BK != 0 (i.e. K < BK) or V % BV != 0. + mkh = o_bk < K + mvh = (i_v * BV + o_bv) < V - Aqk_base = bos * HV * BT - Akk_base = bos * HV * BC - beta_base = bos * HV + if USE_INITIAL_STATE: + if STATE_V_FIRST: + row = (i_v * BV + o_bv)[:, None] + b_h = b_h + ct.astype( + ct.gather( + h0, + h0_base + row * K + o_bk[None, :], + mask=mvh[:, None] & mkh[None, :], + check_bounds=True, + padding_value=0.0, + ), + ct.float32, + ) + else: + col = (i_v * BV + o_bv)[None, :] + b_h = b_h + ct.astype( + ct.gather( + h0, + h0_base + o_bk[:, None] * V + col, + mask=mkh[:, None] & mvh[None, :], + check_bounds=True, + padding_value=0.0, + ), + ct.float32, + ) - # cuTile gather/scatter require the index tuple rank to match the array rank - # (no raw pointer arithmetic). Arrays arrive pre-flattened from the - # host: Q/K -> (B*T*H, K), G -> (B*T*HV, K), Beta/Aqk/Akk -> 1-D. - o_hv = i_hg * BH + ct.arange(BH, dtype=ct.int32) - o_h = o_hv // G - o_k = ct.arange(BK, dtype=ct.int32) - m_hv = o_hv < HV - m_k = o_k < K - m_hk = m_hv[:, None] & m_k[None, :] + for i_t in range(NT): + h_chunk = h_base + i_t * HV * K * V - col = ct.broadcast_to(o_k[None, :], (BH, BK)) + if STATE_V_FIRST: + row = (i_v * BV + o_bv)[:, None] + ct.scatter( + h, + h_chunk + row * K + o_bk[None, :], + ct.astype(b_h, h.dtype), + mask=mvh[:, None] & mkh[None, :], + check_bounds=True, + ) + else: + col = (i_v * BV + o_bv)[None, :] + ct.scatter( + h, + h_chunk + o_bk[:, None] * V + col, + ct.astype(b_h, h.dtype), + mask=mkh[:, None] & mvh[None, :], + check_bounds=True, + ) - # Q/K: row = (bos + token) * H + head; col = key - qk_row = ct.broadcast_to(((bos + i_t) * H + o_h)[:, None], (BH, BK)) - b_q = ct.astype(ct.gather(q, (qk_row, col), mask=m_hk, check_bounds=False, padding_value=0.0), ct.float32) - b_k = ct.astype(ct.gather(k, (qk_row, col), mask=m_hk, check_bounds=False, padding_value=0.0), ct.float32) + w_row = (i_t * BT + o_bt)[:, None] + wmask_r = (i_t * BT + o_bt) < T + # Full-width K: padded cols [K:BK] gather OOB (masked to 0) and row-mask + # zeros the partial-chunk tail, so both contribute 0 to the MMA. + bw = ct.gather(w, w_base + w_row * (HV * K) + o_bk[None, :], mask=mkh[None, :], check_bounds=True, padding_value=0.0) + bw = ct.where(wmask_r[:, None], bw, ct.zeros((BT, BK), dtype=bw.dtype)) + b_v = ct.zeros((BT, BV), dtype=ct.float32) + bmat = ct.astype(ct.transpose(b_h), bw.dtype) if STATE_V_FIRST else ct.astype(b_h, bw.dtype) + b_v = safe_mma(bw, bmat, b_v) - # G: row = (bos + token) * HV + head; Beta: idx = (bos + token) * HV + head - g_row = ct.broadcast_to(((bos + i_t) * HV + o_hv)[:, None], (BH, BK)) - b_g = ct.astype(ct.gather(g, (g_row, col), mask=m_hk, check_bounds=False, padding_value=0.0), ct.float32) - b_beta = ct.astype(ct.gather(beta, beta_base + i_t * HV + o_hv, mask=m_hv, check_bounds=False, padding_value=0.0), ct.float32) - b_k = b_k * b_beta[:, None] + v_col = (i_v * BV + o_bv)[None, :] + vmask_c = (i_v * BV + o_bv) < V + v_off = v_base + (i_t * BT + o_bt)[:, None] * (HV * V) + v_col + v_full = (wmask_r[:, None]) & (vmask_c[None, :]) + b_v_load = ct.gather(v, v_off, check_bounds=True, padding_value=0.0) + b_v_load = ct.where(v_full, b_v_load, ct.zeros((BT, BV), dtype=b_v_load.dtype)) + b_v = ct.astype(b_v_load, ct.float32) - b_v - # Counted loop over the (static) BC-wide sub-chunk window. A runtime-bounded - # `for j in range(i_ts, j_hi)` lowers to per-iteration branches that tileiras - # cannot unroll/predicate; a counted `for jj in range(BC)` with a runtime - # guard derives `j` from `jj` and stays fully unrolled. - for jj in range(BC): - j = i_ts + jj - if j < i_t + 1 and j < T_eff and j < i_ts + BC: - kj_row = ct.broadcast_to(((bos + j) * H + o_h)[:, None], (BH, BK)) - gj_row = ct.broadcast_to(((bos + j) * HV + o_hv)[:, None], (BH, BK)) - b_kj = ct.astype(ct.gather(k, (kj_row, col), mask=m_hk, check_bounds=False, padding_value=0.0), ct.float32) - b_gj = ct.astype(ct.gather(g, (gj_row, col), mask=m_hk, check_bounds=False, padding_value=0.0), ct.float32) + if SAVE_NEW_VALUE: + vn_off = vnew_base + (i_t * BT + o_bt)[:, None] * (HV * V) + v_col + vn_oob = ct.full((BT, BV), v.shape[0], dtype=ct.int32) + vn_off = ct.where(v_full, vn_off, vn_oob) + ct.scatter(v_new, vn_off, ct.astype(b_v, v_new.dtype), check_bounds=True) + + last_idx = ct_min((i_t + 1) * BT, T) - 1 + + if USE_G: + m_t = (i_t * BT + o_bt) < T + b_g_last = ct.astype(ct.gather(g, (bos * HV + last_idx * HV + i_h,), check_bounds=True, padding_value=0.0), ct.float32) + g_off = (bos * HV + i_h) + (i_t * BT + o_bt) * HV + b_g = ct.astype(ct.gather(g, g_off, mask=m_t, check_bounds=True, padding_value=0.0), ct.float32) + decay = ct.where(m_t, exp2(b_g_last - b_g), ct.zeros((BT,), dtype=ct.float32)) + b_v = b_v * decay[:, None] + b_g_last = exp2(b_g_last) + b_h = b_h * b_g_last + + if USE_GK: + # Padded tail [K:BK] loads as 0 -> exp2(0)=1 leaves zero state cols unchanged. + gk_base = (bos + last_idx) * HV * K + i_h * K + b_gk = ct.astype(ct.gather(gk, gk_base + o_bk, mask=o_bk < K, check_bounds=True, padding_value=0.0), ct.float32) + if STATE_V_FIRST: + b_h = b_h * exp2(b_gk)[None, :] + else: + b_h = b_h * exp2(b_gk)[:, None] - b_kgj = ct.where(m_k[None, :], b_kj * exp2(b_g - b_gj), ct.zeros((BH, BK), dtype=ct.float32)) - b_Aqk = ct.sum(b_q * b_kgj, axis=1) * scale - b_Akk = ct.sum(b_k * b_kgj, axis=1) * (1.0 if j < i_t else 0.0) + b_v = ct.astype(b_v, k.dtype) + + k_row = o_bk[:, None] + k_time = (i_t * BT + o_bt)[None, :] + kmask = (o_bk[:, None] < K) & ((i_t * BT + o_bt)[None, :] < T) + bk = ct.gather(k, k_base + k_row * 1 + k_time * (H * K), check_bounds=True, padding_value=0.0) + bk = ct.where(kmask, bk, ct.zeros((BK, BT), dtype=bk.dtype)) + prod = safe_mma(bk, b_v, ct.zeros((BK, BV), dtype=ct.float32)) + if STATE_V_FIRST: + b_h = b_h + ct.transpose(prod) + else: + b_h = b_h + prod + if STORE_FINAL_STATE: + if STATE_V_FIRST: + row = (i_v * BV + o_bv)[:, None] ct.scatter( - Aqk, - Aqk_base + i_t * HV * BT + o_hv * BT + (j % BT), - ct.astype(b_Aqk, Aqk.dtype), - mask=m_hv, - check_bounds=False, + ht, + ht_base + row * K + o_bk[None, :], + ct.astype(b_h, ht.dtype), + mask=mvh[:, None] & mkh[None, :], + check_bounds=True, ) + else: + col = (i_v * BV + o_bv)[None, :] ct.scatter( - Akk, - Akk_base + i_t * HV * BC + o_hv * BC + (j - i_ts), - ct.astype(b_Akk, Akk.dtype), - mask=m_hv, - check_bounds=False, + ht, + ht_base + o_bk[:, None] * V + col, + ct.astype(b_h, ht.dtype), + mask=mkh[:, None] & mvh[None, :], + check_bounds=True, ) -# =========================================================================== -# chunk_intra (diag compute+solve / intra sub-chunk) -# =========================================================================== -# One BC sub-chunk per block: compute the diagonal Aqk/Akk blocks then do an -# in-place forward-substitution triangular solve. GATHER RATIONALE: head- -# interleaved Q/K/G sub-views + a transposed b_kgt. - - @ct.kernel -def chunk_kda_fwd_kernel_inter_diag_compute_solve( +def chunk_gated_delta_rule_bwd_kernel_dhu_blockdim64( q, k, + w, g, - beta, - Aqk, - Akk, - scale, + gk, + dstate_in, + dh0, + do, + dh, + dv, + dv2, cu_seqlens, - chunk_indices, + chunk_offsets, + scale, H: ConstInt, HV: ConstInt, K: ConstInt, + V: ConstInt, BT: ConstInt, - BC: ConstInt, - BK: ConstInt, + BV: ConstInt, + BK: ConstInt, # next_pow2(K) -- full-width state/K tile (general for any K) + USE_G: ConstInt, + USE_GK: ConstInt, + USE_INITIAL_STATE: ConstInt, + USE_FINAL_STATE_GRADIENT: ConstInt, + STATE_V_FIRST: ConstInt, ): - i_t = ct.bid(0) - i_i = ct.bid(1) - i_hv = ct.bid(2) - i_h = i_hv // (HV // H) - - i_n = ct.load(chunk_indices, (i_t * 2,), shape=()).item() - i_t = ct.load(chunk_indices, (i_t * 2 + 1,), shape=()).item() + # dH state is a SINGLE full-width tile ((BV, BK) when STATE_V_FIRST else + # (BK, BV), BK=next_pow2(K)), so the kernel is general for ANY K (no K<=256 + # cap). Every flat gather/scatter over the K axis uses oK=arange(BK) with a + # `(oK < K)` mask, so rows/cols [K:BK] are zero on load, contribute 0 to + # every MMA, and are masked out on store; a Gk padding tail loads as 0 so + # exp2(0)=1 leaves those zero state cols unchanged. + i_v = ct.bid(0) + i_nh = ct.bid(1) + i_n = i_nh // HV + i_h = i_nh % HV bos = ct.load(cu_seqlens, (i_n,), shape=()).item() eos = ct.load(cu_seqlens, (i_n + 1,), shape=()).item() - T_eff = eos - bos + T = eos - bos + NT = ct.cdiv(T, BT) + boh = ct.load(chunk_offsets, (i_n,), shape=()).item() + if STATE_V_FIRST: + b_dh = ct.zeros((BV, BK), dtype=ct.float32) + else: + b_dh = ct.zeros((BK, BV), dtype=ct.float32) - i_ti = i_t * BT + i_i * BC - if i_ti >= T_eff: - return + q_base = (bos * H + i_h // (HV // H)) * K + k_base = (bos * H + i_h // (HV // H)) * K + w_base = (bos * HV + i_h) * K + do_base = (bos * HV + i_h) * V + dv_base = (bos * HV + i_h) * V + dv2_base = (bos * HV + i_h) * V + dh_base = (boh * HV + i_h) * K * V + gk_base = (bos * HV + i_h) * K + dh0_base = i_nh * K * V + dht_base = i_nh * K * V - o_bc = ct.arange(BC, dtype=ct.int32) o_bk = ct.arange(BK, dtype=ct.int32) - o_c = i_ti + o_bc - m_c = o_c < T_eff - m_k = o_bk < K - - q_base = (bos * H + i_h) * K - k_base = (bos * H + i_h) * K - g_base = (bos * HV + i_hv) * K - beta_base = bos * HV + i_hv - Aqk_base = (bos * HV + i_hv) * BT - Akk_base = (bos * HV + i_hv) * BC - - qk_rows = o_c[:, None] - qk_cols = o_bk[None, :] - m_qk = m_c[:, None] & m_k[None, :] - q_off = q_base + qk_rows * (H * K) + qk_cols - k_off = k_base + qk_rows * (H * K) + qk_cols - g_off = g_base + qk_rows * (HV * K) + qk_cols - b_q = ct.gather(q, q_off, mask=m_qk, check_bounds=False, padding_value=0.0) - b_k = ct.gather(k, k_off, mask=m_qk, check_bounds=False, padding_value=0.0) - b_g = ct.gather(g, g_off, mask=m_qk, check_bounds=False, padding_value=0.0) - b_beta = ct.gather(beta, beta_base + o_c * HV, mask=m_c, check_bounds=False, padding_value=0.0) + o_bt = ct.arange(BT, dtype=ct.int32) + o_bv = ct.arange(BV, dtype=ct.int32) - gn_row = i_ti + ct_min(BC // 2, T_eff - i_ti - 1) - b_gn = ct.gather(g, g_base + gn_row * (HV * K) + o_bk, mask=m_k, check_bounds=False, padding_value=0.0) - b_gn = ct.astype(b_gn, ct.float32)[None, :] + # V-boundary mask for the state-gradient (dH/dH0) scatters: a partial + # trailing V tile (V not a multiple of BV) must not write rows/cols >= V, + # else the store spills into the neighbouring chunk/head's state-gradient + # (boundary-checked store semantics). + m_v = (i_v * BV + o_bv) < V - b_gm = ct.astype(b_g, ct.float32) - b_gn - b_gq = ct.where(m_c[:, None], exp2(b_gm), ct.zeros((BC, BK), dtype=ct.float32)) - b_gk = ct.where(m_c[:, None], exp2(-b_gm), ct.zeros((BC, BK), dtype=ct.float32)) + # --- Load final state gradient dHt -> b_dh (single full-width tile; [K:BK] loads 0) --- + if USE_FINAL_STATE_GRADIENT: + if STATE_V_FIRST: + row = (i_v * BV + o_bv)[:, None] + b_dh = b_dh + ct.gather( + dstate_in, + dht_base + row * K + o_bk[None, :], + mask=(o_bk < K)[None, :], + check_bounds=True, + padding_value=0.0, + ) + else: + col = (i_v * BV + o_bv)[None, :] + b_dh = b_dh + ct.gather( + dstate_in, + dht_base + o_bk[:, None] * V + col, + mask=(o_bk < K)[:, None], + check_bounds=True, + padding_value=0.0, + ) - if K < 256: - b_kgt = ct.transpose(ct.astype(ct.astype(b_k, ct.float32) * b_gk, b_k.dtype)) - b_qg = ct.astype(ct.astype(b_q, ct.float32) * b_gq, b_q.dtype) - b_kg = ct.astype(ct.astype(b_k, ct.float32) * b_gq, b_k.dtype) - else: - b_kgt = ct.transpose(ct.astype(b_k, ct.float32) * b_gk) - b_qg = ct.astype(b_q, ct.float32) * b_gq - b_kg = ct.astype(b_k, ct.float32) * b_gq + # cuTile range() requires a positive step; iterate forward and reverse the + # index to preserve the backward (NT-1 .. 0) chunk traversal. + for _i_t in range(NT): + i_t = NT - 1 - _i_t + dh_chunk = dh_base + i_t * HV * K * V - b_Aqk = safe_matmul(b_qg, b_kgt) * scale - b_Akk = safe_matmul(b_kg, b_kgt) * ct.astype(b_beta, ct.float32)[:, None] + # Store current b_dh to dH[i_t] (single full-width tile; [K:BK] masked out) + if STATE_V_FIRST: + row = (i_v * BV + o_bv)[:, None] + ct.scatter( + dh, + dh_chunk + row * K + o_bk[None, :], + ct.astype(b_dh, dh.dtype), + mask=m_v[:, None] & (o_bk < K)[None, :], + check_bounds=False, + ) + else: + col = (i_v * BV + o_bv)[None, :] + ct.scatter( + dh, + dh_chunk + o_bk[:, None] * V + col, + ct.astype(b_dh, dh.dtype), + mask=(o_bk < K)[:, None] & m_v[None, :], + check_bounds=False, + ) - o_i = o_bc - m_Aqk = o_i[:, None] >= o_i[None, :] - m_Akk = o_i[:, None] > o_i[None, :] - m_I = o_i[:, None] == o_i[None, :] + last_idx = ct_min((i_t + 1) * BT, T) - 1 - b_Aqk = ct.where(m_Aqk, b_Aqk, ct.zeros((BC, BC), dtype=ct.float32)) - b_Akk = ct.where(m_Akk, b_Akk, ct.zeros((BC, BC), dtype=ct.float32)) + m_t = (i_t * BT + o_bt) < T + if USE_G: + bg_last = ct.astype(ct.gather(g, ((bos + last_idx) * HV + i_h,), check_bounds=True, padding_value=0.0), ct.float32) + g_off = (bos * HV + i_h) + (i_t * BT + o_bt) * HV + b_g = ct.astype(ct.gather(g, g_off, mask=m_t, check_bounds=True, padding_value=0.0), ct.float32) + bg_last_exp = exp2(bg_last) + b_g_exp = exp2(b_g) + else: + bg_last_exp = ct.astype(0.0, ct.float32) + b_g_exp = ct.zeros((BT,), dtype=ct.float32) - # store Aqk (Akk for this kernel writes to the fp32 diagonal buffer) - aqk_rows = o_c[:, None] - aqk_cols = (i_i * BC + o_bc)[None, :] - aqk_off = Aqk_base + aqk_rows * (HV * BT) + aqk_cols - m_aqk = m_c[:, None] & ((i_i * BC + o_bc) < BT)[None, :] - ct.scatter(Aqk, aqk_off, ct.astype(b_Aqk, Aqk.dtype), mask=m_aqk, check_bounds=False) + # dO, dV, dV2 tiles + v_col = (i_v * BV + o_bv)[None, :] + vmask_c = (i_v * BV + o_bv) < V + v_full = (m_t[:, None]) & (vmask_c[None, :]) + do_off = do_base + (i_t * BT + o_bt)[:, None] * (HV * V) + v_col + b_do = ct.gather(do, do_off, check_bounds=True, padding_value=0.0) + b_do = ct.where(v_full, b_do, ct.zeros((BT, BV), dtype=b_do.dtype)) - # diagonal Akk -> inverse via Neumann series by squaring, written to Akk(diag buf). - # - # Akk is strictly-lower-triangular and nilpotent (Akk^BC = 0); with - # N := -Akk, (I + Akk)^-1 = (I - N)^-1 = sum_{k=0}^{BC-1} N^k - # = prod_{j=0}^{log2(BC)-1} (I + N^(2^j)) -- log2(BC) block matmuls, no - # serial row dependency. Rows beyond T_eff have N = 0 (b_gq/b_gk were - # zeroed by m_c) and converge to the identity. - # - # Precision: the squarings run at tf32 via safe_matmul (fp32 operands - # would fall back to SIMT even at M=32). - b_N = -b_Akk # N := -Akk, so (I + Akk)^-1 = (I - N)^-1 = sum_k N^k - b_Ai = ct.astype(m_I, ct.float32) + b_N # (I + N) - # Squaring stages: `range(2, BC)` traces as a compile-time-bounded loop with - # concrete Python `i` during unroll (BC is a ConstInt). We do one squaring - # each time `i` is a power of two (i = 2, 4, 8, 16 for BC=32), giving exactly - # ceil(log2(BC)) - 1 stages (factors N^2, N^4, N^8, N^16). All the branch - # decisions are host-side (Python int `i`), so nothing lowers to a device - # branch -- the body is fully unrolled into log2(BC) block matmuls. - for i in range(2, BC): - if (i & (i - 1)) == 0: # i is a power of two (host-time test) - b_N = safe_matmul(b_N, b_N) # N -> N^2 -> N^4 ... - b_Ai = b_Ai + safe_matmul(b_Ai, b_N) # (prod so far) @ (I + N^(2^j)) + # b_dv = b_k @ b_dh (single full-width K MMA). b_k tile is (BT, BK); [K:BK] + # is masked to 0 so it contributes nothing. + bk = ct.gather( + k, + k_base + (i_t * BT + o_bt)[:, None] * (H * K) + o_bk[None, :], + mask=(o_bk < K)[None, :], + check_bounds=True, + padding_value=0.0, + ) + bk = ct.where(m_t[:, None], bk, ct.zeros((BT, BK), dtype=bk.dtype)) + if USE_GK: + # Gk offset: base (bos*HV+i_h)*K ; then + last_idx*HV*K + o_k. + # Padded tail [K:BK] loads as 0 -> exp2(0)=1 leaves those zero cols unchanged. + gkl = (bos * HV + i_h) * K + last_idx * HV * K + b_gk_last = ct.astype(ct.gather(gk, gkl + o_bk, mask=o_bk < K, check_bounds=True, padding_value=0.0), ct.float32) + bmat = ct.astype(ct.transpose(b_dh), bk.dtype) if STATE_V_FIRST else ct.astype(b_dh, bk.dtype) + b_dv = safe_mma(bk, bmat, ct.zeros((BT, BV), dtype=ct.float32)) - akk_rows = o_c[:, None] - akk_cols = o_bc[None, :] - akk_off = Akk_base + akk_rows * (HV * BC) + akk_cols - m_akk = m_c[:, None] & (o_bc < BC)[None, :] - ct.scatter(Akk, akk_off, ct.astype(b_Ai, Akk.dtype), mask=m_akk, check_bounds=False) + if USE_G: + decay = ct.where(m_t, exp2(bg_last - b_g), ct.zeros((BT,), dtype=ct.float32)) + b_dv = b_dv * decay[:, None] + # b_dv += dV ; store to dV2 + dv_off = dv_base + (i_t * BT + o_bt)[:, None] * (HV * V) + v_col + b_dv_load = ct.gather(dv, dv_off, check_bounds=True, padding_value=0.0) + b_dv_load = ct.where(v_full, b_dv_load, ct.zeros((BT, BV), dtype=b_dv_load.dtype)) + b_dv = b_dv + ct.astype(b_dv_load, ct.float32) + dv2_off = dv2_base + (i_t * BT + o_bt)[:, None] * (HV * V) + v_col + dv2_oob = ct.full((BT, BV), dv2.shape[0], dtype=ct.int32) + dv2_off = ct.where(v_full, dv2_off, dv2_oob) + ct.scatter(dv2, dv2_off, ct.astype(b_dv, dv2.dtype), check_bounds=True) -# --------------------------------------------------------------------------- -# chunk_kda_fwd_kernel_inter_solve_fused -- off-diagonal Akk + merged solve -# --------------------------------------------------------------------------- -def load_bc_bk(arr, base, row0, col0, stride_row, BC: ConstInt, BK: ConstInt, T_eff, K: ConstInt): - # Boundary-checked (BC,BK) block load at (row0,col0) of the (T,K) view - # (row stride stride_row), cast to fp32, on the flattened 1-D view. - o_r = ct.arange(BC, dtype=ct.int32) - o_c = ct.arange(BK, dtype=ct.int32) - rows = row0 + o_r - cols = col0 + o_c - off = base + rows[:, None] * stride_row + cols[None, :] - mask = (rows < T_eff)[:, None] & (cols < K)[None, :] - return ct.astype(ct.gather(arr, off, mask=mask, check_bounds=False, padding_value=0.0), ct.float32) + # b_dh += trans(b_q@b_do*scale - b_w@b_dv) (b_q,b_w are (BK,BT) transposed; + # rows [K:BK] are masked to 0 so contribute nothing to the update) + b_dv_c = ct.astype(b_dv, do.dtype) + time = (i_t * BT + o_bt)[None, :] + tmask = (i_t * BT + o_bt)[None, :] < T + kr = o_bk[:, None] + wmask = (o_bk[:, None] < K) & tmask + b_w = ct.gather(w, w_base + kr * 1 + time * (HV * K), check_bounds=True, padding_value=0.0) + b_w = ct.where(wmask, b_w, ct.zeros((BK, BT), dtype=b_w.dtype)) + b_q = ct.gather(q, q_base + kr * 1 + time * (H * K), check_bounds=True, padding_value=0.0) + b_q = ct.where(wmask, b_q, ct.zeros((BK, BT), dtype=b_q.dtype)) + if USE_G: + b_dh = b_dh * bg_last_exp + b_q = b_q * b_g_exp[None, :] + if USE_GK: + if STATE_V_FIRST: + b_dh = b_dh * exp2(b_gk_last)[None, :] + else: + b_dh = b_dh * exp2(b_gk_last[:, None]) + term = safe_matmul(b_q, b_do) * scale - safe_matmul(b_w, b_dv_c) + if STATE_V_FIRST: + b_dh = b_dh + ct.transpose(term) + else: + b_dh = b_dh + term + + if USE_INITIAL_STATE: + if STATE_V_FIRST: + row = (i_v * BV + o_bv)[:, None] + ct.scatter( + dh0, + dh0_base + row * K + o_bk[None, :], + ct.astype(b_dh, dh0.dtype), + mask=m_v[:, None] & (o_bk < K)[None, :], + check_bounds=False, + ) + else: + col = (i_v * BV + o_bv)[None, :] + ct.scatter( + dh0, + dh0_base + o_bk[:, None] * V + col, + ct.astype(b_dh, dh0.dtype), + mask=(o_bk < K)[:, None] & m_v[None, :], + check_bounds=False, + ) -def store_bc_bc(arr, base, row0, col0, stride_row, blk, BC: ConstInt, BT_or_BC: ConstInt, T_eff): - o_r = ct.arange(BC, dtype=ct.int32) - o_c = ct.arange(BC, dtype=ct.int32) - rows = row0 + o_r - cols = col0 + o_c - off = base + rows[:, None] * stride_row + cols[None, :] - mask = (rows < T_eff)[:, None] & (cols < BT_or_BC)[None, :] - ct.scatter(arr, off, ct.astype(blk, arr.dtype), mask=mask, check_bounds=False) +# --- Kernels: attention and gradients ------------------------------------------------------------- @ct.kernel -def chunk_kda_fwd_kernel_inter_solve_fused( +def chunk_gla_fwd_kernel_o( q, - k, + v, g, - beta, - Aqk, - Akkd, - Akk, - scale, + h, + o, + A, cu_seqlens, chunk_indices, + scale, H: ConstInt, HV: ConstInt, K: ConstInt, + V: ConstInt, BT: ConstInt, - BC: ConstInt, - NC: ConstInt, BK: ConstInt, - USE_SAFE_GATE: ConstInt, + BV: ConstInt, + STATE_V_FIRST: ConstInt, ): - i_t = ct.bid(0) - i_hv = ct.bid(1) + i_v = ct.bid(0) + i_t = ct.bid(1) + i_hv = ct.bid(2) i_h = i_hv // (HV // H) + # grid dim-1 is the GLOBAL chunk index; H is laid out per global chunk + # (H-state kernel writes slot chunk_offsets[i_n] + local). Capture it + # before i_t is reassigned to the per-sequence (local) chunk index. + i_tg = i_t i_n = ct.load(chunk_indices, (i_t * 2,), shape=()).item() i_t = ct.load(chunk_indices, (i_t * 2 + 1,), shape=()).item() bos = ct.load(cu_seqlens, (i_n,), shape=()).item() eos = ct.load(cu_seqlens, (i_n + 1,), shape=()).item() - T_eff = eos - bos - - if i_t * BT >= T_eff: - return - - i_tc0 = i_t * BT - i_tc1 = i_t * BT + BC - i_tc2 = i_t * BT + 2 * BC - i_tc3 = i_t * BT + 3 * BC + T = eos - bos + NT = ct.cdiv(T, BT) + o_bt = ct.arange(BT, dtype=ct.int32) + o_bk = ct.arange(BK, dtype=ct.int32) + o_bv = ct.arange(BV, dtype=ct.int32) + m_s = o_bt[:, None] >= o_bt[None, :] q_base = (bos * H + i_h) * K - k_base = (bos * H + i_h) * K g_base = (bos * HV + i_hv) * K - Aqk_base = (bos * HV + i_hv) * BT - Akk_base = (bos * HV + i_hv) * BT - Akkd_base = (bos * HV + i_hv) * BC - beta_base = bos * HV + i_hv - - o_i = ct.arange(BC, dtype=ct.int32) - o_k = ct.arange(BK, dtype=ct.int32) - m_tc1 = (i_tc1 + o_i) < T_eff - m_tc2 = (i_tc2 + o_i) < T_eff - m_tc3 = (i_tc3 + o_i) < T_eff - - z = ct.zeros((BC, BC), dtype=ct.float32) - b_Aqk10 = z - b_Akk10 = z - b_Aqk20 = z - b_Akk20 = z - b_Aqk21 = z - b_Akk21 = z - b_Aqk30 = z - b_Akk30 = z - b_Aqk31 = z - b_Akk31 = z - b_Aqk32 = z - b_Akk32 = z - - # ---- off-diagonal blocks ----------------------------------------------------- - num_k = (K + BK - 1) // BK - for i_k in range(num_k): - kk = i_k * BK + o_k - m_k = kk < K - b_k0 = load_bc_bk(k, k_base, i_tc0, i_k * BK, H * K, BC, BK, T_eff, K) - b_g0 = load_bc_bk(g, g_base, i_tc0, i_k * BK, HV * K, BC, BK, T_eff, K) - - if i_tc1 < T_eff: - b_q1 = load_bc_bk(q, q_base, i_tc1, i_k * BK, H * K, BC, BK, T_eff, K) - b_k1 = load_bc_bk(k, k_base, i_tc1, i_k * BK, H * K, BC, BK, T_eff, K) - b_g1 = load_bc_bk(g, g_base, i_tc1, i_k * BK, HV * K, BC, BK, T_eff, K) - b_gn1 = ct.astype( - ct.gather(g, g_base + i_tc1 * (HV * K) + kk, mask=m_k, check_bounds=False, padding_value=0.0), - ct.float32, - ) - b_gqn = ct.where(m_tc1[:, None], exp2(b_g1 - b_gn1[None, :]), ct.zeros((BC, BK), dtype=ct.float32)) - b_kgt = ct.transpose(b_k0 * exp2(b_gn1[None, :] - b_g0)) - b_Aqk10 = safe_mma(b_q1 * b_gqn, b_kgt, b_Aqk10) - b_Akk10 = safe_mma(b_k1 * b_gqn, b_kgt, b_Akk10) - - if NC >= 3 and i_tc2 < T_eff: - b_q2 = load_bc_bk(q, q_base, i_tc2, i_k * BK, H * K, BC, BK, T_eff, K) - b_k2 = load_bc_bk(k, k_base, i_tc2, i_k * BK, H * K, BC, BK, T_eff, K) - b_g2 = load_bc_bk(g, g_base, i_tc2, i_k * BK, HV * K, BC, BK, T_eff, K) - b_gn2 = ct.astype( - ct.gather(g, g_base + i_tc2 * (HV * K) + kk, mask=m_k, check_bounds=False, padding_value=0.0), - ct.float32, - ) - b_gqn2 = ct.where(m_tc2[:, None], exp2(b_g2 - b_gn2[None, :]), ct.zeros((BC, BK), dtype=ct.float32)) - b_qg2 = b_q2 * b_gqn2 - b_kg2 = b_k2 * b_gqn2 - b_kgt = ct.transpose(b_k0 * exp2(b_gn2[None, :] - b_g0)) - b_Aqk20 = safe_mma(b_qg2, b_kgt, b_Aqk20) - b_Akk20 = safe_mma(b_kg2, b_kgt, b_Akk20) - b_kgt = ct.transpose(b_k1 * exp2(b_gn2[None, :] - b_g1)) - b_Aqk21 = safe_mma(b_qg2, b_kgt, b_Aqk21) - b_Akk21 = safe_mma(b_kg2, b_kgt, b_Akk21) - - if NC >= 4 and i_tc3 < T_eff: - b_q3 = load_bc_bk(q, q_base, i_tc3, i_k * BK, H * K, BC, BK, T_eff, K) - b_k3 = load_bc_bk(k, k_base, i_tc3, i_k * BK, H * K, BC, BK, T_eff, K) - b_g3 = load_bc_bk(g, g_base, i_tc3, i_k * BK, HV * K, BC, BK, T_eff, K) - b_gn3 = ct.astype( - ct.gather(g, g_base + i_tc3 * (HV * K) + kk, mask=m_k, check_bounds=False, padding_value=0.0), - ct.float32, - ) - b_gqn3 = ct.where(m_tc3[:, None], exp2(b_g3 - b_gn3[None, :]), ct.zeros((BC, BK), dtype=ct.float32)) - b_qg3 = b_q3 * b_gqn3 - b_kg3 = b_k3 * b_gqn3 - b_kgt = ct.transpose(b_k0 * exp2(b_gn3[None, :] - b_g0)) - b_Aqk30 = safe_mma(b_qg3, b_kgt, b_Aqk30) - b_Akk30 = safe_mma(b_kg3, b_kgt, b_Akk30) - b_kgt = ct.transpose(b_k1 * exp2(b_gn3[None, :] - b_g1)) - b_Aqk31 = safe_mma(b_qg3, b_kgt, b_Aqk31) - b_Akk31 = safe_mma(b_kg3, b_kgt, b_Akk31) - b_kgt = ct.transpose(b_k2 * exp2(b_gn3[None, :] - b_g2)) - b_Aqk32 = safe_mma(b_qg3, b_kgt, b_Aqk32) - b_Akk32 = safe_mma(b_kg3, b_kgt, b_Akk32) + v_base = (bos * HV + i_hv) * V + o_base = (bos * HV + i_hv) * V + h_base = (i_tg * HV + i_hv) * K * V + A_base = (bos * HV + i_hv) * BT - # ---- save off-diagonal Aqk blocks, scale Akk by Beta ------------------------- - if i_tc1 < T_eff: - store_bc_bc(Aqk, Aqk_base, i_tc1, 0, HV * BT, b_Aqk10 * scale, BC, BT, T_eff) - b_b1 = ct.astype( - ct.gather(beta, beta_base + (i_tc1 + o_i) * HV, mask=m_tc1, check_bounds=False, padding_value=0.0), - ct.float32, - ) - b_Akk10 = b_Akk10 * b_b1[:, None] - if NC >= 3 and i_tc2 < T_eff: - store_bc_bc(Aqk, Aqk_base, i_tc2, 0, HV * BT, b_Aqk20 * scale, BC, BT, T_eff) - store_bc_bc(Aqk, Aqk_base, i_tc2, BC, HV * BT, b_Aqk21 * scale, BC, BT, T_eff) - b_b2 = ct.astype( - ct.gather(beta, beta_base + (i_tc2 + o_i) * HV, mask=m_tc2, check_bounds=False, padding_value=0.0), - ct.float32, - ) - b_Akk20 = b_Akk20 * b_b2[:, None] - b_Akk21 = b_Akk21 * b_b2[:, None] - if NC >= 4 and i_tc3 < T_eff: - store_bc_bc(Aqk, Aqk_base, i_tc3, 0, HV * BT, b_Aqk30 * scale, BC, BT, T_eff) - store_bc_bc(Aqk, Aqk_base, i_tc3, BC, HV * BT, b_Aqk31 * scale, BC, BT, T_eff) - store_bc_bc(Aqk, Aqk_base, i_tc3, 2 * BC, HV * BT, b_Aqk32 * scale, BC, BT, T_eff) - b_b3 = ct.astype( - ct.gather(beta, beta_base + (i_tc3 + o_i) * HV, mask=m_tc3, check_bounds=False, padding_value=0.0), - ct.float32, - ) - b_Akk30 = b_Akk30 * b_b3[:, None] - b_Akk31 = b_Akk31 * b_b3[:, None] - b_Akk32 = b_Akk32 * b_b3[:, None] + t_row = i_t * BT + o_bt + m_t = t_row < T - # ---- load diagonal inverse blocks from Akkd (fp32) --------------------------- - b_Ai00 = load_bc_bk(Akkd, Akkd_base, i_tc0, 0, HV * BC, BC, BC, T_eff, BC) - b_Ai11 = load_bc_bk(Akkd, Akkd_base, i_tc1, 0, HV * BC, BC, BC, T_eff, BC) - b_Ai22 = load_bc_bk(Akkd, Akkd_base, i_tc2, 0, HV * BC, BC, BC, T_eff, BC) if NC >= 3 else z - b_Ai33 = load_bc_bk(Akkd, Akkd_base, i_tc3, 0, HV * BC, BC, BC, T_eff, BC) if NC >= 4 else z + b_o = ct.zeros((BT, BV), dtype=ct.float32) + num_k = (K + BK - 1) // BK + for i_k in range(num_k): + k_col = i_k * BK + o_bk + m_k = k_col < K + v_col = i_v * BV + o_bv + m_v = v_col < V - # ---- forward substitution on diagonals (only when gate not pre-solved) ------- - if not USE_SAFE_GATE: - m_A = o_i[:, None] > o_i[None, :] - m_I = o_i[:, None] == o_i[None, :] - b_Ai00 = -ct.where(m_A, b_Ai00, z) - b_Ai11 = -ct.where(m_A, b_Ai11, z) - if NC >= 3: - b_Ai22 = -ct.where(m_A, b_Ai22, z) - if NC >= 4: - b_Ai33 = -ct.where(m_A, b_Ai33, z) + # b_h: STATE_V_FIRST -> view (V,K) block (BV,BK) at (i_v*BV, i_k*BK); + # else -> view (K,V) block (BK,BV) at (i_k*BK, i_v*BV). + if STATE_V_FIRST: + h_off = h_base + v_col[:, None] * K + k_col[None, :] + h_mask = m_v[:, None] & m_k[None, :] + b_h = ct.gather(h, h_off, mask=h_mask, check_bounds=True, padding_value=0.0) + else: + h_off = h_base + k_col[:, None] * V + v_col[None, :] + h_mask = m_k[:, None] & m_v[None, :] + b_h = ct.gather(h, h_off, mask=h_mask, check_bounds=True, padding_value=0.0) - # Counted loops over the static [BC] forward-substitution window with a - # runtime guard (i < T_eff - i_tc0). A runtime upper bound (ct_min(..., - # T_eff - i_tc0)) lowers to per-iteration branches tileiras can't - # unroll/predicate; the counted form stays fully unrolled. - for i in range(2, BC): - if i < T_eff - i_tc0: - b_a00 = -ct.astype( - ct.gather(Akkd, Akkd_base + (i_tc0 + i) * (HV * BC) + o_i, check_bounds=False, padding_value=0.0), - ct.float32, - ) - b_a00 = ct.where(o_i < i, b_a00, ct.zeros((BC,), dtype=ct.float32)) - b_a00 = b_a00 + ct.sum(b_a00[:, None] * b_Ai00, axis=0) - b_Ai00 = ct.where((o_i == i)[:, None], b_a00, b_Ai00) - for i in range(BC + 2, 2 * BC): - if i < T_eff - i_tc0: - b_a11 = -ct.astype( - ct.gather(Akkd, Akkd_base + (i_tc0 + i) * (HV * BC) + o_i, check_bounds=False, padding_value=0.0), - ct.float32, - ) - b_a11 = ct.where(o_i < i - BC, b_a11, ct.zeros((BC,), dtype=ct.float32)) - b_a11 = b_a11 + ct.sum(b_a11[:, None] * b_Ai11, axis=0) - b_Ai11 = ct.where((o_i == i - BC)[:, None], b_a11, b_Ai11) - if NC >= 3: - for i in range(2 * BC + 2, 3 * BC): - if i < T_eff - i_tc0: - b_a22 = -ct.astype( - ct.gather(Akkd, Akkd_base + (i_tc0 + i) * (HV * BC) + o_i, check_bounds=False, padding_value=0.0), - ct.float32, - ) - b_a22 = ct.where(o_i < i - 2 * BC, b_a22, ct.zeros((BC,), dtype=ct.float32)) - b_a22 = b_a22 + ct.sum(b_a22[:, None] * b_Ai22, axis=0) - b_Ai22 = ct.where((o_i == i - 2 * BC)[:, None], b_a22, b_Ai22) - if NC >= 4: - for i in range(3 * BC + 2, 4 * BC): - if i < T_eff - i_tc0: - b_a33 = -ct.astype( - ct.gather(Akkd, Akkd_base + (i_tc0 + i) * (HV * BC) + o_i, check_bounds=False, padding_value=0.0), - ct.float32, - ) - b_a33 = ct.where(o_i < i - 3 * BC, b_a33, ct.zeros((BC,), dtype=ct.float32)) - b_a33 = b_a33 + ct.sum(b_a33[:, None] * b_Ai33, axis=0) - b_Ai33 = ct.where((o_i == i - 3 * BC)[:, None], b_a33, b_Ai33) + q_off = q_base + t_row[:, None] * (H * K) + k_col[None, :] + g_off = g_base + t_row[:, None] * (HV * K) + k_col[None, :] + qg_mask = m_t[:, None] & m_k[None, :] + b_q = ct.gather(q, q_off, mask=qg_mask, check_bounds=True, padding_value=0.0) + b_g = ct.astype(ct.gather(g, g_off, mask=qg_mask, check_bounds=True, padding_value=0.0), ct.float32) + b_qg = ct.astype(b_q * exp2(b_g), b_q.dtype) - b_Ai00 = b_Ai00 + ct.astype(m_I, ct.float32) - b_Ai11 = b_Ai11 + ct.astype(m_I, ct.float32) - if NC >= 3: - b_Ai22 = b_Ai22 + ct.astype(m_I, ct.float32) - if NC >= 4: - b_Ai33 = b_Ai33 + ct.astype(m_I, ct.float32) + if STATE_V_FIRST: + b_o = safe_mma(b_qg, ct.astype(ct.transpose(b_h), b_qg.dtype), b_o) + else: + b_o = safe_mma(b_qg, ct.astype(b_h, b_qg.dtype), b_o) - # ---- merged inverse using off-diagonals (tf32) ------------------------------- - b_Ai10 = -safe_matmul(safe_matmul(b_Ai11, b_Akk10), b_Ai00) - b_Ai20 = z - b_Ai21 = z - b_Ai30 = z - b_Ai31 = z - b_Ai32 = z - if NC >= 3: - b_Ai21 = -safe_matmul(safe_matmul(b_Ai22, b_Akk21), b_Ai11) - b_Ai20 = -safe_matmul(b_Ai22, safe_matmul(b_Akk20, b_Ai00) + safe_matmul(b_Akk21, b_Ai10)) - if NC >= 4: - b_Ai32 = -safe_matmul(safe_matmul(b_Ai33, b_Akk32), b_Ai22) - b_Ai31 = -safe_matmul(b_Ai33, safe_matmul(b_Akk31, b_Ai11) + safe_matmul(b_Akk32, b_Ai21)) - b_Ai30 = -safe_matmul(b_Ai33, safe_matmul(b_Akk30, b_Ai00) + safe_matmul(b_Akk31, b_Ai10) + safe_matmul(b_Akk32, b_Ai20)) + b_o = b_o * scale - # ---- store full Akk_inv to Akk ----------------------------------------------- - store_bc_bc(Akk, Akk_base, i_tc0, 0, HV * BT, b_Ai00, BC, BT, T_eff) - store_bc_bc(Akk, Akk_base, i_tc1, 0, HV * BT, b_Ai10, BC, BT, T_eff) - store_bc_bc(Akk, Akk_base, i_tc1, BC, HV * BT, b_Ai11, BC, BT, T_eff) - if NC >= 3: - store_bc_bc(Akk, Akk_base, i_tc2, 0, HV * BT, b_Ai20, BC, BT, T_eff) - store_bc_bc(Akk, Akk_base, i_tc2, BC, HV * BT, b_Ai21, BC, BT, T_eff) - store_bc_bc(Akk, Akk_base, i_tc2, 2 * BC, HV * BT, b_Ai22, BC, BT, T_eff) - if NC >= 4: - store_bc_bc(Akk, Akk_base, i_tc3, 0, HV * BT, b_Ai30, BC, BT, T_eff) - store_bc_bc(Akk, Akk_base, i_tc3, BC, HV * BT, b_Ai31, BC, BT, T_eff) - store_bc_bc(Akk, Akk_base, i_tc3, 2 * BC, HV * BT, b_Ai32, BC, BT, T_eff) - store_bc_bc(Akk, Akk_base, i_tc3, 3 * BC, HV * BT, b_Ai33, BC, BT, T_eff) + v_col = i_v * BV + o_bv + m_v = v_col < V + v_off = v_base + t_row[:, None] * (HV * V) + v_col[None, :] + v_mask = m_t[:, None] & m_v[None, :] + b_v = ct.gather(v, v_off, mask=v_mask, check_bounds=True, padding_value=0.0) + A_col = o_bt + A_off = A_base + t_row[:, None] * (HV * BT) + A_col[None, :] + A_mask = m_t[:, None] & (A_col[None, :] < BT) + b_A = ct.gather(A, A_off, mask=A_mask, check_bounds=True, padding_value=0.0) + b_A = ct.astype(ct.where(m_s, b_A, ct.zeros((BT, BT), dtype=b_A.dtype)), b_v.dtype) + b_o = safe_mma(b_A, b_v, b_o) -# =========================================================================== -# chunk_bwd (dAv) -- dAqk = dO @ V^T ; dV = A @ dO -# =========================================================================== + o_off = o_base + t_row[:, None] * (HV * V) + v_col[None, :] + ct.scatter(o, o_off, ct.astype(b_o, o.dtype), mask=v_mask, check_bounds=True) @ct.kernel @@ -1859,11 +1619,6 @@ def chunk_kda_bwd_kernel_dAv( ct.scatter(dA, dA_off, ct.astype(b_dA, dA.dtype), mask=dA_mask, check_bounds=False) -# =========================================================================== -# chunk_bwd (wy + dqkg fused) -- main backward kernel -# =========================================================================== - - @ct.kernel def chunk_kda_bwd_kernel_wy_dqkg_fused( q, @@ -2100,11 +1855,6 @@ def chunk_kda_bwd_kernel_wy_dqkg_fused( ct.scatter(dA, dA_base + dA_rows * (HV * BT) + dA_cols, ct.astype(b_dA, dA.dtype), mask=dA_m, check_bounds=False) -# =========================================================================== -# chunk_bwd (intra) -- dQ/dK/dG/dBeta from dAqk/dAkk -# =========================================================================== - - @ct.kernel def chunk_kda_bwd_kernel_intra( q, @@ -2378,247 +2128,22 @@ def chunk_kda_bwd_kernel_intra( ) m_i = o_i[:, None] <= j b_gkq = exp2(b_gkj[None, :] - b_g) - b_dkt = b_dkt + ct.where(m_i, b_dAqk[:, None] * b_qj[None, :] * b_gkq, ct.zeros((BC, BK), dtype=ct.float32)) - b_dkt = b_dkt + ct.where(m_i, b_dAkk[:, None] * b_kbj[None, :] * b_gkq, ct.zeros((BC, BK), dtype=ct.float32)) - - dk_off = dk_base + cur_rows * (HV * K) + cur_cols - dg_off = dg_base + cur_rows * (HV * K) + cur_cols - b_dg_prev = ct.astype(ct.gather(dg, dg_off, mask=m_cur, check_bounds=False, padding_value=0.0), ct.float32) - b_dk_prev = ct.astype(ct.gather(dk, dk_off, mask=m_cur, check_bounds=False, padding_value=0.0), ct.float32) - b_dg2 = b_dg2 + (b_dk2 - b_dkt) * ct.astype(b_k, ct.float32) + b_dg_prev - b_dk2_out = b_dk2 + b_dk_prev + b_dkt - ct.scatter(dk2, dk2_base + cur_rows * (HV * K) + cur_cols, ct.astype(b_dk2_out, dk2.dtype), mask=m_cur, check_bounds=False) - ct.scatter(dg2, dg2_base + cur_rows * (HV * K) + cur_cols, ct.astype(b_dg2, dg2.dtype), mask=m_cur, check_bounds=False) - - -# =========================================================================== -# wy_fast (prepare_wy_repr_bwd) -- NOT on the reachable chunk_kda -# path (no launch site); kept for module parity. -# =========================================================================== - - -@ct.kernel -def prepare_wy_repr_bwd_kda_kernel( - k, - v, - beta, - gk, - A, - dA, - dw, - du, - dk, - dk2, - dv, - db, - dg, - dg2, - cu_seqlens, - chunk_indices, - H: ConstInt, - HV: ConstInt, - K: ConstInt, - V: ConstInt, - BT: ConstInt, - BK: ConstInt, - BV: ConstInt, -): - i_t = ct.bid(0) - i_hv = ct.bid(1) - i_h = i_hv // (HV // H) - - i_n = ct.load(chunk_indices, (i_t * 2,), shape=()).item() - i_t = ct.load(chunk_indices, (i_t * 2 + 1,), shape=()).item() - bos = ct.load(cu_seqlens, (i_n,), shape=()).item() - eos = ct.load(cu_seqlens, (i_n + 1,), shape=()).item() - T_eff = eos - bos - - k_base = (bos * H + i_h) * K - v_base = (bos * HV + i_hv) * V - beta_base = bos * HV + i_hv - gk_base = (bos * HV + i_hv) * K - A_base = (bos * HV + i_hv) * BT - dA_base = (bos * HV + i_hv) * BT - dk_base = (bos * HV + i_hv) * K - dk2_base = (bos * HV + i_hv) * K - dw_base = (bos * HV + i_hv) * K - du_base = (bos * HV + i_hv) * V - dv_base = (bos * HV + i_hv) * V - db_base = bos * HV + i_hv - dg_base = (bos * HV + i_hv) * K - dg2_base = (bos * HV + i_hv) * K - - o_bt = ct.arange(BT, dtype=ct.int32) - o_bk = ct.arange(BK, dtype=ct.int32) - o_bv = ct.arange(BV, dtype=ct.int32) - o_t = i_t * BT + o_bt - m_t = o_t < T_eff - - b_b = ct.astype(ct.gather(beta, beta_base + o_t * HV, mask=m_t, check_bounds=False, padding_value=0.0), ct.float32) - # b_A: view (BT, T) stride (1, HV*BT), block (BT, BT) at (0, i_t*BT) - A_off = A_base + o_bt[:, None] * 1 + o_t[None, :] * (HV * BT) - A_mask = (o_bt[:, None] < BT) & m_t[None, :] - b_A = ct.gather(A, A_off, mask=A_mask, check_bounds=False, padding_value=0.0) - - b_db = ct.zeros((BT,), dtype=ct.float32) - b_dA = ct.zeros((BT, BT), dtype=ct.float32) - - for i_k in range(ct.cdiv(K, BK)): - kk = i_k * BK + o_bk - m_k = kk < K - rows = o_t[:, None] - m_tk = m_t[:, None] & m_k[None, :] - b_k = ct.gather(k, k_base + rows * (H * K) + kk[None, :], mask=m_tk, check_bounds=False, padding_value=0.0) - b_gk_exp = exp2( - ct.astype( - ct.gather(gk, gk_base + rows * (HV * K) + kk[None, :], mask=m_tk, check_bounds=False, padding_value=0.0), - ct.float32, - ) - ) - b_dw = ct.gather(dw, dw_base + rows * (HV * K) + kk[None, :], mask=m_tk, check_bounds=False, padding_value=0.0) - b_kbg = ct.astype(b_k, ct.float32) * b_b[:, None] * b_gk_exp - - b_dA = safe_mma(b_dw, ct.transpose(ct.astype(b_kbg, b_dw.dtype)), b_dA) - b_dkbg = safe_matmul(b_A, b_dw) - b_dk_prev = ct.astype( - ct.gather(dk, dk_base + rows * (HV * K) + kk[None, :], mask=m_tk, check_bounds=False, padding_value=0.0), - ct.float32, - ) - b_dg_prev = ct.astype( - ct.gather(dg, dg_base + rows * (HV * K) + kk[None, :], mask=m_tk, check_bounds=False, padding_value=0.0), - ct.float32, - ) - b_dk = ct.astype(b_dkbg, ct.float32) * b_gk_exp * b_b[:, None] + b_dk_prev - b_dg = b_kbg * ct.astype(b_dkbg, ct.float32) + b_dg_prev - b_db = b_db + ct.sum(ct.astype(b_dkbg, ct.float32) * ct.astype(b_k, ct.float32) * b_gk_exp, axis=1) - ct.scatter(dk2, dk2_base + rows * (HV * K) + kk[None, :], ct.astype(b_dk, dk2.dtype), mask=m_tk, check_bounds=False) - ct.scatter(dg2, dg2_base + rows * (HV * K) + kk[None, :], ct.astype(b_dg, dg2.dtype), mask=m_tk, check_bounds=False) - - for i_v in range(ct.cdiv(V, BV)): - vv = i_v * BV + o_bv - m_v = vv < V - rows = o_t[:, None] - m_tv = m_t[:, None] & m_v[None, :] - b_v = ct.gather(v, v_base + rows * (HV * V) + vv[None, :], mask=m_tv, check_bounds=False, padding_value=0.0) - b_du = ct.gather(du, du_base + rows * (HV * V) + vv[None, :], mask=m_tv, check_bounds=False, padding_value=0.0) - b_vb = ct.astype(ct.astype(b_v, ct.float32) * b_b[:, None], b_v.dtype) - b_dA = safe_mma(b_du, ct.transpose(b_vb), b_dA) - b_dvb = safe_matmul(b_A, b_du) - b_dv = ct.astype(b_dvb, ct.float32) * b_b[:, None] - b_db = b_db + ct.sum(ct.astype(b_dvb, ct.float32) * ct.astype(b_v, ct.float32), axis=1) - ct.scatter(dv, dv_base + rows * (HV * V) + vv[None, :], ct.astype(b_dv, dv.dtype), mask=m_tv, check_bounds=False) - - m_A = (o_t[:, None] > o_t[None, :]) & (m_t[:, None] & m_t[None, :]) - b_dA = ct.where(m_A, b_dA, ct.zeros((BT, BT), dtype=ct.float32)) - b_dA = safe_matmul(ct.astype(b_dA, b_A.dtype), b_A) - b_dA = safe_matmul(b_A, ct.astype(b_dA, b_A.dtype)) - b_dA = ct.where(m_A, -b_dA, ct.zeros((BT, BT), dtype=ct.float32)) - - dA_off = dA_base + o_t[:, None] * (HV * BT) + o_bt[None, :] - dA_mask = m_t[:, None] & (o_bt[None, :] < BT) - ct.scatter(dA, dA_off, ct.astype(b_dA, dA.dtype), mask=dA_mask, check_bounds=False) - ct.scatter(db, db_base + o_t * HV, ct.astype(b_db, db.dtype), mask=m_t, check_bounds=False) - - -# =========================================================================== -# Host wrappers (fixed reasonable tile sizes chosen; -# None optionals -> dummy tensor + flag). -# =========================================================================== - - -# --------------------------------------------------------------------------- -# l2norm -# --------------------------------------------------------------------------- -def l2norm_fwd( - x, - eps: float = 1e-6, - out=None, - rstd_out=None, - stream=None, -): - stream = 0 if stream is None else stream - x_shape_og = x.shape - x = reshaped(x, (-1, x.shape[-1])) - y = reshaped(out, tuple(x.shape)) - T, D = x.shape[0], x.shape[-1] - MAX_FUSED_SIZE = 65536 // x.element_size() - BD = min(MAX_FUSED_SIZE, next_power_of_2(D)) - rstd = reshaped(rstd_out, (T,)) - if D <= 512: - BT = 32 - grid = (cdiv(T, BT),) - # l2norm_fwd is a MEMORY-bound row reduction (each row loaded once, - # reduced, written once); higher occupancy hides DRAM latency (do NOT - # force occ=1 here). - _l2_key = ("l2norm_fwd_kernel", int(D), int(BD), int(BT), str(x.dtype), dev_id(x)) - autotuned_launch(l2norm_fwd_kernel, _l2_key, grid, (x, y, rstd, float(eps), T, D, BD, BT), occ_choices=(1, 2, 4, 8), nww_choices=(4,), stream=stream) - else: - ct.launch(stream, (T,), l2norm_fwd_kernel1, (x, y, rstd, float(eps), D, BD)) - return y.view(x_shape_og), rstd.view(x_shape_og[:-1]) - - -def l2norm_bwd(y, rstd, dy, eps: float = 1e-6, out=None, stream=None): - stream = 0 if stream is None else stream - y_shape_og = y.shape - y = y.reshape(-1, dy.shape[-1]) - dy = dy.reshape(-1, dy.shape[-1]) - dx = reshaped(out, tuple(y.shape)) - T, D = y.shape[0], y.shape[-1] - MAX_FUSED_SIZE = 65536 // y.element_size() - BD = min(MAX_FUSED_SIZE, next_power_of_2(D)) - rstd_flat = rstd.reshape(-1) - if D <= 512: - BT = 32 - grid = (cdiv(T, BT),) - ct.launch(stream, grid, l2norm_bwd_kernel, (y, rstd_flat, dy, dx, float(eps), T, D, BD, BT)) - else: - ct.launch(stream, (T,), l2norm_bwd_kernel1, (y, rstd_flat, dy, dx, float(eps), D, BD)) - return dx.view(y_shape_og) - - -# --------------------------------------------------------------------------- -# fused_beta_sigmoid -# --------------------------------------------------------------------------- -BETA_SIGMOID_BLOCK_SIZE = 2048 - - -def fused_beta_sigmoid_fwd(x, scale: float = 1.0, out=None, stream=None): - stream = 0 if stream is None else stream - y = reshaped(out, tuple(x.shape)) - n = x.numel() - grid = (cdiv(n, BETA_SIGMOID_BLOCK_SIZE),) - ct.launch( - stream, - grid, - fused_beta_sigmoid_fwd_kernel, - (reshaped(x, (-1,)), y.reshape(-1), float(scale), n, BETA_SIGMOID_BLOCK_SIZE), - ) - return y - - -def fused_beta_sigmoid_bwd(x, dy, scale: float = 1.0, out=None, stream=None): - stream = 0 if stream is None else stream - dx = reshaped(out, tuple(x.shape)) - n = x.numel() - grid = (cdiv(n, BETA_SIGMOID_BLOCK_SIZE),) - ct.launch( - stream, - grid, - fused_beta_sigmoid_bwd_kernel, - (reshaped(x, (-1,)), dy.reshape(-1), dx.reshape(-1), float(scale), n, BETA_SIGMOID_BLOCK_SIZE), - ) - return dx + b_dkt = b_dkt + ct.where(m_i, b_dAqk[:, None] * b_qj[None, :] * b_gkq, ct.zeros((BC, BK), dtype=ct.float32)) + b_dkt = b_dkt + ct.where(m_i, b_dAkk[:, None] * b_kbj[None, :] * b_gkq, ct.zeros((BC, BK), dtype=ct.float32)) + + dk_off = dk_base + cur_rows * (HV * K) + cur_cols + dg_off = dg_base + cur_rows * (HV * K) + cur_cols + b_dg_prev = ct.astype(ct.gather(dg, dg_off, mask=m_cur, check_bounds=False, padding_value=0.0), ct.float32) + b_dk_prev = ct.astype(ct.gather(dk, dk_off, mask=m_cur, check_bounds=False, padding_value=0.0), ct.float32) + b_dg2 = b_dg2 + (b_dk2 - b_dkt) * ct.astype(b_k, ct.float32) + b_dg_prev + b_dk2_out = b_dk2 + b_dk_prev + b_dkt + ct.scatter(dk2, dk2_base + cur_rows * (HV * K) + cur_cols, ct.astype(b_dk2_out, dk2.dtype), mask=m_cur, check_bounds=False) + ct.scatter(dg2, dg2_base + cur_rows * (HV * K) + cur_cols, ct.astype(b_dg2, dg2.dtype), mask=m_cur, check_bounds=False) -def fused_beta_sigmoid(x, scale: float = 1.0, out=None, stream=None): - """Fused ``scale * sigmoid(x)`` (fp32, written to ``out``).""" - stream = 0 if stream is None else stream - return fused_beta_sigmoid_fwd(x, scale, out=out, stream=stream) +# --- Launchers: normalization and gates ----------------------------------------------------------- -# --------------------------------------------------------------------------- -# chunk_local_cumsum (vector gate) / kda_gate_chunk_cumsum / kda_gate_bwd -# --------------------------------------------------------------------------- BS_LIST_DEFAULT = 32 @@ -2639,417 +2164,156 @@ def chunk_local_cumsum_vector( assert chunk_size == 2 ** (chunk_size.bit_length() - 1), "chunk_size must be a power of 2" BS = min(BS_LIST_DEFAULT, next_power_of_2(S)) g_org = g - g_out = reshaped(out, (T, H, S)) + g_out = out.reshape((T, H, S)) scale_val = float(scale) if scale is not None else 0.0 has_scale = int(scale is not None) - cu_arg = i32_flat(cu_seqlens) - ci_arg = i32_flat(chunk_indices) - grid = (cdiv(S, BS), NT, H) - ct.launch( - stream, - grid, - chunk_local_cumsum_vector_kernel, - ( - reshaped(g_org, (-1,)), - g_out.reshape(-1), - scale_val, - cu_arg, - ci_arg, - H, - S, - BT, - BS, - int(bool(reverse)), - has_scale, - ), - ) - return g_out - - -def chunk_local_cumsum( - g, - chunk_size, - reverse=False, - scale=None, - cu_seqlens=None, - chunk_indices=None, - out=None, - stream=None, -): - stream = 0 if stream is None else stream - return chunk_local_cumsum_vector( - g=g, - chunk_size=chunk_size, - reverse=reverse, - scale=scale, - cu_seqlens=cu_seqlens, - chunk_indices=chunk_indices, - out=out, - stream=stream, - ) - - -def kda_gate_chunk_cumsum( - g, - A_log, - chunk_size, - scale=None, - dt_bias=None, - cu_seqlens=None, - chunk_indices=None, - lower_bound=None, - out=None, - bufs=None, - stream=None, -): - stream = 0 if stream is None else stream - T, H, S = g.shape - BT = chunk_size - NT = len(chunk_indices) - assert chunk_size == 2 ** (chunk_size.bit_length() - 1), "chunk_size must be a power of 2" - BS = min(BS_LIST_DEFAULT, next_power_of_2(S)) - g_out = reshaped(out, (T, H, S)) - dt_arg = reshaped(opt(dt_bias, bufs, dtname(A_log)), (-1,)) - scale_val = float(scale) if scale is not None else 0.0 - lb_val = float(lower_bound) if lower_bound is not None else 0.0 - cu_arg = i32_flat(cu_seqlens) - ci_arg = i32_flat(chunk_indices) + cu_arg = cu_seqlens.reshape(-1) + ci_arg = chunk_indices.reshape(-1) grid = (cdiv(S, BS), NT, H) - ct.launch( - stream, - grid, - kda_gate_chunk_cumsum_vector_kernel, - ( - reshaped(g, (-1,)), - reshaped(A_log, (-1,)), - dt_arg, - g_out.reshape(-1), - scale_val, - cu_arg, - ci_arg, - lb_val, - H, - S, - BT, - BS, - 0, - int(dt_bias is not None), - int(scale is not None), - int(lower_bound is not None), - ), - ) - return g_out - - -def kda_gate_bwd(g, A_log, dt_bias=None, dyg=None, lower_bound=None, dg_out=None, dA_out=None, dbias_out=None, bufs=None, stream=None): - """Vector-gate backward. ``dg_out`` (g-shaped), ``dA_out`` (A_log-shaped) and — - with ``dt_bias`` — ``dbias_out`` (H*K) are written in place; ``bufs['dA_gate']`` - / ``bufs['db_gate']`` hold the (NT, H) / (NT, H*K) fp32 chunk partials.""" - stream = 0 if stream is None else stream - H, K = g.shape[-2:] - T = g.numel() // (H * K) - BT = 32 - NT = cdiv(T, BT) - BD = next_power_of_2(K) - dg = reshaped(dg_out, tuple(g.shape)) - dA_nt = reshaped(bufs["dA_gate"], (NT, H)) - db_nt = reshaped(bufs["db_gate"], (NT, H * K)) if dt_bias is not None else None - dt_arg = reshaped(opt(dt_bias, bufs, dtname(A_log)), (-1,)) - db_arg = db_nt.reshape(-1) if db_nt is not None else dummy("float32", bufs) - lb_val = float(lower_bound) if lower_bound is not None else 0.0 - grid = (NT, H) - ct.launch( - stream, - grid, - kda_gate_bwd_kernel, - ( - reshaped(g, (-1,)), - reshaped(A_log, (-1,)), - dt_arg, - reshaped(dyg, (-1,)), - dg.reshape(-1), - dA_nt.reshape(-1), - db_arg, - lb_val, - T, - H, - K, - BT, - BD, - int(dt_bias is not None), - int(lower_bound is not None), - ), - ) - sum_leading(reshaped(dA_out, (H,)), dA_nt, NT, H, stream=stream) - if dt_bias is not None: - sum_leading(reshaped(dbias_out, (H * K,)), db_nt, NT, H * K, stream=stream) - return dg, dA_out, (dbias_out if dt_bias is not None else None) - - -# --------------------------------------------------------------------------- -# chunk_gla_fwd_o_gk -# --------------------------------------------------------------------------- -def chunk_gla_fwd_o_gk(q, v, g, A, h, scale, state_v_first=False, cu_seqlens=None, chunk_size=64, chunk_indices=None, bufs=None, stream=None): - stream = 0 if stream is None else stream - T, H, K, HV, V = *q.shape, v.shape[1], v.shape[-1] - BT = chunk_size - NT = len(chunk_indices) - o = bufs["o"] - zero_fill(o, stream=stream) - BK = min(max(next_power_of_2(K), 16), 64) - # BV=128 when V<=128 removes the V grid-split - # (grid dim0 -> 1), doubling work/block but halving launched blocks. - BV = 128 if V <= 128 else 64 - dev = dev_id(q) - cu_arg = i32_flat(cu_seqlens) - ci_arg = i32_flat(chunk_indices) - - grid = (cdiv(V, BV), NT, HV) - - # Kernel uses flat element-offset gather/scatter; pass 1-D views so the - # cuTile index-tuple rank (1) matches the array rank. - _q_arg = reshaped(q, (-1,)) - _v_arg = v.reshape(-1) - _g_arg = g.reshape(-1) - _h_arg = h.reshape(-1) - _o_arg = reshaped(o, (-1,)) - _A_arg = A.reshape(-1) - _o_args = ( - _q_arg, - _v_arg, - _g_arg, - _h_arg, - _o_arg, - _A_arg, - cu_arg, - ci_arg, - float(scale), - H, - HV, - K, - V, - BT, - BK, - BV, - int(state_v_first), - ) - # Launch-hint autotune (occupancy x num_worker_warps) on this - # output-projection kernel. - _o_key = ( - "chunk_gla_fwd_kernel_o", - int(H), - int(HV), - int(K), - int(V), - int(BT), - int(BK), - int(BV), - int(state_v_first), - str(q.dtype), - str(dev), - ) - autotuned_launch(chunk_gla_fwd_kernel_o, _o_key, grid, _o_args, occ_choices=(1, 2, 4), nww_choices=(4,), stream=stream) - return o - - -# --------------------------------------------------------------------------- -# chunk_gated_delta_rule_fwd_h / bwd_dhu (shared chunk_delta_h) -# --------------------------------------------------------------------------- -def chunk_gated_delta_rule_fwd_h( - k, - w, - u, - g=None, - gk=None, - initial_state=None, - output_final_state=False, - chunk_size=64, - save_new_value=True, - state_v_first=False, - cu_seqlens=None, - cu_seqlens_cpu=None, - chunk_indices=None, - bufs=None, - stream=None, -): - stream = 0 if stream is None else stream - T, H, K, V, HV = *k.shape, u.shape[-1], u.shape[1] - BT = chunk_size - N, NT = len(cu_seqlens) - 1, len(chunk_indices) - chunk_offsets = bufs["chunk_offsets"] - # Full-width state/K tile: BK = next_pow2(K). The blockdim64 kernel carries the - # KV state as a single (BV, BK)/(BK, BV) tile (no K<=256 cap); K-axis loads - # zero-pad [K:BK] and stores drop it, so any K is supported. - BK = next_power_of_2(K) - - state_shape = (N, HV, V, K) if state_v_first else (N, HV, K, V) - h = reshaped(bufs["state_checkpoints"], (NT, HV) + state_shape[2:]) - final_state = reshaped(bufs["final_state"], state_shape) if output_final_state else None - if final_state is not None: - zero_fill(final_state, stream=stream) - v_new = reshaped(bufs["v_new"], (T, HV, V)) if save_new_value else None - - dev = dev_id(k) - vnew_arg = v_new if v_new is not None else dummy(dtname(u), bufs) - g_arg = opt(g, bufs) - gk_arg = opt(gk, bufs) - h0_arg = initial_state if initial_state is not None else dummy("float32", bufs) - ht_arg = final_state if final_state is not None else dummy("float32", bufs) - cu_arg = i32_flat(cu_seqlens) - co_arg = i32_flat(chunk_offsets) - - # Kernel uses flat element-offset gather/scatter; pass 1-D views so the - # cuTile index-tuple rank (1) matches the array rank. - _k_arg = k.reshape(-1) - _u_arg = u.reshape(-1) - _w_arg = w.reshape(-1) - _vnew_arg = vnew_arg.reshape(-1) - _h_arg = h.reshape(-1) - _h01d = reshaped(h0_arg, (-1,)) - _ht1d = reshaped(ht_arg, (-1,)) - _g1d = g_arg.reshape(-1) - _gk1d = gk_arg.reshape(-1) - - # BV is a V-tiling block width (drives the grid V-fan-out and the V-tile - # shapes only; the K axis is always split into fixed 64-wide blocks, so BV - # never changes the numerics). The tuner picks the per-shape grid fill that - # best hides the inter-chunk latency. - def _grid_fn(bv): - return (cdiv(V, bv), N * HV) - - def _args_fn(bv): - return ( - _k_arg, - _u_arg, - _w_arg, - _vnew_arg, - _g1d, - _gk1d, - _h_arg, - _h01d, - _ht1d, + ct.launch( + stream, + grid, + chunk_local_cumsum_vector_kernel, + ( + g_org.reshape((-1,)), + g_out.reshape(-1), + scale_val, cu_arg, - co_arg, + ci_arg, H, - HV, - K, - V, + S, BT, - bv, - BK, - int(g is not None), - int(gk is not None), - int(initial_state is not None), - int(output_final_state), - int(save_new_value), - int(state_v_first), - ) + BS, + int(bool(reverse)), + has_scale, + ), + ) + return g_out - _h_key = ( - "chunk_gated_delta_rule_fwd_kernel_h_blockdim64", - int(H), - int(HV), - int(K), - int(V), - int(BT), - int(g is not None), - int(gk is not None), - int(initial_state is not None), - int(output_final_state), - int(save_new_value), - int(state_v_first), - str(k.dtype), - str(dev), + +def chunk_local_cumsum( + g, + chunk_size, + reverse=False, + scale=None, + cu_seqlens=None, + chunk_indices=None, + out=None, + stream=None, +): + stream = 0 if stream is None else stream + return chunk_local_cumsum_vector( + g=g, + chunk_size=chunk_size, + reverse=reverse, + scale=scale, + cu_seqlens=cu_seqlens, + chunk_indices=chunk_indices, + out=out, + stream=stream, ) - # BV candidates: divisors of the V tile that keep N>=8 on the MMA V-axis, - # capped at <= V so a single tile is never larger than V. The grid-fill - # tradeoff is shape-dependent, not monotone: shrinking BV multiplies the - # V-tile CTA count but each CTA is register-bound (255 reg -> ~1 block/SM) - # and redundantly reloads K/W/G, so past the point where the grid already - # fills the SMs a *larger* BV wins (fewer, fatter CTAs, higher warp - # occupancy); 16 stays first as the safe DISABLE_TUNE/failure fallback. - _bv_choices = tuple(bv for bv in (16, 8, 32, 64) if bv <= max(V, 8)) - if not _bv_choices: - _bv_choices = (min(32, V),) - autotuned_launch_bv(chunk_gated_delta_rule_fwd_kernel_h_blockdim64, _h_key, _bv_choices, _grid_fn, _args_fn, stream=stream) - return h, v_new, final_state -def chunk_gated_delta_rule_bwd_dhu( - q, - k, - w, - do, - dv, - g=None, - gk=None, - h0=None, - dstate_in=None, - scale: float | None = None, - state_v_first: bool = False, +def kda_gate_chunk_cumsum( + g, + A_log, + chunk_size, + scale=None, + dt_bias=None, cu_seqlens=None, - chunk_size: int = 64, chunk_indices=None, - bufs: dict | None = None, + lower_bound=None, + out=None, + bufs=None, stream=None, ): stream = 0 if stream is None else stream - T, H, K, V, HV = *q.shape, do.shape[-1], do.shape[1] + T, H, S = g.shape BT = chunk_size - # Full-width state/K tile: BK = next_pow2(K). The blockdim64 dhu kernel carries - # the dH state as a single (BV, BK)/(BK, BV) tile (no K<=256 cap); K-axis loads - # zero-pad/mask [K:BK] and stores drop it, so any K is supported. dH/dH0 - # carves stay K-wide. - BK = next_power_of_2(K) - N, NT = len(cu_seqlens) - 1, len(chunk_indices) - chunk_offsets = bufs["chunk_offsets"] - - dh = reshaped(bufs["dstate"], (NT, HV, V, K) if state_v_first else (NT, HV, K, V)) - dh0 = reshaped(bufs["dstate0"], tuple(h0.shape)) if h0 is not None else None - dv2 = reshaped(bufs["dv_dstate_u"], (T, HV, V)) + NT = len(chunk_indices) + assert chunk_size == 2 ** (chunk_size.bit_length() - 1), "chunk_size must be a power of 2" + BS = min(BS_LIST_DEFAULT, next_power_of_2(S)) + g_out = out.reshape((T, H, S)) + dt_arg = opt(dt_bias, bufs, dtname(A_log)).reshape((-1,)) + scale_val = float(scale) if scale is not None else 0.0 + lb_val = float(lower_bound) if lower_bound is not None else 0.0 + cu_arg = cu_seqlens.reshape(-1) + ci_arg = chunk_indices.reshape(-1) + grid = (cdiv(S, BS), NT, H) + ct.launch( + stream, + grid, + kda_gate_chunk_cumsum_vector_kernel, + ( + g.reshape((-1,)), + A_log.reshape((-1,)), + dt_arg, + g_out.reshape(-1), + scale_val, + cu_arg, + ci_arg, + lb_val, + H, + S, + BT, + BS, + 0, + int(dt_bias is not None), + int(scale is not None), + int(lower_bound is not None), + ), + ) + return g_out - BV = 64 - grid = (cdiv(V, BV), N * HV) +def kda_gate_bwd(g, A_log, dt_bias=None, dyg=None, lower_bound=None, dg_out=None, dA_out=None, dbias_out=None, bufs=None, stream=None): + """Vector-gate backward. ``dg_out`` (g-shaped), ``dA_out`` (A_log-shaped) and — + with ``dt_bias`` — ``dbias_out`` (H*K) are written in place; ``bufs['dA_gate']`` + / ``bufs['db_gate']`` hold the (NT, H) / (NT, H*K) fp32 chunk partials.""" + stream = 0 if stream is None else stream + H, K = g.shape[-2:] + T = g.numel() // (H * K) + BT = 32 + NT = cdiv(T, BT) + BD = next_power_of_2(K) + dg = dg_out.reshape(tuple(g.shape)) + dA_nt = bufs["dA_gate"].reshape((NT, H)) + db_nt = bufs["db_gate"].reshape((NT, H * K)) if dt_bias is not None else None + dt_arg = opt(dt_bias, bufs, dtname(A_log)).reshape((-1,)) + db_arg = opt(db_nt, bufs).reshape(-1) + lb_val = float(lower_bound) if lower_bound is not None else 0.0 + grid = (NT, H) ct.launch( stream, grid, - chunk_gated_delta_rule_bwd_kernel_dhu_blockdim64, + kda_gate_bwd_kernel, ( - q.reshape(-1), - k.reshape(-1), - w.reshape(-1), - opt(g, bufs).reshape(-1), - opt(gk, bufs).reshape(-1), - reshaped(opt(dstate_in, bufs), (-1,)), - reshaped(dh0 if dh0 is not None else dummy("float32", bufs), (-1,)), - reshaped(do, (-1,)), - dh.reshape(-1), - dv.reshape(-1), - dv2.reshape(-1), - i32_flat(cu_seqlens), - i32_flat(chunk_offsets), - float(scale), + g.reshape((-1,)), + A_log.reshape((-1,)), + dt_arg, + dyg.reshape((-1,)), + dg.reshape(-1), + dA_nt.reshape(-1), + db_arg, + lb_val, + T, H, - HV, K, - V, BT, - BV, - BK, - int(g is not None), - int(gk is not None), - int(h0 is not None), - int(dstate_in is not None), - int(state_v_first), + BD, + int(dt_bias is not None), + int(lower_bound is not None), ), ) - return dh, dh0, dv2 + sum_leading(dA_out.reshape((H,)), dA_nt, NT, H, stream=stream) + if dt_bias is not None: + sum_leading(dbias_out.reshape((H * K,)), db_nt, NT, H * K, stream=stream) + return dg, dA_out, (dbias_out if dt_bias is not None else None) + + +# --- Launchers: WY representation ----------------------------------------------------------------- -# --------------------------------------------------------------------------- -# recompute_w_u_fwd -# --------------------------------------------------------------------------- def recompute_w_u_fwd(k, v, beta, A, gk, q=None, cu_seqlens=None, chunk_indices=None, bufs=None, stream=None): stream = 0 if stream is None else stream T, H, K, V = *k.shape, v.shape[-1] @@ -3058,28 +2322,28 @@ def recompute_w_u_fwd(k, v, beta, A, gk, q=None, cu_seqlens=None, chunk_indices= BK = 32 BV = 32 NT = len(chunk_indices) - dev = dev_id(k) + dev = current_device_id() - w = reshaped(bufs["w"], (T, HV, K)) - u = reshaped(bufs["u"], (T, HV, V)) - qg = reshaped(bufs["qg"], (T, HV, K)) if q is not None else None - kg = reshaped(bufs["kg"], (T, HV, K)) + w = bufs["w"].reshape((T, HV, K)) + u = bufs["u"].reshape((T, HV, V)) + qg = bufs["qg"].reshape((T, HV, K)) if q is not None else None + kg = bufs["kg"].reshape((T, HV, K)) - cu_arg = i32_flat(cu_seqlens) - ci_arg = i32_flat(chunk_indices) + cu_arg = cu_seqlens.reshape(-1) + ci_arg = chunk_indices.reshape(-1) q_arg = opt(q, bufs, dtname(k)) - qg_arg = qg if qg is not None else dummy(dtname(k), bufs) + qg_arg = opt(qg, bufs, dtname(k)) # Kernel uses flat element-offset gather/scatter; pass 1-D views so the # cuTile index-tuple rank (1) matches the array rank. reshape(-1) on these # contiguous tensors yields storage-aliasing views (outputs W/U/QG/KG too). _wu_grid = (NT, HV) _wu_args = ( - reshaped(q_arg, (-1,)), - reshaped(k, (-1,)), + q_arg.reshape((-1,)), + k.reshape((-1,)), qg_arg.reshape(-1), kg.reshape(-1), - reshaped(v, (-1,)), - reshaped(beta, (-1,)), + v.reshape((-1,)), + beta.reshape((-1,)), w.reshape(-1), u.reshape(-1), A.reshape(-1), @@ -3117,9 +2381,6 @@ def recompute_w_u_fwd(k, v, beta, A, gk, q=None, cu_seqlens=None, chunk_indices= return w, u, qg, kg -# --------------------------------------------------------------------------- -# chunk_kda_fwd_intra_token_parallel / chunk_kda_fwd_intra -# --------------------------------------------------------------------------- def chunk_kda_fwd_intra_token_parallel(q, k, gk, beta, Aqk, Akk, scale, cu_seqlens=None, chunk_size=64, sub_chunk_size=16, stream=None): stream = 0 if stream is None else stream T, H, K, HV = *q.shape, gk.shape[1] @@ -3131,7 +2392,7 @@ def chunk_kda_fwd_intra_token_parallel(q, k, gk, beta, Aqk, Akk, scale, cu_seqle # BH in {1,2,4,8}; HV must be divisible for the grid split. BH = 4 if (HV % 4 == 0) else (2 if (HV % 2 == 0) else 1) BK = next_power_of_2(K) - cu_arg = i32_flat(cu_seqlens) + cu_arg = cu_seqlens.reshape(-1) grid = (T, cdiv(HV, BH)) # cuTile gather/scatter index-tuple rank must match the array rank, so pass # pre-flattened views: Q/K -> (T*H, K), Gk -> (T*HV, K), Beta/Aqk/Akk -> 1-D. @@ -3141,10 +2402,10 @@ def chunk_kda_fwd_intra_token_parallel(q, k, gk, beta, Aqk, Akk, scale, cu_seqle grid, chunk_kda_fwd_kernel_intra_token_parallel, ( - reshaped(q, (-1, K)), - reshaped(k, (-1, K)), + q.reshape((-1, K)), + k.reshape((-1, K)), gk.reshape(-1, K), - reshaped(beta, (-1,)), + beta.reshape((-1,)), Aqk.reshape(-1), Akk.reshape(-1), float(scale), @@ -3195,15 +2456,15 @@ def chunk_kda_fwd_intra( # inter_diag_compute_solve so inter_solve_fused can SKIP forward-substitution. use_split_diag_compute_solve = (not safe_gate) and BT == 64 and K >= 64 use_solved_diagonal = safe_gate or use_split_diag_compute_solve - dev = dev_id(k) + dev = current_device_id() - Aqk = reshaped(bufs["Aqk"], (T, HV, BT)) - Akk = reshaped(bufs["Akk"], (T, HV, BT)) + Aqk = bufs["Aqk"].reshape((T, HV, BT)) + Akk = bufs["Akk"].reshape((T, HV, BT)) zero_fill(Akk, stream=stream) - Akkd = reshaped(bufs["Akkd"], (T, HV, BC)) + Akkd = bufs["Akkd"].reshape((T, HV, BC)) - cu_arg = i32_flat(cu_seqlens) - ci_arg = i32_flat(chunk_indices) + cu_arg = cu_seqlens.reshape(-1) + ci_arg = chunk_indices.reshape(-1) # Step 1: diagonal blocks into Akkd (fp32). When use_solved_diagonal is set # (safe_gate OR the split path) the diagonals are PRE-SOLVED @@ -3214,10 +2475,10 @@ def chunk_kda_fwd_intra( # cuTile index-tuple rank (1) matches the array rank. _diag_grid = (NT, NC, HV) _diag_args = ( - reshaped(q, (-1,)), - reshaped(k, (-1,)), + q.reshape((-1,)), + k.reshape((-1,)), gk.reshape(-1), - reshaped(beta, (-1,)), + beta.reshape((-1,)), Aqk.reshape(-1), Akkd.reshape(-1), float(scale), @@ -3260,10 +2521,10 @@ def chunk_kda_fwd_intra( # solve kernel. _isf_grid = (NT, HV) _isf_args = ( - reshaped(q, (-1,)), - reshaped(k, (-1,)), + q.reshape((-1,)), + k.reshape((-1,)), gk.reshape(-1), - reshaped(beta, (-1,)), + beta.reshape((-1,)), Aqk.reshape(-1), Akkd.reshape(-1), Akk.reshape(-1), @@ -3285,87 +2546,281 @@ def chunk_kda_fwd_intra( int(HV), int(K), int(BT), - int(BC), - int(NC), - 1, - int(use_solved_diagonal), + int(BC), + int(NC), + 1, + int(use_solved_diagonal), + str(k.dtype), + str(dev), + ) + autotuned_launch(chunk_kda_fwd_kernel_inter_solve_fused, _isf_key, _isf_grid, _isf_args, occ_choices=(1, 2, 4), nww_choices=(4,), stream=stream) + w, u, qg, kg = recompute_w_u_fwd( + k=k, v=v, beta=beta, A=Akk, q=q if disable_recompute else None, gk=gk, cu_seqlens=cu_seqlens, chunk_indices=chunk_indices, bufs=bufs, stream=stream + ) + return w, u, qg, kg, Aqk, Akk + + +# --- Launchers: state scan ------------------------------------------------------------------------ + + +def chunk_gated_delta_rule_fwd_h( + k, + w, + u, + g=None, + gk=None, + initial_state=None, + output_final_state=False, + chunk_size=64, + save_new_value=True, + state_v_first=False, + cu_seqlens=None, + cu_seqlens_cpu=None, + chunk_indices=None, + bufs=None, + stream=None, +): + stream = 0 if stream is None else stream + T, H, K, V, HV = *k.shape, u.shape[-1], u.shape[1] + BT = chunk_size + N, NT = len(cu_seqlens) - 1, len(chunk_indices) + chunk_offsets = bufs["chunk_offsets"] + # Full-width state/K tile: BK = next_pow2(K). The blockdim64 kernel carries the + # KV state as a single (BV, BK)/(BK, BV) tile (no K<=256 cap); K-axis loads + # zero-pad [K:BK] and stores drop it, so any K is supported. + BK = next_power_of_2(K) + + state_shape = (N, HV, V, K) if state_v_first else (N, HV, K, V) + h = bufs["state_checkpoints"].reshape((NT, HV) + state_shape[2:]) + final_state = bufs["final_state"].reshape(state_shape) if output_final_state else None + if final_state is not None: + zero_fill(final_state, stream=stream) + v_new = bufs["v_new"].reshape((T, HV, V)) if save_new_value else None + + dev = current_device_id() + vnew_arg = opt(v_new, bufs, dtname(u)) + g_arg = opt(g, bufs) + gk_arg = opt(gk, bufs) + h0_arg = opt(initial_state, bufs) + ht_arg = opt(final_state, bufs) + cu_arg = cu_seqlens.reshape(-1) + co_arg = chunk_offsets.reshape(-1) + + # Kernel uses flat element-offset gather/scatter; pass 1-D views so the + # cuTile index-tuple rank (1) matches the array rank. + _k_arg = k.reshape(-1) + _u_arg = u.reshape(-1) + _w_arg = w.reshape(-1) + _vnew_arg = vnew_arg.reshape(-1) + _h_arg = h.reshape(-1) + _h01d = h0_arg.reshape((-1,)) + _ht1d = ht_arg.reshape((-1,)) + _g1d = g_arg.reshape(-1) + _gk1d = gk_arg.reshape(-1) + + # BV is a V-tiling block width (drives the grid V-fan-out and the V-tile + # shapes only; the K axis is always split into fixed 64-wide blocks, so BV + # never changes the numerics). The tuner picks the per-shape grid fill that + # best hides the inter-chunk latency. + def _grid_fn(bv): + return (cdiv(V, bv), N * HV) + + def _args_fn(bv): + return ( + _k_arg, + _u_arg, + _w_arg, + _vnew_arg, + _g1d, + _gk1d, + _h_arg, + _h01d, + _ht1d, + cu_arg, + co_arg, + H, + HV, + K, + V, + BT, + bv, + BK, + int(g is not None), + int(gk is not None), + int(initial_state is not None), + int(output_final_state), + int(save_new_value), + int(state_v_first), + ) + + _h_key = ( + "chunk_gated_delta_rule_fwd_kernel_h_blockdim64", + int(H), + int(HV), + int(K), + int(V), + int(BT), + int(g is not None), + int(gk is not None), + int(initial_state is not None), + int(output_final_state), + int(save_new_value), + int(state_v_first), str(k.dtype), str(dev), ) - autotuned_launch(chunk_kda_fwd_kernel_inter_solve_fused, _isf_key, _isf_grid, _isf_args, occ_choices=(1, 2, 4), nww_choices=(4,), stream=stream) - w, u, qg, kg = recompute_w_u_fwd( - k=k, v=v, beta=beta, A=Akk, q=q if disable_recompute else None, gk=gk, cu_seqlens=cu_seqlens, chunk_indices=chunk_indices, bufs=bufs, stream=stream - ) - return w, u, qg, kg, Aqk, Akk + # BV candidates: divisors of the V tile that keep N>=8 on the MMA V-axis, + # capped at <= V so a single tile is never larger than V. The grid-fill + # tradeoff is shape-dependent, not monotone: shrinking BV multiplies the + # V-tile CTA count but each CTA is register-bound (255 reg -> ~1 block/SM) + # and redundantly reloads K/W/G, so past the point where the grid already + # fills the SMs a *larger* BV wins (fewer, fatter CTAs, higher warp + # occupancy); 16 stays first as the safe fallback when tuning fails. + _bv_choices = tuple(bv for bv in (16, 8, 32, 64) if bv <= max(V, 8)) + if not _bv_choices: + _bv_choices = (min(32, V),) + autotuned_launch_bv(chunk_gated_delta_rule_fwd_kernel_h_blockdim64, _h_key, _bv_choices, _grid_fn, _args_fn, stream=stream) + return h, v_new, final_state -def chunk_kda_bwd_intra(q, k, g, beta, dAqk, dAkk, dq, dk, db, dg, cu_seqlens=None, chunk_indices=None, chunk_size=64, safe_gate=False, bufs=None, stream=None): +def chunk_gated_delta_rule_bwd_dhu( + q, + k, + w, + do, + dv, + g=None, + gk=None, + h0=None, + dstate_in=None, + scale: float | None = None, + state_v_first: bool = False, + cu_seqlens=None, + chunk_size: int = 64, + chunk_indices=None, + bufs: dict | None = None, + stream=None, +): stream = 0 if stream is None else stream - T, H, K, HV = *k.shape, g.shape[1] + T, H, K, V, HV = *q.shape, do.shape[-1], do.shape[1] BT = chunk_size - # Fast path: for BT >= 64 use larger - # BC=32 / BK=64 sub-tiles and route through the SAFE_GATE matmul branch - # instead of the BC-iteration scalar `for j` loops. The scalar path is the - # bwd bottleneck. - use_fast_path = BT >= 64 - BC = 32 if use_fast_path else min(16, BT) - BK = min(64 if use_fast_path else 32, next_power_of_2(K)) - safe_gate = safe_gate or use_fast_path - NT = len(chunk_indices) - NC = cdiv(BT, BC) - NK = cdiv(K, BK) + # Full-width state/K tile: BK = next_pow2(K). The blockdim64 dhu kernel carries + # the dH state as a single (BV, BK)/(BK, BV) tile (no K<=256 cap); K-axis loads + # zero-pad/mask [K:BK] and stores drop it, so any K is supported. dH/dH0 + # carves stay K-wide. + BK = next_power_of_2(K) + N, NT = len(cu_seqlens) - 1, len(chunk_indices) + chunk_offsets = bufs["chunk_offsets"] - dq2 = reshaped(bufs["dq2"], (T, HV, K)) - dk2 = reshaped(bufs["dk2"], (T, HV, K)) - db2 = reshaped(bufs["db2"], (NK, T, HV)) - dg2 = reshaped(bufs["dg2"], (T, HV, K)) + dh = bufs["dstate"].reshape((NT, HV, V, K) if state_v_first else (NT, HV, K, V)) + dh0 = bufs["dstate0"].reshape(tuple(h0.shape)) if h0 is not None else None + dv2 = bufs["dv_dstate_u"].reshape((T, HV, V)) + + BV = 64 + grid = (cdiv(V, BV), N * HV) - cu_arg = i32_flat(cu_seqlens) - ci_arg = i32_flat(chunk_indices) - grid = (NK * NC, NT, HV) - # Kernel uses flat element-offset gather/scatter; pass 1-D views. ct.launch( stream, grid, - chunk_kda_bwd_kernel_intra, + chunk_gated_delta_rule_bwd_kernel_dhu_blockdim64, ( - reshaped(q, (-1,)), - reshaped(k, (-1,)), - g.reshape(-1), - reshaped(beta, (-1,)), - dAqk.reshape(-1), - dAkk.reshape(-1), - dq.reshape(-1), - dq2.reshape(-1), - dk.reshape(-1), - dk2.reshape(-1), - dg.reshape(-1), - dg2.reshape(-1), - db2.reshape(-1), - cu_arg, - ci_arg, - T, + q.reshape(-1), + k.reshape(-1), + w.reshape(-1), + opt(g, bufs).reshape(-1), + opt(gk, bufs).reshape(-1), + opt(dstate_in, bufs).reshape((-1,)), + opt(dh0, bufs).reshape((-1,)), + do.reshape((-1,)), + dh.reshape(-1), + dv.reshape(-1), + dv2.reshape(-1), + cu_seqlens.reshape(-1), + chunk_offsets.reshape(-1), + float(scale), H, HV, K, + V, BT, - BC, + BV, BK, - NC, - int(safe_gate), + int(g is not None), + int(gk is not None), + int(h0 is not None), + int(dstate_in is not None), + int(state_v_first), ), ) - dq = dq2 - dk = dk2 - # dBeta += sum_nk dBeta2 (fp32 acc); the fan-in NK is a compile-time constant - sum_leading(reshaped(db, (T * HV,)), reshaped(db2, (NK, T * HV)), NK, T * HV, stream=stream, accumulate=True) - dg = dg2 - return dq, dk, db, dg + return dh, dh0, dv2 + + +# --- Launchers: attention and gradients ----------------------------------------------------------- + + +def chunk_gla_fwd_o_gk(q, v, g, A, h, scale, state_v_first=False, cu_seqlens=None, chunk_size=64, chunk_indices=None, bufs=None, stream=None): + stream = 0 if stream is None else stream + T, H, K, HV, V = *q.shape, v.shape[1], v.shape[-1] + BT = chunk_size + NT = len(chunk_indices) + o = bufs["o"] + zero_fill(o, stream=stream) + BK = min(max(next_power_of_2(K), 16), 64) + # BV=128 when V<=128 removes the V grid-split + # (grid dim0 -> 1), doubling work/block but halving launched blocks. + BV = 128 if V <= 128 else 64 + dev = current_device_id() + cu_arg = cu_seqlens.reshape(-1) + ci_arg = chunk_indices.reshape(-1) + + grid = (cdiv(V, BV), NT, HV) + + # Kernel uses flat element-offset gather/scatter; pass 1-D views so the + # cuTile index-tuple rank (1) matches the array rank. + _q_arg = q.reshape((-1,)) + _v_arg = v.reshape(-1) + _g_arg = g.reshape(-1) + _h_arg = h.reshape(-1) + _o_arg = o.reshape((-1,)) + _A_arg = A.reshape(-1) + _o_args = ( + _q_arg, + _v_arg, + _g_arg, + _h_arg, + _o_arg, + _A_arg, + cu_arg, + ci_arg, + float(scale), + H, + HV, + K, + V, + BT, + BK, + BV, + int(state_v_first), + ) + # Launch-hint autotune (occupancy x num_worker_warps) on this + # output-projection kernel. + _o_key = ( + "chunk_gla_fwd_kernel_o", + int(H), + int(HV), + int(K), + int(V), + int(BT), + int(BK), + int(BV), + int(state_v_first), + str(q.dtype), + str(dev), + ) + autotuned_launch(chunk_gla_fwd_kernel_o, _o_key, grid, _o_args, occ_choices=(1, 2, 4), nww_choices=(4,), stream=stream) + return o -# --------------------------------------------------------------------------- -# chunk_kda_bwd_dAv / chunk_kda_bwd_wy_dqkg_fused -# --------------------------------------------------------------------------- def chunk_kda_bwd_dAv(q, k, v, do, A=None, scale=None, cu_seqlens=None, chunk_size=64, chunk_indices=None, bufs=None, stream=None): stream = 0 if stream is None else stream T, H, K, HV, V = *k.shape, do.shape[1], do.shape[-1] @@ -3375,21 +2830,21 @@ def chunk_kda_bwd_dAv(q, k, v, do, A=None, scale=None, cu_seqlens=None, chunk_si BV = min(max(next_power_of_2(V), 16), CONST_TILING) NT = len(chunk_indices) - dA = reshaped(bufs["dAqk"], (T, HV, BT)) - dv = reshaped(bufs["dv_dAv"], (T, HV, V)) - cu_arg = i32_flat(cu_seqlens) - ci_arg = i32_flat(chunk_indices) + dA = bufs["dAqk"].reshape((T, HV, BT)) + dv = bufs["dv_dAv"].reshape((T, HV, V)) + cu_arg = cu_seqlens.reshape(-1) + ci_arg = chunk_indices.reshape(-1) # Kernel uses flat element-offset gather/scatter; pass 1-D views. ct.launch( stream, (NT, HV), chunk_kda_bwd_kernel_dAv, ( - reshaped(q, (-1,)), - reshaped(k, (-1,)), + q.reshape((-1,)), + k.reshape((-1,)), v.reshape(-1), A.reshape(-1), - reshaped(do, (-1,)), + do.reshape((-1,)), dv.reshape(-1), dA.reshape(-1), cu_arg, @@ -3435,12 +2890,12 @@ def chunk_kda_bwd_wy_dqkg_fused( BK = min(max(next_power_of_2(K), 16), CONST_TILING) BV = min(max(next_power_of_2(V), 16), CONST_TILING) - dq = reshaped(bufs["dq"], (T, HV, K)) - dk = reshaped(bufs["dk"], (T, HV, K)) - dv2 = reshaped(bufs["dv2"], (T, HV, V)) - dg = reshaped(bufs["dg"], (T, HV, K)) - db = reshaped(bufs["db"], (T, HV)) - dA = reshaped(bufs["dAkk"], (T, HV, BT)) + dq = bufs["dq"].reshape((T, HV, K)) + dk = bufs["dk"].reshape((T, HV, K)) + dv2 = bufs["dv2"].reshape((T, HV, V)) + dg = bufs["dg"].reshape((T, HV, K)) + db = bufs["db"].reshape((T, HV)) + dA = bufs["dAkk"].reshape((T, HV, BT)) grid = (NT, HV) @@ -3449,15 +2904,15 @@ def chunk_kda_bwd_wy_dqkg_fused( grid, chunk_kda_bwd_kernel_wy_dqkg_fused, ( - reshaped(q, (-1,)), - reshaped(k, (-1,)), - reshaped(v, (-1,)), + q.reshape((-1,)), + k.reshape((-1,)), + v.reshape((-1,)), v_new.reshape(-1), g.reshape(-1), - reshaped(beta, (-1,)), + beta.reshape((-1,)), A.reshape(-1), h.reshape(-1), - reshaped(do, (-1,)), + do.reshape((-1,)), dh.reshape(-1), dq.reshape(-1), dk.reshape(-1), @@ -3466,8 +2921,8 @@ def chunk_kda_bwd_wy_dqkg_fused( dg.reshape(-1), db.reshape(-1), dA.reshape(-1), - i32_flat(cu_seqlens), - i32_flat(chunk_indices), + cu_seqlens.reshape(-1), + chunk_indices.reshape(-1), float(scale), H, HV, @@ -3483,9 +2938,71 @@ def chunk_kda_bwd_wy_dqkg_fused( return dq, dk, dv, db, dg, dA -# =========================================================================== -# chunk_fwd / chunk_bwd orchestration -# =========================================================================== +def chunk_kda_bwd_intra(q, k, g, beta, dAqk, dAkk, dq, dk, db, dg, cu_seqlens=None, chunk_indices=None, chunk_size=64, safe_gate=False, bufs=None, stream=None): + stream = 0 if stream is None else stream + T, H, K, HV = *k.shape, g.shape[1] + BT = chunk_size + # Fast path: for BT >= 64 use larger + # BC=32 / BK=64 sub-tiles and route through the SAFE_GATE matmul branch + # instead of the BC-iteration scalar `for j` loops. The scalar path is the + # bwd bottleneck. + use_fast_path = BT >= 64 + BC = 32 if use_fast_path else min(16, BT) + BK = min(64 if use_fast_path else 32, next_power_of_2(K)) + safe_gate = safe_gate or use_fast_path + NT = len(chunk_indices) + NC = cdiv(BT, BC) + NK = cdiv(K, BK) + + dq2 = bufs["dq2"].reshape((T, HV, K)) + dk2 = bufs["dk2"].reshape((T, HV, K)) + db2 = bufs["db2"].reshape((NK, T, HV)) + dg2 = bufs["dg2"].reshape((T, HV, K)) + + cu_arg = cu_seqlens.reshape(-1) + ci_arg = chunk_indices.reshape(-1) + grid = (NK * NC, NT, HV) + # Kernel uses flat element-offset gather/scatter; pass 1-D views. + ct.launch( + stream, + grid, + chunk_kda_bwd_kernel_intra, + ( + q.reshape((-1,)), + k.reshape((-1,)), + g.reshape(-1), + beta.reshape((-1,)), + dAqk.reshape(-1), + dAkk.reshape(-1), + dq.reshape(-1), + dq2.reshape(-1), + dk.reshape(-1), + dk2.reshape(-1), + dg.reshape(-1), + dg2.reshape(-1), + db2.reshape(-1), + cu_arg, + ci_arg, + T, + H, + HV, + K, + BT, + BC, + BK, + NC, + int(safe_gate), + ), + ) + dq = dq2 + dk = dk2 + # dBeta += sum_nk dBeta2 (fp32 acc); the fan-in NK is a compile-time constant + sum_leading(db.reshape((T * HV,)), db2.reshape((NK, T * HV)), NK, T * HV, stream=stream, accumulate=True) + dg = dg2 + return dq, dk, db, dg + + +# --- Pipelines ------------------------------------------------------------------------------------ def chunk_kda_fwd( @@ -3729,8 +3246,8 @@ def chunk_kda_bwd( # For GVA, reduce dQ and dK from [T, HV, K] back to [T, H, K] if HV > H: T_, K_ = dq.shape[0], dq.shape[-1] - dq_r = reshaped(bufs["dq_hred"], (T_, H, K_)) - dk_r = reshaped(bufs["dk_hred"], (T_, H, K_)) + dq_r = bufs["dq_hred"].reshape((T_, H, K_)) + dk_r = bufs["dk_hred"].reshape((T_, H, K_)) head_group_sum(dq_r, dq, T_, H, G, K_, stream=stream) head_group_sum(dk_r, dk, T_, H, G, K_, stream=stream) dq, dk = dq_r, dk_r @@ -3754,11 +3271,6 @@ def chunk_kda_bwd( return dq, dk, dv, db, dg, dh0, dA, dbias -# =========================================================================== -# ChunkKDAFunction + chunk_kda -# =========================================================================== - - def chunk_kda_grad( q, k, @@ -3866,8 +3378,8 @@ def chunk_kda_grad( stream=stream, ) if use_qk_l2norm_in_kernel: - dq = l2norm_bwd(q_in, q_rstd, dq, out=bufs["dq_l2"], stream=stream) - dk = l2norm_bwd(k_in, k_rstd, dk, out=bufs["dk_l2"], stream=stream) + dq = l2norm_bwd(q_in, q_rstd, dq, out=bufs["dq_l2"], bufs=bufs, stream=stream) + dk = l2norm_bwd(k_in, k_rstd, dk, out=bufs["dk_l2"], bufs=bufs, stream=stream) if use_beta_sigmoid_in_kernel: db = fused_beta_sigmoid_bwd(beta_raw, db, scale=2.0 if allow_neg_eigval else 1.0, out=db, stream=stream) return ( diff --git a/python/cudnn/linear_attention/frost/common/downcast.py b/python/cudnn/linear_attention/frost/common/downcast.py index f6b19a9a4..0bb77d1c4 100644 --- a/python/cudnn/linear_attention/frost/common/downcast.py +++ b/python/cudnn/linear_attention/frost/common/downcast.py @@ -3,28 +3,38 @@ """Initial-state staging for the FROST LA backward kernels: copy the caller's ``[N, HO, K, V]`` state (fp32 or io dtype, padded outer strides fine) into -the compact io-dtype buffer the per-(b,h) state descriptors read.""" +the compact io-dtype buffer the per-(b,h) state descriptors read. + +Alignment is the caller's contract, as everywhere TMA is involved: 16-byte +aligned bases, and outer strides that keep every 8-element V chunk address +16-byte aligned (compact buffers trivially qualify).""" import functools import cuda.bindings.driver as cuda import cutlass import cutlass.cute as cute +from cutlass.cute.arch.nvvm_wrappers import inline_ptx from cutlass.cute.runtime import from_dlpack +from cudnn.frost.tile_dsl.pointwise import f16x2_to_f32, fp32_to_fp16 + @cute.kernel -def downcast_state_f16_kernel( +def downcast_state_kernel( mState0: cute.Tensor, mOut: cute.Tensor, n_k: cutlass.Int32, threads_per_row: cutlass.Int32, rows_per_cta: cutlass.Int32, ) -> None: - """Row-chunk copy of the ``[N, HO, K, V]`` initial state into the io-dtype - buffer the backward's static state descriptor reads: grid (K-tiles, HO, N), - one 8-element V chunk per thread, source read through its (dynamic) - strides so padded outer layouts stage zero-copy.""" + """Vectorized copy of the ``[N, HO, K, V]`` initial state into the + io-dtype buffer the backward's static state descriptor reads: grid + (K-tiles, HO, N), one 8-element V chunk per thread as 128-bit loads and + one 128-bit store, source read through its (dynamic) strides so padded + outer layouts stage zero-copy. fp32 sources convert through packed + ``cvt.rn.{f16,bf16}x2.f32``; same-dtype io sources copy words verbatim; + a 16-bit cross convert unpacks to fp32 pairs and repacks.""" bid = cute.arch.block_idx() tidx = cutlass.Int32(cute.arch.thread_idx()[0]) k_idx = cutlass.Int32(bid[0]) * rows_per_cta + tidx // threads_per_row @@ -32,12 +42,59 @@ def downcast_state_f16_kernel( n_idx = cutlass.Int32(bid[2]) h_idx = cutlass.Int32(bid[1]) if k_idx < n_k: - for i in cutlass.range_constexpr(8): - mOut[n_idx, h_idx, k_idx, v0 + i] = mState0[n_idx, h_idx, k_idx, v0 + i].to(mOut.element_type) + src_elems = ( + cutlass.Int64(n_idx) * cutlass.Int64(mState0.stride[0]) + + cutlass.Int64(h_idx) * cutlass.Int64(mState0.stride[1]) + + cutlass.Int64(k_idx) * cutlass.Int64(mState0.stride[2]) + + cutlass.Int64(v0) + ) + dst_elems = ( + cutlass.Int64(n_idx) * cutlass.Int64(mOut.stride[0]) + + cutlass.Int64(h_idx) * cutlass.Int64(mOut.stride[1]) + + cutlass.Int64(k_idx) * cutlass.Int64(mOut.stride[2]) + + cutlass.Int64(v0) + ) + dst_addr = mOut.iterator.toint() + dst_elems * cutlass.Int64(2) + if cutlass.const_expr(mState0.element_type == cutlass.Float32): + src_addr = mState0.iterator.toint() + src_elems * cutlass.Int64(4) + f0, f1, f2, f3 = inline_ptx( + "ld.global.v4.f32 {$0, $1, $2, $3}, [$4];", + write_only_types=[cutlass.Float32, cutlass.Float32, cutlass.Float32, cutlass.Float32], + read_only_args=[src_addr], + ) + f4, f5, f6, f7 = inline_ptx( + "ld.global.v4.f32 {$0, $1, $2, $3}, [$4];", + write_only_types=[cutlass.Float32, cutlass.Float32, cutlass.Float32, cutlass.Float32], + read_only_args=[src_addr + cutlass.Int64(16)], + ) + w0 = fp32_to_fp16(f0, f1, dtype=mOut.element_type) + w1 = fp32_to_fp16(f2, f3, dtype=mOut.element_type) + w2 = fp32_to_fp16(f4, f5, dtype=mOut.element_type) + w3 = fp32_to_fp16(f6, f7, dtype=mOut.element_type) + else: + src_addr = mState0.iterator.toint() + src_elems * cutlass.Int64(2) + w0, w1, w2, w3 = inline_ptx( + "ld.global.v4.b32 {$0, $1, $2, $3}, [$4];", + write_only_types=[cutlass.Int32, cutlass.Int32, cutlass.Int32, cutlass.Int32], + read_only_args=[src_addr], + ) + if cutlass.const_expr(mState0.element_type != mOut.element_type): + lo0, hi0 = f16x2_to_f32(w0, dtype=mState0.element_type) + lo1, hi1 = f16x2_to_f32(w1, dtype=mState0.element_type) + lo2, hi2 = f16x2_to_f32(w2, dtype=mState0.element_type) + lo3, hi3 = f16x2_to_f32(w3, dtype=mState0.element_type) + w0 = fp32_to_fp16(lo0, hi0, dtype=mOut.element_type) + w1 = fp32_to_fp16(lo1, hi1, dtype=mOut.element_type) + w2 = fp32_to_fp16(lo2, hi2, dtype=mOut.element_type) + w3 = fp32_to_fp16(lo3, hi3, dtype=mOut.element_type) + inline_ptx( + "st.global.v4.b32 [$0], {$1, $2, $3, $4};", + read_only_args=[dst_addr, w0, w1, w2, w3], + ) @cute.jit -def downcast_state_f16( +def downcast_state_launch( state0: cute.Tensor, out: cute.Tensor, n_k: cutlass.Int32, @@ -48,7 +105,7 @@ def downcast_state_f16( n_seq: cutlass.Int32, stream: cuda.CUstream, ): - downcast_state_f16_kernel( + downcast_state_kernel( state0, out, n_k, @@ -65,9 +122,7 @@ def downcast_state_cache(key): def downcast_state(initial_state, out, *, stream): """Copy the initial state ``[N, HO, K, V]`` (fp32 or io dtype, padded outer strides fine) into ``out`` (io dtype, same shape, compact) — the - buffer the backward's per-(b,h) state descriptors read. Stride-aware: - the source is read through its own layout, never reshaped or copied - host-side.""" + buffer the backward's per-(b,h) state descriptors read.""" if tuple(int(s_) for s_ in initial_state.shape) != tuple(int(s_) for s_ in out.shape): raise ValueError(f"initial_state must match the io state buffer shape {tuple(out.shape)}; got {tuple(initial_state.shape)}") n_seq, ho, k, v = (int(s_) for s_ in out.shape) @@ -85,7 +140,7 @@ def downcast_state(initial_state, out, *, stream): state0_c = from_dlpack(initial_state, assumed_align=16).mark_layout_dynamic(leading_dim=3) out_c = from_dlpack(out, assumed_align=16).mark_layout_dynamic(leading_dim=3) cache["compiled"] = cute.compile( - downcast_state_f16, + downcast_state_launch, state0_c, out_c, cutlass.Int32(k), diff --git a/python/cudnn/linear_attention/frost/common/elementwise.py b/python/cudnn/linear_attention/frost/common/elementwise.py new file mode 100644 index 000000000..ad1b5c04f --- /dev/null +++ b/python/cudnn/linear_attention/frost/common/elementwise.py @@ -0,0 +1,48 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Element-wise device helpers shared by the FROST LA kernels: a lane-group +butterfly reduction, the inverse L2 norm with its epsilon floor, and the +sigmoid / softplus activations behind the safe gate and beta.""" + +import cutlass +import cutlass.cute as cute +import cutlass.experimental.primitives as nvvm + +L2_NORM_EPS = 1.0e-12 + + +@cute.jit +def lane_group_sum(value: cutlass.Float32, lanes: cutlass.Constexpr[int]) -> cutlass.Float32: + """Sum ``value`` across a power-of-two group of consecutive lanes via + butterfly shuffles (every lane ends up holding the group total).""" + offset = lanes // 2 + while offset >= 1: + value = value + cutlass.Float32(nvvm.shfl_sync(0xFFFFFFFF, value, offset, 31, kind=nvvm.Shfl.BFLY)) + offset = offset // 2 + return value + + +@cute.jit +def l2norm_inv(sum_sq: cutlass.Float32) -> cutlass.Float32: + """Inverse L2 norm with the shared epsilon floor: rows at or below the + floor normalize by ``1 / L2_NORM_EPS`` instead of dividing by zero.""" + norm_floor_sq = cutlass.Float32(L2_NORM_EPS * L2_NORM_EPS) + return cute.math.rsqrt(cute.math.max(sum_sq, norm_floor_sq), fastmath=True) + + +@cute.jit +def sigmoid(x: cutlass.Float32) -> cutlass.Float32: + """sigmoid(x) via the tanh identity (single MUFU on Blackwell).""" + half = cutlass.Float32(0.5) + return cute.math.tanh(x * half, approx=True) * half + half + + +@cute.jit +def softplus(x: cutlass.Float32) -> cutlass.Float32: + """log(1 + exp(x)) with the linear tail (x > 20 returns x: exp saturates + fp32 there and log1p(exp(x)) == x to fp32 precision).""" + result = x + if x < cutlass.Float32(20.0): + result = cute.math.log(cutlass.Float32(1.0) + cute.math.exp(x, fastmath=True), fastmath=True) + return result diff --git a/python/cudnn/linear_attention/frost/common/gate_bwd.py b/python/cudnn/linear_attention/frost/common/gate_bwd.py new file mode 100644 index 000000000..e86550385 --- /dev/null +++ b/python/cudnn/linear_attention/frost/common/gate_bwd.py @@ -0,0 +1,403 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Safe-gate backward helper: convert the main backward kernels' dGate +(gradient wrt the transformed log-decay) into the raw-logit gradient in +place, and produce the gate-parameter gradients dA_log / ddt_bias. + +Two forms: + - scalar (GDN): g = -exp(A_log[h]) * softplus(g_raw + dt_bias[h]) + - per-channel (KDA): g = lb * sigmoid(exp(a_log) * (g_raw + dt_bias)) + +The parameter gradients are cross-token sums, done deterministically: every +fp32 accumulator owns a statically assigned (head[, channel], token-slice) +and partials are combined through fixed-shape trees only — the bracketing is +a pure function of the shapes, never of scheduling, so results are bitwise +stable run to run. Two stages with a launch boundary as the grid barrier: +a partial pass over token stripes (the same pass rewrites dGate -> dg_raw), +then a finisher that folds the stripe partials. + +Scalar stripes scale with the token axis (:func:`scalar_gate_blocks`, capped) +so the partial pass fills the machine; channel stripes stay at +``GATE_BWD_BLOCKS`` (the partials carve is ``128 * HO * 128`` fp32). +""" + +import functools + +import cuda.bindings.driver as cuda +import cutlass +import cutlass.cute as cute +import cutlass.experimental.primitives as nvvm +from cutlass.cute.arch.nvvm_wrappers import inline_ptx +from cutlass.cute.runtime import from_dlpack + +from cudnn.frost.buffers import data_ptr +from cudnn.frost.tile_dsl.pointwise import fadd2, ffma2, fmul2 + +from .elementwise import lane_group_sum, sigmoid, softplus + +GATE_BWD_BLOCKS = 128 # channel-gate token stripes (partials carve = 128 * HO * 128 fp32) +SCALAR_BLOCK_CAP = 8192 # scalar-gate stripe ceiling +SCALAR_SLICE_TOKENS = 16 # scalar-gate target tokens per stripe +SCALAR_HEAD_TILE = 32 # scalar-gate heads per block: one warp, tiled over grid.y + + +def scalar_gate_blocks(n_tokens: int) -> int: + """Scalar-gate stripe count: shape-only (so the summation bracketing is + deterministic) and scaling with the token axis so the partial pass runs + memory-bound instead of 128 * HO threads walking long serial slices.""" + return min(SCALAR_BLOCK_CAP, max(1, -(-n_tokens // SCALAR_SLICE_TOKENS))) + + +@cute.kernel +def scalar_gate_bwd_partial_kernel( + mDg: cute.Tensor, + mG: cute.Tensor, + mALog: cute.Tensor, + mDtBias: cute.Tensor, + mPartA: cute.Tensor, + mPartDt: cute.Tensor, + n_tokens: cutlass.Int32, + h_o: cutlass.Int32, + slice_len: cutlass.Int32, +) -> None: + """Grid (stripes, head tiles), block (SCALAR_HEAD_TILE,): thread h walks + stripe g's token slice in order, rewriting dGate -> dg_raw in place and + accumulating the (g, h) partials for dA_log (dGate * g_transformed) and + ddt_bias (dg_raw). Tiling heads over grid.y rather than one thread per + head takes any HO and keeps a small HO down to a single warp; each + (stripe, head) still has exactly one owner, so the partial layout and the + finisher's summation order are unchanged.""" + bid = cute.arch.block_idx() + tidx = cutlass.Int32(cute.arch.thread_idx()[0]) + g_blk = cutlass.Int32(bid[0]) + h = cutlass.Int32(bid[1]) * cutlass.Int32(SCALAR_HEAD_TILE) + tidx + if h < h_o: + neg_exp_a = -cute.math.exp(mALog[h], fastmath=True) + bias = mDtBias[h] + t = g_blk * slice_len + t_end = t + slice_len + if t_end > n_tokens: + t_end = n_tokens + acc_a = cutlass.Float32(0.0) + acc_dt = cutlass.Float32(0.0) + while t < t_end: + d_gate = mDg[t, h] + y = mG[t, h] + bias + dg_raw = d_gate * neg_exp_a * sigmoid(y) + acc_a += d_gate * (neg_exp_a * softplus(y)) + acc_dt += dg_raw + mDg[t, h] = dg_raw + t += 1 + mPartA[g_blk * h_o + h] = acc_a + mPartDt[g_blk * h_o + h] = acc_dt + + +@cute.kernel +def scalar_gate_bwd_finish_kernel( + mPartA: cute.Tensor, + mPartDt: cute.Tensor, + mDA: cute.Tensor, + mDDt: cute.Tensor, + h_o: cutlass.Int32, + n_blocks: cutlass.Int32, +) -> None: + """Grid (HO,), block (32,): head h's stripe partials fold as 8 fixed + interleaved chains per lane (independent accumulators so the loads + pipeline instead of serializing at memory latency), then the 8 chains + pairwise and one butterfly tree across the 32 lanes — a fixed-shape + bracketing regardless of the stripe count.""" + bid = cute.arch.block_idx() + lane = cutlass.Int32(cute.arch.thread_idx()[0]) + h = cutlass.Int32(bid[0]) + a8 = cutlass.Array(cutlass.Float32, 8) + d8 = cutlass.Array(cutlass.Float32, 8) + for j in cutlass.range_constexpr(8): + a8[j] = cutlass.Float32(0.0) + d8[j] = cutlass.Float32(0.0) + s = lane + while s + cutlass.Int32(224) < n_blocks: + for j in cutlass.range_constexpr(8): + idx = (s + cutlass.Int32(32 * j)) * h_o + h + a8[j], d8[j] = fadd2(a8[j], d8[j], mPartA[idx], mPartDt[idx]) + s += cutlass.Int32(256) + while s < n_blocks: + idx = s * h_o + h + a8[0], d8[0] = fadd2(a8[0], d8[0], mPartA[idx], mPartDt[idx]) + s += cutlass.Int32(32) + # the A and dt chains are independent, so each packed add folds both + p0a, p0d = fadd2(a8[0], d8[0], a8[1], d8[1]) + p1a, p1d = fadd2(a8[2], d8[2], a8[3], d8[3]) + p2a, p2d = fadd2(a8[4], d8[4], a8[5], d8[5]) + p3a, p3d = fadd2(a8[6], d8[6], a8[7], d8[7]) + q0a, q0d = fadd2(p0a, p0d, p1a, p1d) + q1a, q1d = fadd2(p2a, p2d, p3a, p3d) + acc_a, acc_dt = fadd2(q0a, q0d, q1a, q1d) + acc_a = lane_group_sum(acc_a, 32) + acc_dt = lane_group_sum(acc_dt, 32) + if lane == 0: + mDA[h] = acc_a + mDDt[h] = acc_dt + + +@cute.kernel +def channel_gate_bwd_partial_kernel( + mDg: cute.Tensor, + mG: cute.Tensor, + mALog: cute.Tensor, + mDtBias: cute.Tensor, + mPartA: cute.Tensor, + mPartDt: cute.Tensor, + n_tokens: cutlass.Int32, + h_o: cutlass.Int32, + slice_len: cutlass.Int32, + lower_bound: cutlass.Float32, +) -> None: + """Grid (GATE_BWD_BLOCKS, HO), block (128,) = 4 token-phased warps of 32 + lanes; lane owns channels [lane*4, lane*4 + 4) through 128-bit loads and + stores. Warp w walks tokens t0 + w, t0 + w + 4, ... of the stripe: with + z = exp(a_log) * (g_raw + dt_bias) and w_ = dGate * lb * sig(z)(1-sig(z)), + dg_raw = w_ * exp(a_log) (rewritten in place), the dA_log integrand is + w_ * z, the ddt_bias one dg_raw. The four warp partials per channel fold + w = 0..3 through SMEM (fixed order) into the (g, h) partial slots. + a_log is one scalar per head; dt_bias is per (head, channel).""" + bid = cute.arch.block_idx() + tidx = cutlass.Int32(cute.arch.thread_idx()[0]) + g_blk = cutlass.Int32(bid[0]) + h = cutlass.Int32(bid[1]) + wrp = tidx // cutlass.Int32(32) + lane = tidx % cutlass.Int32(32) + d0 = lane * cutlass.Int32(4) + exp_a = cute.math.exp(mALog[h], fastmath=True) + bias = cutlass.Array(cutlass.Float32, 4) + for q in cutlass.range_constexpr(4): + bias[q] = mDtBias[h, d0 + cutlass.Int32(q)] + acc_a = cutlass.Array(cutlass.Float32, 4) + acc_dt = cutlass.Array(cutlass.Float32, 4) + out = cutlass.Array(cutlass.Float32, 4) + for q in cutlass.range_constexpr(4): + acc_a[q] = cutlass.Float32(0.0) + acc_dt[q] = cutlass.Float32(0.0) + dg_base = mDg.iterator.toint() + g_base = mG.iterator.toint() + dg_s0 = cutlass.Int64(mDg.stride[0]) + dg_s1 = cutlass.Int64(mDg.stride[1]) + g_s0 = cutlass.Int64(mG.stride[0]) + g_s1 = cutlass.Int64(mG.stride[1]) + t = g_blk * slice_len + wrp + t_end = g_blk * slice_len + slice_len + if t_end > n_tokens: + t_end = n_tokens + while t < t_end: + dg_addr = dg_base + (cutlass.Int64(t) * dg_s0 + cutlass.Int64(h) * dg_s1 + cutlass.Int64(d0)) * cutlass.Int64(4) + g_addr = g_base + (cutlass.Int64(t) * g_s0 + cutlass.Int64(h) * g_s1 + cutlass.Int64(d0)) * cutlass.Int64(4) + dg0, dg1, dg2, dg3 = inline_ptx( + "ld.global.v4.f32 {$0, $1, $2, $3}, [$4];", + write_only_types=[cutlass.Float32, cutlass.Float32, cutlass.Float32, cutlass.Float32], + read_only_args=[dg_addr], + ) + gv0, gv1, gv2, gv3 = inline_ptx( + "ld.global.v4.f32 {$0, $1, $2, $3}, [$4];", + write_only_types=[cutlass.Float32, cutlass.Float32, cutlass.Float32, cutlass.Float32], + read_only_args=[g_addr], + ) + dgv = (dg0, dg1, dg2, dg3) + gvv = (gv0, gv1, gv2, gv3) + for p in cutlass.range_constexpr(2): + i = 2 * p + j = i + 1 + one = cutlass.Float32(1.0) + y_lo, y_hi = fadd2(gvv[i], gvv[j], bias[i], bias[j]) + z_lo, z_hi = fmul2(exp_a, exp_a, y_lo, y_hi) + sig_lo = sigmoid(z_lo) + sig_hi = sigmoid(z_hi) + c_lo, c_hi = fadd2(one, one, -sig_lo, -sig_hi) + k_lo, k_hi = fmul2(sig_lo, sig_hi, c_lo, c_hi) + b_lo, b_hi = fmul2(dgv[i], dgv[j], lower_bound, lower_bound) + w_lo, w_hi = fmul2(b_lo, b_hi, k_lo, k_hi) + raw_lo, raw_hi = fmul2(w_lo, w_hi, exp_a, exp_a) + acc_a[i], acc_a[j] = ffma2(w_lo, w_hi, z_lo, z_hi, acc_a[i], acc_a[j]) + acc_dt[i], acc_dt[j] = fadd2(acc_dt[i], acc_dt[j], raw_lo, raw_hi) + out[i] = raw_lo + out[j] = raw_hi + inline_ptx( + "st.global.v4.f32 [$0], {$1, $2, $3, $4};", + read_only_args=[dg_addr, out[0], out[1], out[2], out[3]], + ) + t += cutlass.Int32(4) + sA = cutlass.Array(cutlass.Float32, 512, space=cutlass.AddressSpace.smem, alignment=16) + sDt = cutlass.Array(cutlass.Float32, 512, space=cutlass.AddressSpace.smem, alignment=16) + for q in cutlass.range_constexpr(4): + sA[wrp * cutlass.Int32(128) + d0 + cutlass.Int32(q)] = acc_a[q] + sDt[wrp * cutlass.Int32(128) + d0 + cutlass.Int32(q)] = acc_dt[q] + nvvm.barrier_cta_sync() + if wrp == 0: + for q in cutlass.range_constexpr(4): + d = d0 + cutlass.Int32(q) + fa = ((sA[d] + sA[cutlass.Int32(128) + d]) + sA[cutlass.Int32(256) + d]) + sA[cutlass.Int32(384) + d] + fdt = ((sDt[d] + sDt[cutlass.Int32(128) + d]) + sDt[cutlass.Int32(256) + d]) + sDt[cutlass.Int32(384) + d] + base = (g_blk * h_o + h) * cutlass.Int32(128) + d + mPartA[base] = fa + mPartDt[base] = fdt + + +@cute.kernel +def channel_gate_bwd_finish_kernel( + mPartA: cute.Tensor, + mPartDt: cute.Tensor, + mDA: cute.Tensor, + mDDt: cute.Tensor, + h_o: cutlass.Int32, +) -> None: + """Grid (HO,), block (128,): thread d folds its channel column over the + GATE_BWD_BLOCKS stripe partials as 8 fixed interleaved chains + (independent accumulators so the loads pipeline) folded pairwise. + ddt_bias is per (head, channel), so thread d stores its column outright; + dA_log is per head, so the channel axis folds through a fixed tree (warp + butterflies, then the 4 warp sums in index order).""" + bid = cute.arch.block_idx() + tidx = cutlass.Int32(cute.arch.thread_idx()[0]) + h = cutlass.Int32(bid[0]) + d = tidx + wrp = tidx // cutlass.Int32(32) + lane = tidx % cutlass.Int32(32) + a8 = cutlass.Array(cutlass.Float32, 8) + d8 = cutlass.Array(cutlass.Float32, 8) + for j in cutlass.range_constexpr(8): + a8[j] = cutlass.Float32(0.0) + d8[j] = cutlass.Float32(0.0) + c = cutlass.Int32(0) + while c < cutlass.Int32(GATE_BWD_BLOCKS // 8): + for j in cutlass.range_constexpr(8): + idx = ((c * cutlass.Int32(8) + cutlass.Int32(j)) * h_o + h) * cutlass.Int32(128) + d + a8[j] = a8[j] + mPartA[idx] + d8[j] = d8[j] + mPartDt[idx] + c += 1 + col_a = ((a8[0] + a8[1]) + (a8[2] + a8[3])) + ((a8[4] + a8[5]) + (a8[6] + a8[7])) + col_dt = ((d8[0] + d8[1]) + (d8[2] + d8[3])) + ((d8[4] + d8[5]) + (d8[6] + d8[7])) + mDDt[h, d] = col_dt + sWa = cutlass.Array(cutlass.Float32, 4, space=cutlass.AddressSpace.smem, alignment=16) + va = lane_group_sum(col_a, 32) + if lane == 0: + sWa[wrp] = va + nvvm.barrier_cta_sync() + if tidx == 0: + mDA[h] = ((sWa[0] + sWa[1]) + sWa[2]) + sWa[3] + + +@cute.jit +def scalar_gate_bwd_launch( + d_gate: cute.Tensor, + g_raw: cute.Tensor, + a_log: cute.Tensor, + dt_bias: cute.Tensor, + part_a: cute.Tensor, + part_dt: cute.Tensor, + d_a_log: cute.Tensor, + d_dt_bias: cute.Tensor, + n_tokens: cutlass.Int32, + h_o: cutlass.Int32, + slice_len: cutlass.Int32, + n_blocks: cutlass.Int32, + head_tiles: cutlass.Int32, + stream: cuda.CUstream, +): + scalar_gate_bwd_partial_kernel(d_gate, g_raw, a_log, dt_bias, part_a, part_dt, n_tokens, h_o, slice_len).launch( + grid=(n_blocks, head_tiles, 1), block=(SCALAR_HEAD_TILE, 1, 1), stream=stream + ) + scalar_gate_bwd_finish_kernel(part_a, part_dt, d_a_log, d_dt_bias, h_o, n_blocks).launch(grid=(h_o, 1, 1), block=(32, 1, 1), stream=stream) + + +@cute.jit +def channel_gate_bwd_launch( + d_gate: cute.Tensor, + g_raw: cute.Tensor, + a_log: cute.Tensor, + dt_bias: cute.Tensor, + part_a: cute.Tensor, + part_dt: cute.Tensor, + d_a_log: cute.Tensor, + d_dt_bias: cute.Tensor, + n_tokens: cutlass.Int32, + h_o: cutlass.Int32, + slice_len: cutlass.Int32, + lower_bound: cutlass.Float32, + stream: cuda.CUstream, +): + channel_gate_bwd_partial_kernel(d_gate, g_raw, a_log, dt_bias, part_a, part_dt, n_tokens, h_o, slice_len, lower_bound).launch( + grid=(GATE_BWD_BLOCKS, h_o, 1), block=(128, 1, 1), stream=stream + ) + channel_gate_bwd_finish_kernel(part_a, part_dt, d_a_log, d_dt_bias, h_o).launch(grid=(h_o, 1, 1), block=(128, 1, 1), stream=stream) + + +@functools.cache +def gate_bwd_cache(key): + return {} + + +def scalar_gate_bwd(d_gate, g_raw, a_log, dt_bias, d_a_log, d_dt_bias, part_a, part_dt, *, stream): + """Scalar-gate backward: rewrite d_gate [total, HO] fp32 in place from + transformed-space to raw-logit space and fill d_a_log/d_dt_bias (HO,). + part_a/part_dt are (scalar_gate_blocks(total) * HO,) fp32 workspace carves.""" + n_tokens, h_o = (int(s_) for s_ in d_gate.shape) + n_blocks = scalar_gate_blocks(n_tokens) + head_tiles = -(-h_o // SCALAR_HEAD_TILE) + slice_len = (n_tokens + n_blocks - 1) // n_blocks + for name, buf in (("part_a", part_a), ("part_dt", part_dt)): + if int(buf.shape[0]) < n_blocks * h_o: + raise ValueError(f"{name} must hold scalar_gate_blocks({n_tokens}) * {h_o} = {n_blocks * h_o} fp32; got {int(buf.shape[0])}") + cache = gate_bwd_cache(("gdn",)) + cu_stream = cuda.CUstream(int(stream)) + tensors = (d_gate, g_raw, a_log, dt_bias, part_a, part_dt, d_a_log, d_dt_bias) + if "compiled" not in cache: + traced = [from_dlpack(t, assumed_align=4) for t in tensors[:2]] + traced = [tr.mark_layout_dynamic(leading_dim=1) for tr in traced] + traced += [from_dlpack(t, assumed_align=4).mark_layout_dynamic(leading_dim=0) for t in tensors[2:]] + cache["compiled"] = cute.compile( + scalar_gate_bwd_launch, + *traced, + cutlass.Int32(n_tokens), + cutlass.Int32(h_o), + cutlass.Int32(slice_len), + cutlass.Int32(n_blocks), + cutlass.Int32(head_tiles), + cu_stream, + options="--enable-tvm-ffi", + ) + cache["compiled"](*tensors, n_tokens, h_o, slice_len, n_blocks, head_tiles, cu_stream) + + +def channel_gate_bwd(d_gate, g_raw, a_log, dt_bias, d_a_log, d_dt_bias, part_a, part_dt, gate_lower_bound, *, stream): + """Per-channel-gate backward: rewrite d_gate [total, HO, 128] fp32 in + place and fill d_a_log/d_dt_bias at their parameter shapes ((HO,) params + get the channel axis folded in the finisher).""" + n_tokens, h_o, d_k = (int(s_) for s_ in d_gate.shape) + if d_k != 128: + raise ValueError(f"per-channel gate backward requires 128 channels; got {d_k}") + for name, t in (("d_gate", d_gate), ("g_raw", g_raw)): + st = tuple(int(s_) for s_ in t.stride()) + if st[2] != 1 or st[0] % 4 != 0 or st[1] % 4 != 0 or data_ptr(t) % 16 != 0: + raise ValueError(f"{name} needs a 16B-aligned base, unit channel stride, and outer strides in multiples of 4; got {st}") + for name, t in (("a_log", a_log), ("d_a_log", d_a_log)): + if tuple(int(s_) for s_ in t.shape) != (h_o,) or tuple(int(s_) for s_ in t.stride()) != (1,): + raise ValueError(f"{name} must be a contiguous ({h_o},) per-head fp32 parameter; got shape {tuple(t.shape)}") + for name, t in (("dt_bias", dt_bias), ("d_dt_bias", d_dt_bias)): + if tuple(int(s_) for s_ in t.shape) != (h_o, d_k) or tuple(int(s_) for s_ in t.stride()) != (d_k, 1): + raise ValueError(f"{name} must be a contiguous ({h_o}, {d_k}) per-channel fp32 parameter; got shape {tuple(t.shape)}") + slice_len = (n_tokens + GATE_BWD_BLOCKS - 1) // GATE_BWD_BLOCKS + cache = gate_bwd_cache(("channel",)) + cu_stream = cuda.CUstream(int(stream)) + tensors = (d_gate, g_raw, a_log, dt_bias, part_a, part_dt, d_a_log, d_dt_bias) + args = (n_tokens, h_o, slice_len) + if "compiled" not in cache: + traced = [from_dlpack(t, assumed_align=16).mark_layout_dynamic(leading_dim=2) for t in tensors[:2]] + traced += [from_dlpack(t, assumed_align=4).mark_layout_dynamic(leading_dim=len(t.shape) - 1) for t in tensors[2:]] + cache["compiled"] = cute.compile( + channel_gate_bwd_launch, + *traced, + *(cutlass.Int32(a) for a in args), + cutlass.Float32(gate_lower_bound), + cu_stream, + options="--enable-tvm-ffi", + ) + cache["compiled"](*tensors, *args, float(gate_lower_bound), cu_stream) diff --git a/python/cudnn/linear_attention/frost/common/l2norm.py b/python/cudnn/linear_attention/frost/common/l2norm.py new file mode 100644 index 000000000..795662ed0 --- /dev/null +++ b/python/cudnn/linear_attention/frost/common/l2norm.py @@ -0,0 +1,331 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Q/K row L2 normalization helpers for GDN (the main kernels stay unchanged +and consume normalized workspace copies through their usual descriptors). + +Forward: normalize every 128-element q/k row into compact io-dtype workspace +buffers and stash the fp32 inverse norms. Backward: project the main +kernel's dq_n/dk_n (gradients wrt the normalized rows) back through the +normalize Jacobian in place: ``dq = inv_norm * (dq_n - (dq_n . q_n) q_n)``. +""" + +import functools + +import cuda.bindings.driver as cuda +import cutlass +import cutlass.cute as cute +from cutlass.cute.arch.nvvm_wrappers import inline_ptx +from cutlass.cute.runtime import from_dlpack + +from cudnn.frost.buffers import data_ptr +from cudnn.frost.tile_dsl.pointwise import f16x2_to_f32, ffma2, fmul2, fp32_to_fp16 +from .elementwise import l2norm_inv, lane_group_sum + +THREADS_PER_ROW = 16 +ROWS_PER_CTA = 8 +FWD_LANES = 4 # fwd lanes per row: 32 elems/lane (4 x 128-bit loads), 2-step butterfly +FWD_ROWS_PER_GROUP = 2 # fwd rows batched per lane group: 8 independent loads in flight per thread + + +@cute.kernel +def l2norm_qk_kernel( + mQ: cute.Tensor, + mK: cute.Tensor, + mQn: cute.Tensor, + mKn: cute.Tensor, + mInvQ: cute.Tensor, + mInvK: cute.Tensor, + n_q_rows: cutlass.Int32, + n_rows: cutlass.Int32, + h_q: cutlass.Int32, + h_k: cutlass.Int32, +) -> None: + """Grid over all q rows then all k rows, FWD_LANES lanes x 32 elements + per row (4 x 128-bit loads per lane), FWD_ROWS_PER_GROUP consecutive rows + batched per lane group. All rows' loads issue before the reductions and + the 4-lane butterfly is 2 steps, so memory latency pipelines (the + 16-lane single-row shape measured 75.7% of DRAM peak on GB200, this one + 88.3%; FLA's autotuned Triton tile is 91%). fp32 sums of squares, + rsqrt with the shared epsilon floor, normalized rows to the compact io + workspace, fp32 inverse norms to their slots. Tail rows clamp their + loads and skip stores.""" + bid = cute.arch.block_idx() + tidx = cutlass.Int32(cute.arch.thread_idx()[0]) + grp = tidx // cutlass.Int32(FWD_LANES) + lane = tidx % cutlass.Int32(FWD_LANES) + row0 = (cutlass.Int32(bid[0]) * cutlass.Int32(128 // FWD_LANES) + grp) * cutlass.Int32(FWD_ROWS_PER_GROUP) + v0 = lane * cutlass.Int32(32) + rows = [] + ws_addrs = [] + nrm_addrs = [] + vals = [] + for r in cutlass.range_constexpr(FWD_ROWS_PER_GROUP): + row = row0 + cutlass.Int32(r) + row_r = row if row < n_rows else n_rows - cutlass.Int32(1) + # q rows first, then k rows, head-fastest; sources honor their own + # [T, H, 128] strides while workspace rows sit at row * 128. Both + # branches are traced, so the addresses have to exist beforehand. + src_addr = cutlass.Int64(0) + ws_addr = cutlass.Int64(0) + nrm_addr = cutlass.Int64(0) + if row_r < n_q_rows: + t = row_r // h_q + h = row_r % h_q + src_elems = t * cutlass.Int64(mQ.stride[0]) + h * cutlass.Int64(mQ.stride[1]) + cutlass.Int64(v0) + src_addr = mQ.iterator.toint() + src_elems * cutlass.Int64(2) + ws_addr = mQn.iterator.toint() + (cutlass.Int64(row_r) * cutlass.Int64(128) + cutlass.Int64(v0)) * cutlass.Int64(2) + nrm_addr = mInvQ.iterator.toint() + cutlass.Int64(row_r) * cutlass.Int64(4) + else: + k_row = row_r - n_q_rows + t = k_row // h_k + h = k_row % h_k + src_elems = t * cutlass.Int64(mK.stride[0]) + h * cutlass.Int64(mK.stride[1]) + cutlass.Int64(v0) + src_addr = mK.iterator.toint() + src_elems * cutlass.Int64(2) + ws_addr = mKn.iterator.toint() + (cutlass.Int64(k_row) * cutlass.Int64(128) + cutlass.Int64(v0)) * cutlass.Int64(2) + nrm_addr = mInvK.iterator.toint() + cutlass.Int64(k_row) * cutlass.Int64(4) + chunks = [] + for c in cutlass.range_constexpr(4): + w0, w1, w2, w3 = inline_ptx( + "ld.global.v4.b32 {$0, $1, $2, $3}, [$4];", + write_only_types=[cutlass.Int32, cutlass.Int32, cutlass.Int32, cutlass.Int32], + read_only_args=[src_addr + cutlass.Int64(16 * c)], + ) + f0, f1 = f16x2_to_f32(w0, dtype=mQ.element_type) + f2, f3 = f16x2_to_f32(w1, dtype=mQ.element_type) + f4, f5 = f16x2_to_f32(w2, dtype=mQ.element_type) + f6, f7 = f16x2_to_f32(w3, dtype=mQ.element_type) + chunks.append((f0, f1, f2, f3, f4, f5, f6, f7)) + rows.append(row) + ws_addrs.append(ws_addr) + nrm_addrs.append(nrm_addr) + vals.append(chunks) + for r in cutlass.range_constexpr(FWD_ROWS_PER_GROUP): + acc = cutlass.Float32(0.0) + for c in cutlass.range_constexpr(4): + f = vals[r][c] + acc = acc + ((f[0] * f[0] + f[1] * f[1]) + (f[2] * f[2] + f[3] * f[3])) + ((f[4] * f[4] + f[5] * f[5]) + (f[6] * f[6] + f[7] * f[7])) + inv = l2norm_inv(lane_group_sum(acc, FWD_LANES)) + if rows[r] < n_rows: + for c in cutlass.range_constexpr(4): + f = vals[r][c] + s0, s1 = fmul2(f[0], f[1], inv, inv) + s2, s3 = fmul2(f[2], f[3], inv, inv) + s4, s5 = fmul2(f[4], f[5], inv, inv) + s6, s7 = fmul2(f[6], f[7], inv, inv) + w0 = fp32_to_fp16(s0, s1, dtype=mQ.element_type) + w1 = fp32_to_fp16(s2, s3, dtype=mQ.element_type) + w2 = fp32_to_fp16(s4, s5, dtype=mQ.element_type) + w3 = fp32_to_fp16(s6, s7, dtype=mQ.element_type) + inline_ptx( + "st.global.v4.b32 [$0], {$1, $2, $3, $4};", + read_only_args=[ws_addrs[r] + cutlass.Int64(16 * c), w0, w1, w2, w3], + ) + if lane == cutlass.Int32(0): + inline_ptx("st.global.f32 [$0], $1;", read_only_args=[nrm_addrs[r], inv]) + + +@cute.kernel +def l2norm_qk_bwd_kernel( + mDq: cute.Tensor, + mDk: cute.Tensor, + mQn: cute.Tensor, + mKn: cute.Tensor, + mInvQ: cute.Tensor, + mInvK: cute.Tensor, + n_q_rows: cutlass.Int32, + n_rows: cutlass.Int32, + h_q: cutlass.Int32, + h_k: cutlass.Int32, +) -> None: + """In-place normalize-Jacobian projection of dq/dk: per row, fp32 dot of + the incoming gradient with the saved normalized row (butterfly reduce), + then ``inv_norm * (grad - dot * row_n)`` back to the caller's buffer. + Single row per lane group: the dot+projection already fills the memory + latency (92.9% SOL measured; the forward's row batching REGRESSED this + kernel to 86% on GB200, lyris job 2694953).""" + bid = cute.arch.block_idx() + tidx = cutlass.Int32(cute.arch.thread_idx()[0]) + row = cutlass.Int32(bid[0]) * cutlass.Int32(ROWS_PER_CTA) + tidx // cutlass.Int32(THREADS_PER_ROW) + lane = tidx % cutlass.Int32(THREADS_PER_ROW) + if row < n_rows: + v0 = lane * cutlass.Int32(8) + grad_addr = cutlass.Int64(0) + ws_addr = cutlass.Int64(0) + nrm_addr = cutlass.Int64(0) + if row < n_q_rows: + t = row // h_q + h = row % h_q + grad_elems = t * cutlass.Int64(mDq.stride[0]) + h * cutlass.Int64(mDq.stride[1]) + cutlass.Int64(v0) + grad_addr = mDq.iterator.toint() + grad_elems * cutlass.Int64(2) + ws_addr = mQn.iterator.toint() + (cutlass.Int64(row) * cutlass.Int64(128) + cutlass.Int64(v0)) * cutlass.Int64(2) + nrm_addr = mInvQ.iterator.toint() + cutlass.Int64(row) * cutlass.Int64(4) + else: + k_row = row - n_q_rows + t = k_row // h_k + h = k_row % h_k + grad_elems = t * cutlass.Int64(mDk.stride[0]) + h * cutlass.Int64(mDk.stride[1]) + cutlass.Int64(v0) + grad_addr = mDk.iterator.toint() + grad_elems * cutlass.Int64(2) + ws_addr = mKn.iterator.toint() + (cutlass.Int64(k_row) * cutlass.Int64(128) + cutlass.Int64(v0)) * cutlass.Int64(2) + nrm_addr = mInvK.iterator.toint() + cutlass.Int64(k_row) * cutlass.Int64(4) + gw0, gw1, gw2, gw3 = inline_ptx( + "ld.global.v4.b32 {$0, $1, $2, $3}, [$4];", + write_only_types=[cutlass.Int32, cutlass.Int32, cutlass.Int32, cutlass.Int32], + read_only_args=[grad_addr], + ) + d0, d1 = f16x2_to_f32(gw0, dtype=mDq.element_type) + d2, d3 = f16x2_to_f32(gw1, dtype=mDq.element_type) + d4, d5 = f16x2_to_f32(gw2, dtype=mDq.element_type) + d6, d7 = f16x2_to_f32(gw3, dtype=mDq.element_type) + nw0, nw1, nw2, nw3 = inline_ptx( + "ld.global.v4.b32 {$0, $1, $2, $3}, [$4];", + write_only_types=[cutlass.Int32, cutlass.Int32, cutlass.Int32, cutlass.Int32], + read_only_args=[ws_addr], + ) + n0, n1 = f16x2_to_f32(nw0, dtype=mQn.element_type) + n2, n3 = f16x2_to_f32(nw1, dtype=mQn.element_type) + n4, n5 = f16x2_to_f32(nw2, dtype=mQn.element_type) + n6, n7 = f16x2_to_f32(nw3, dtype=mQn.element_type) + dot_lo, dot_hi = fmul2(d0, d1, n0, n1) + dot_lo, dot_hi = ffma2(d2, d3, n2, n3, dot_lo, dot_hi) + dot_lo, dot_hi = ffma2(d4, d5, n4, n5, dot_lo, dot_hi) + dot_lo, dot_hi = ffma2(d6, d7, n6, n7, dot_lo, dot_hi) + dot = lane_group_sum(dot_lo + dot_hi, THREADS_PER_ROW) + inv = inline_ptx( + "ld.global.f32 $0, [$1];", + write_only_types=[cutlass.Float32], + read_only_args=[nrm_addr], + ) + neg_dot = -dot + p0, p1 = ffma2(neg_dot, neg_dot, n0, n1, d0, d1) + p2, p3 = ffma2(neg_dot, neg_dot, n2, n3, d2, d3) + p4, p5 = ffma2(neg_dot, neg_dot, n4, n5, d4, d5) + p6, p7 = ffma2(neg_dot, neg_dot, n6, n7, d6, d7) + q0, q1 = fmul2(p0, p1, inv, inv) + q2, q3 = fmul2(p2, p3, inv, inv) + q4, q5 = fmul2(p4, p5, inv, inv) + q6, q7 = fmul2(p6, p7, inv, inv) + w0 = fp32_to_fp16(q0, q1, dtype=mDq.element_type) + w1 = fp32_to_fp16(q2, q3, dtype=mDq.element_type) + w2 = fp32_to_fp16(q4, q5, dtype=mDq.element_type) + w3 = fp32_to_fp16(q6, q7, dtype=mDq.element_type) + inline_ptx( + "st.global.v4.b32 [$0], {$1, $2, $3, $4};", + read_only_args=[grad_addr, w0, w1, w2, w3], + ) + + +@cute.jit +def l2norm_qk_launch( + q: cute.Tensor, + k: cute.Tensor, + q_n: cute.Tensor, + k_n: cute.Tensor, + inv_q: cute.Tensor, + inv_k: cute.Tensor, + n_q_rows: cutlass.Int32, + n_rows: cutlass.Int32, + h_q: cutlass.Int32, + h_k: cutlass.Int32, + n_blocks: cutlass.Int32, + stream: cuda.CUstream, +): + l2norm_qk_kernel(q, k, q_n, k_n, inv_q, inv_k, n_q_rows, n_rows, h_q, h_k).launch( + grid=(n_blocks, 1, 1), block=(THREADS_PER_ROW * ROWS_PER_CTA, 1, 1), stream=stream + ) + + +@cute.jit +def l2norm_qk_bwd_launch( + dq: cute.Tensor, + dk: cute.Tensor, + q_n: cute.Tensor, + k_n: cute.Tensor, + inv_q: cute.Tensor, + inv_k: cute.Tensor, + n_q_rows: cutlass.Int32, + n_rows: cutlass.Int32, + h_q: cutlass.Int32, + h_k: cutlass.Int32, + n_blocks: cutlass.Int32, + stream: cuda.CUstream, +): + l2norm_qk_bwd_kernel(dq, dk, q_n, k_n, inv_q, inv_k, n_q_rows, n_rows, h_q, h_k).launch( + grid=(n_blocks, 1, 1), block=(THREADS_PER_ROW * ROWS_PER_CTA, 1, 1), stream=stream + ) + + +@functools.cache +def l2norm_cache(key): + return {} + + +def l2norm_qk(q, k, q_n, k_n, inv_q, inv_k, *, stream): + """Normalize q/k rows into the compact io workspace copies and stash the + fp32 inverse norms. Sources are read through their own strides.""" + ROWS = (128 // FWD_LANES) * FWD_ROWS_PER_GROUP + total, h_q, d = (int(s_) for s_ in q.shape) + total_k, h_k, d_k = (int(s_) for s_ in k.shape) + if d != 128 or d_k != 128 or total_k != total: + raise ValueError(f"q/k must be [total, H, 128] with matching totals; got {tuple(q.shape)} / {tuple(k.shape)}") + for name, t in (("q", q), ("k", k)): + st = tuple(int(s_) for s_ in t.stride()) + if st[2] != 1 or st[0] % 8 != 0 or st[1] % 8 != 0 or data_ptr(t) % 16 != 0: + raise ValueError(f"{name} needs a 16B-aligned base, unit channel stride, and outer strides in multiples of 8; got {st}") + for name, buf, h in (("q_n", q_n, h_q), ("k_n", k_n, h_k)): + if tuple(int(s_) for s_ in buf.stride()) != (h * 128, 128, 1): + raise ValueError(f"{name} workspace must be compact [total, {h}, 128]") + for name, buf, h in (("inv_q", inv_q, h_q), ("inv_k", inv_k, h_k)): + if tuple(int(s_) for s_ in buf.stride()) != (h, 1): + raise ValueError(f"{name} workspace must be compact [total, {h}] fp32") + n_q_rows = total * h_q + n_rows = n_q_rows + total * h_k + args = (n_q_rows, n_rows, h_q, h_k, (n_rows + ROWS - 1) // ROWS) + cache = l2norm_cache(("fwd", str(q.dtype))) + cu_stream = cuda.CUstream(int(stream)) + if "compiled" not in cache: + tensors = (q, k, q_n, k_n, inv_q, inv_k) + cache["compiled"] = cute.compile( + l2norm_qk_launch, + *(from_dlpack(t, assumed_align=16).mark_layout_dynamic(leading_dim=lead) for t, lead in zip(tensors, (2, 2, 2, 2, 1, 1))), + *(cutlass.Int32(a) for a in args), + cu_stream, + options="--enable-tvm-ffi", + ) + cache["compiled"](q, k, q_n, k_n, inv_q, inv_k, *args, cu_stream) + + +def l2norm_qk_bwd(dq, dk, q_n, k_n, inv_q, inv_k, *, stream): + """Project dq/dk in place through the normalize Jacobian using the saved + normalized rows and inverse norms. Run after any head-group fold so the + gradients are back at the native q/k head counts.""" + ROWS = ROWS_PER_CTA + total, h_q, d = (int(s_) for s_ in dq.shape) + total_k, h_k, d_k = (int(s_) for s_ in dk.shape) + if d != 128 or d_k != 128 or total_k != total: + raise ValueError(f"dq/dk must be [total, H, 128] with matching totals; got {tuple(dq.shape)} / {tuple(dk.shape)}") + for name, t in (("dq", dq), ("dk", dk)): + st = tuple(int(s_) for s_ in t.stride()) + if st[2] != 1 or st[0] % 8 != 0 or st[1] % 8 != 0 or data_ptr(t) % 16 != 0: + raise ValueError(f"{name} needs a 16B-aligned base, unit channel stride, and outer strides in multiples of 8; got {st}") + for name, buf, h in (("q_n", q_n, h_q), ("k_n", k_n, h_k)): + if tuple(int(s_) for s_ in buf.stride()) != (h * 128, 128, 1): + raise ValueError(f"{name} workspace must be compact [total, {h}, 128]") + for name, buf, h in (("inv_q", inv_q, h_q), ("inv_k", inv_k, h_k)): + if tuple(int(s_) for s_ in buf.stride()) != (h, 1): + raise ValueError(f"{name} workspace must be compact [total, {h}] fp32") + n_q_rows = total * h_q + n_rows = n_q_rows + total * h_k + args = (n_q_rows, n_rows, h_q, h_k, (n_rows + ROWS - 1) // ROWS) + cache = l2norm_cache(("bwd", str(dq.dtype))) + cu_stream = cuda.CUstream(int(stream)) + if "compiled" not in cache: + tensors = (dq, dk, q_n, k_n, inv_q, inv_k) + cache["compiled"] = cute.compile( + l2norm_qk_bwd_launch, + *(from_dlpack(t, assumed_align=16).mark_layout_dynamic(leading_dim=lead) for t, lead in zip(tensors, (2, 2, 2, 2, 1, 1))), + *(cutlass.Int32(a) for a in args), + cu_stream, + options="--enable-tvm-ffi", + ) + cache["compiled"](dq, dk, q_n, k_n, inv_q, inv_k, *args, cu_stream) diff --git a/python/cudnn/linear_attention/frost/common/split_k.py b/python/cudnn/linear_attention/frost/common/split_k.py index 878e20eef..2127b4499 100644 --- a/python/cudnn/linear_attention/frost/common/split_k.py +++ b/python/cudnn/linear_attention/frost/common/split_k.py @@ -55,22 +55,24 @@ parallel (one warp per boundary, one lane per window chunk), then thread 0 walks the probe results and emits work items into the caller's ``item_scratch``. -3. order: one CTA bitonic-sorts the emitted items into ``work_items``, - longest ``[cstart, cend)`` first, so the main kernels' ticket scheduler - consumes them in LPT order — the makespan tail is set by whatever starts - last, so the big items must go first. This is what keeps ragged varlen - batches balanced without cutting them. +3. order (:func:`order_body`, hosted by each kernel module's + prologue kernel alongside its TMA-descriptor build — one launch for + both): bitonic-sort the items into ``work_items``, longest ``[cstart, + cend)`` first, so the ticket scheduler consumes them in LPT order — + the makespan tail is set by whatever starts last, so the big items + must go first. This is what keeps ragged varlen batches balanced + without cutting them. -The order kernel also zeroes the main kernels' scheduler ticket rings +The order body also zeroes the main kernels' scheduler ticket rings (dirty on exit), and with ``split=False`` it replaces the whole pipeline: -scan and walk never launch, and the order kernel synthesizes the uncut +scan and walk never launch, and the prologue kernel synthesizes the uncut whole-sequence item per (batch, head) from ``cu_seqlens`` alone, then LPT-sorts those. That no-cuts table serves batch-invariant mode and -coarse checkpoint cadences (cuts may not cross a checkpoint period), so -ragged batches keep LPT scheduling at the cost of one single-CTA launch. +coarse checkpoint cadences (cuts may not cross a checkpoint period). """ import math +from typing import NamedTuple import cuda.bindings.driver as cuda @@ -82,6 +84,8 @@ from cudnn.frost.buffers import data_ptr +from .elementwise import softplus + WORK_ITEM_FIELDS = 8 WARMUP_CAP_CHUNKS = 32 # hard warmup cap: a cut must saturate within one warp of chunks per side MAX_BLOCKS = 2048 # piece-count ceiling; host clamps ideal_chunks so the per-tile block count fits @@ -92,15 +96,14 @@ SCAN_THREADS = SCAN_WARPS * WARP_SIZE SCAN_ROWS_PER_WARP = 4 # consecutive chunk rows per scan warp SCAN_TOKEN_STRIDE = 4 # sample every Nth token of a chunk: skipped tokens only RAISE the negative horizon sums -ORDER_THREADS = 1024 -ORDER_ELEMS = 4 -ORDER_CAPACITY = ( - ORDER_THREADS * ORDER_ELEMS -) # sort capacity (32 KB SMEM); the kernel always launches — past this the device-side branch copies through unsorted OVERHEAD_TOKENS = 256 # per-item fixed cost for the piece model: state reseed + pipeline refill + typical warmup P_WINDOW = 16 # fill-regime piece-count search width P_BELOW = 8 # how far below the ideal-cap floor the fill-regime search may go +ORDER_THREADS = 1024 +ORDER_ELEMS = 4 +ORDER_CAPACITY = ORDER_THREADS * ORDER_ELEMS # sort capacity (32 KB SMEM); past this the device-side branch copies through unsorted + DEFAULT_LOG2_THRESHOLD = -10.0 / math.log(2.0) # e^-10, in log2 units RCP_LN2 = 1.4426950408889634 # 1/ln(2): natural-log gates -> the scan's log2 domain @@ -379,6 +382,7 @@ def scan_kernel( def scan_scalar_kernel( b_t: cutlass.Constexpr[int], log_gate: cutlass.Constexpr[bool], + safe_gate: cutlass.Constexpr[bool], overhead_chunks: cutlass.Constexpr[int], n_heads_out: cutlass.Int32, n_tiles: cutlass.Int32, @@ -386,6 +390,8 @@ def scan_scalar_kernel( ideal_chunks: cutlass.Int32, batch_size: cutlass.Int32, mGate: cute.Tensor, + mALog: cute.Tensor | None, + mDtBias: cute.Tensor | None, mCuSeqlens: cute.Tensor, mChunkVals: cute.Tensor, mCount: cute.Tensor, @@ -393,8 +399,11 @@ def scan_scalar_kernel( """Scalar-gate scan (GDN): CTA ``(x, hg)`` covers 16 chunk-scratch rows for heads ``[hg*32, (hg+1)*32)``; lane ``l`` owns head ``hg*32 + l``, so gate reads and chunk-value writes are coalesced across lanes and every - lane accumulates its own head — no reduction. CTA (0, 0) zeroes the - item count (the scheduler rings are the order kernel's job).""" + lane accumulates its own head — no reduction. With ``safe_gate`` the + gate holds raw logits and each token contributes the GDN transform in + log2 domain: ``-exp(A_log[h]) * softplus(g + dt_bias[h]) * RCP_LN2``. + CTA (0, 0) zeroes the item count (the scheduler rings are the order + kernel's job).""" tidx, _, _ = cute.arch.thread_idx() bidx = cute.arch.block_idx() tidx = cutlass.Int32(tidx) @@ -405,6 +414,12 @@ def scan_scalar_kernel( h = cutlass.Int32(bidx[1]) * cutlass.Int32(WARP_SIZE) + lidx h_ok = h < n_heads_out h_r = h if h_ok else n_heads_out - cutlass.Int32(1) + a_l2 = cutlass.Float32(0.0) + bias = cutlass.Float32(0.0) + if cutlass.const_expr(safe_gate): + # per-head transform constants, fixed for the lane's whole sweep + a_l2 = -cute.math.exp2(mALog[h_r].to(cutlass.Float32) * cutlass.Float32(RCP_LN2), fastmath=True) * cutlass.Float32(RCP_LN2) + bias = mDtBias[h_r].to(cutlass.Float32) row0 = (cutlass.Int32(bidx[0]) * cutlass.Int32(SCAN_WARPS) + widx) * cutlass.Int32(SCAN_ROWS_PER_WARP) # batch of the warp's first row: largest b with cu[b] // b_t + b <= row0 @@ -442,8 +457,12 @@ def scan_scalar_kernel( inb = pos < batch_end pos_r = pos if inb else batch_start gv = (mGate.iterator + cutlass.Int64(pos_r) * cutlass.Int64(mGate.stride[0]) + h_r).load() - gv = gv if inb else oob - acc = acc + clamped_log2(log_gate, gv) + if cutlass.const_expr(safe_gate): + contrib = a_l2 * softplus(gv + bias) + acc = acc + (contrib if inb else cutlass.Float32(0.0)) + else: + gv = gv if inb else oob + acc = acc + clamped_log2(log_gate, gv) if h_ok: mChunkVals[cv_base + c, h] = acc @@ -579,11 +598,14 @@ def write_item( mWorkItems[dst, cutlass.Int32(f)] = mStaging[src, cutlass.Int32(f)] -@cute.kernel -def order_kernel( +@cute.jit +def order_body( gen: cutlass.Constexpr[bool], has_sched: cutlass.Constexpr[bool], b_t: cutlass.Constexpr[int], + n_threads: cutlass.Constexpr[int], + order_elems: cutlass.Constexpr[int], + tidx, n_heads_out: cutlass.Int32, n_tiles: cutlass.Int32, mCuSeqlens: cute.Tensor, @@ -591,16 +613,21 @@ def order_kernel( mCount: cute.Tensor, mWorkItems: cute.Tensor, mSched: cute.Tensor | None, + sKey, + sIdx, + sSpread, ): - """LPT ordering (single CTA): bitonic-sort the items by span ``cend - - cstart``, longest first, into the final table, so the ticket scheduler - starts the big items before the filler. Sorts the walk's staged items, - or with ``gen`` synthesizes the uncut whole-sequence item per (batch, - head) from ``cu_seqlens`` directly — the no-cuts table. Thread 0 also - zeroes every ``sched_ctr`` cell (the main kernels' ticket rings, dirty - on exit): this kernel runs on every table build, split or not.""" - tidx, _, _ = cute.arch.thread_idx() - tidx = cutlass.Int32(tidx) + """LPT ordering body over ``n_threads`` CTA threads and caller-owned SMEM + staging (``sKey``/``sIdx`` of ``n_threads * order_elems`` Int32 cells + + a 2-cell ``sSpread``): bitonic-sort the items by span ``cend - cstart``, + longest first, into the final table. Sorts the walk's staged items, or + with ``gen`` synthesizes the uncut whole-sequence item per (batch, head) + from ``cu_seqlens`` directly — the no-cuts table. Thread 0 also zeroes + every ``sched_ctr`` cell (the main kernels' ticket rings, dirty on exit). + Runs on the standalone :func:`order_kernel` CTA, or fused into a main + kernel's CTA 0 prologue. Internally CTA-wide-barriers; every thread of + the calling CTA must reach it.""" + capacity = cutlass.const_expr(n_threads * order_elems) if cutlass.const_expr(has_sched): if tidx == 0: si = cutlass.Int32(0) @@ -613,15 +640,12 @@ def order_kernel( mCount[0] = n_tiles else: n = mCount[0] - if n > cutlass.Int32(ORDER_CAPACITY): + if n > cutlass.Int32(capacity): i = tidx while i < n: write_item(gen, b_t, n_heads_out, mCuSeqlens, mStaging, mWorkItems, i, i) - i = i + cutlass.Int32(ORDER_THREADS) + i = i + cutlass.Int32(n_threads) else: - sKey = cutlass.Array(cutlass.Int32, ORDER_CAPACITY, space=cutlass.AddressSpace.smem, alignment=16) - sIdx = cutlass.Array(cutlass.Int32, ORDER_CAPACITY, space=cutlass.AddressSpace.smem, alignment=16) - sSpread = cutlass.Array(cutlass.Int32, 2, space=cutlass.AddressSpace.smem, alignment=8) if tidx == 0: sSpread[0] = cutlass.Int32(2147483647) sSpread[1] = cutlass.Int32(-2147483648) @@ -631,8 +655,8 @@ def order_kernel( nvvm.barrier_cta_sync() kmin = cutlass.Int32(2147483647) kmax = cutlass.Int32(-2147483648) - for e in cutlass.range_constexpr(ORDER_ELEMS): - i = tidx + cutlass.Int32(e * ORDER_THREADS) + for e in cutlass.range_constexpr(order_elems): + i = tidx + cutlass.Int32(e * n_threads) if i < n: if cutlass.const_expr(gen): batch_idx, head_idx, batch_start, batch_end, num_chunks_b = gen_item_bounds(b_t, n_heads_out, mCuSeqlens, i) @@ -655,14 +679,14 @@ def order_kernel( i2 = tidx while i2 < n: write_item(gen, b_t, n_heads_out, mCuSeqlens, mStaging, mWorkItems, i2, i2) - i2 = i2 + cutlass.Int32(ORDER_THREADS) + i2 = i2 + cutlass.Int32(n_threads) else: k = cutlass.Int32(2) while k <= b_pad: j = k // cutlass.Int32(2) while j > 0: - for e in cutlass.range_constexpr(ORDER_ELEMS): - i = tidx + cutlass.Int32(e * ORDER_THREADS) + for e in cutlass.range_constexpr(order_elems): + i = tidx + cutlass.Int32(e * n_threads) if i < b_pad: l = i ^ j if l > i: @@ -680,8 +704,8 @@ def order_kernel( nvvm.barrier_cta_sync() j = j // cutlass.Int32(2) k = k * cutlass.Int32(2) - for e in cutlass.range_constexpr(ORDER_ELEMS): - i = tidx + cutlass.Int32(e * ORDER_THREADS) + for e in cutlass.range_constexpr(order_elems): + i = tidx + cutlass.Int32(e * n_threads) if i < n: src = sIdx[i] write_item(gen, b_t, n_heads_out, mCuSeqlens, mStaging, mWorkItems, i, src) @@ -745,6 +769,7 @@ def launch( scan_scalar_kernel( b_t, log_gate, + safe_gate, overhead_chunks, n_heads_out, n_tiles, @@ -752,6 +777,8 @@ def launch( ideal_chunks, batch_size, mGate, + mALog, + mDtBias, mCuSeqlens, mChunkVals, mCount, @@ -777,27 +804,56 @@ def launch( block=(THREADS_PER_BLOCK, 1, 1), stream=stream, ) - order_kernel( - not split, - has_sched, - b_t, - n_heads_out, - n_tiles, - mCuSeqlens, - mStaging, - mCount, - mWorkItems, - mSched, - ).launch( - grid=(1, 1, 1), - block=(ORDER_THREADS, 1, 1), - stream=stream, - ) compiled_cache = {} +class TableRecipe(NamedTuple): + """Build-time facts of one split-table launch: everything static is + settled once so :func:`run_table` can replay the call as a straight + line. Produced by :func:`build_split_table`.""" + + compiled: object + split: bool + safe_gate: bool + n_heads_out: int + n_tiles: int + num_sms: int + ideal_chunks: int + batch_size: int + log2_threshold: float + gate_scale_log2: float + n_scan_ctas: int + n_walk_ctas: int + + +def run_table(r, gate, a_log, dt_bias, cu_seqlens, chunk_scratch, item_scratch, work_items, work_count, sched_ctr, stream) -> None: + """The lowered split-table launch: no validation, no key build. Only + buffers move between calls; every scalar comes from the recipe.""" + r.compiled( + r.n_heads_out, + r.n_tiles, + r.num_sms, + r.ideal_chunks, + r.batch_size, + r.log2_threshold, + r.gate_scale_log2, + gate if r.split else None, + a_log if r.safe_gate else None, + dt_bias if r.safe_gate else None, + cu_seqlens, + chunk_scratch if r.split else None, + item_scratch if r.split else None, + work_items, + work_count, + sched_ctr, + r.n_scan_ctas, + r.n_walk_ctas, + cuda.CUstream(int(stream)), + ) + + def build_split_table( gate, cu_seqlens, @@ -819,16 +875,17 @@ def build_split_table( sched_ctr=None, split=True, stream, -) -> None: +) -> "TableRecipe": """Fill ``work_items``/``work_count`` with the split-K partition of ``gate`` / ``cu_seqlens (B+1,) int32``, LPT-ordered (longest item first). A 2-D ``(total_tokens, HO)`` gate is the scalar GDN kind; a 3-D ``(total_tokens, HO, DK)`` gate is the per-key-channel KDA / GDN-2 kind. With ``log_gate`` the gate values are natural-log decay instead - of raw linear alpha. With ``safe_gate`` (per-channel only) the gate - holds RAW logits and the scan applies the KDA safe-gate transform - ``gate_lower_bound * sigmoid(exp(a_log) * (g + dt_bias))`` per element, - so cuts land on true decay values. + of raw linear alpha. With ``safe_gate`` the gate holds RAW logits and + the scan applies the matching transform so cuts land on true decay + values: per-channel (KDA / GDN-2, ``gate_lower_bound`` required) + ``gate_lower_bound * sigmoid(exp(a_log) * (g + dt_bias))`` per element; + scalar (GDN) ``-exp(a_log[h]) * softplus(g + dt_bias[h])`` per head. With ``split=False`` the scan and walk never launch: the order kernel alone synthesizes the no-cuts table — the uncut whole-sequence item per @@ -853,14 +910,14 @@ def build_split_table( if len(gate.shape) not in (2, 3): raise ValueError(f"gate must be (total_tokens, HO) or (total_tokens, HO, DK), got {tuple(gate.shape)}") gate_channels = gate.shape[2] if len(gate.shape) == 3 else 0 - if safe_gate and gate_channels == 0: - raise ValueError("safe_gate applies to per-channel gates only") - if safe_gate and (a_log is None or dt_bias is None or gate_lower_bound is None): - raise ValueError("safe_gate requires a_log, dt_bias, and gate_lower_bound") + if safe_gate and (a_log is None or dt_bias is None): + raise ValueError("safe_gate requires a_log and dt_bias") + if safe_gate and gate_channels > 0 and gate_lower_bound is None: + raise ValueError("per-channel safe_gate requires gate_lower_bound") if not safe_gate: a_log = None dt_bias = None - gate_scale_log2 = float(gate_lower_bound) * RCP_LN2 if safe_gate else 0.0 + gate_scale_log2 = float(gate_lower_bound) * RCP_LN2 if safe_gate and gate_channels > 0 else 0.0 n_heads_out = gate.shape[1] batch_size = cu_seqlens.shape[0] - 1 if split: @@ -880,9 +937,6 @@ def build_split_table( f"item_scratch must match work_items (max_items, {WORK_ITEM_FIELDS}) int32, got {tuple(item_scratch.shape)} vs {tuple(work_items.shape)}" ) else: - # no-cuts table: the order kernel synthesizes it from cu_seqlens; the - # gate and every scan/walk operand are unused (normalized out of the - # compile key so all variants share one specialization per b_t/HO) if work_items.shape[0] < n_tiles or work_items.shape[1] != WORK_ITEM_FIELDS: raise ValueError(f"work_items must be (>= {n_tiles}, {WORK_ITEM_FIELDS}) int32, got {tuple(work_items.shape)}") log_gate = False @@ -969,3 +1023,17 @@ def build_split_table( n_walk_ctas, cu_stream, ) + return TableRecipe( + compiled_cache[key], + bool(split), + bool(safe_gate), + n_heads_out, + n_tiles, + num_sms, + int(ideal_chunks), + batch_size, + float(log2_threshold), + float(gate_scale_log2), + n_scan_ctas, + n_walk_ctas, + ) diff --git a/python/cudnn/linear_attention/frost/common/thd.py b/python/cudnn/linear_attention/frost/common/thd.py index c81f459be..e69ce580b 100644 --- a/python/cudnn/linear_attention/frost/common/thd.py +++ b/python/cudnn/linear_attention/frost/common/thd.py @@ -7,7 +7,7 @@ * :func:`emit_seq_descs` — device helper (one electing thread) that builds a per-BATCH TMA-descriptor array in GMEM for varlen loads/stores over a packed ``[T,H,D]`` tensor whose head axis is a load coordinate. Each op - calls it from a single combined ``build_all_descs_kernel`` launch (one + calls it from a single ``prologue_kernel`` launch (one warp per array). * :func:`emit_checkpoint_seq_descs` — its per-chunk-checkpoint sibling; derives the per-sequence checkpoint offsets from the token ``cu_seqlens`` in place of a diff --git a/python/cudnn/linear_attention/frost/engine.py b/python/cudnn/linear_attention/frost/engine.py new file mode 100644 index 000000000..a20dd52ef --- /dev/null +++ b/python/cudnn/linear_attention/frost/engine.py @@ -0,0 +1,90 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Engine layer shared by the FROST linear-attention families: the +check_support core all three run, and the compiled-plan wrapper their +``build_plan`` returns.""" + +from __future__ import annotations + +import cudnn +from cudnn.engines.base import CompiledPlan, bind_ports +from cudnn.frost import buffers +from cudnn.frost.workspace import Workspace + + +def frost_la_gate(engine: str, facts, op: str) -> None: + """The FROST LA engines' shared check_support core: the analyzer record, + the device/DSL environment, and the gates common to all three kernels.""" + if facts is None or facts.op != op: + raise NotImplementedError(f"{engine} supports exactly one {op}/{op}_BWD node") + if facts.invalid: + raise NotImplementedError(f"{engine}: {facts.invalid}") + sm = buffers.current_sm() + if sm is None or not (100 <= sm <= 103): + raise NotImplementedError(f"{engine} requires SM100-SM103 (found {sm})") + installed, version = buffers.cutedsl_state() + if not installed: + raise NotImplementedError(f"{engine} requires the cutedsl extra (nvidia-cutlass-dsl), which is not installed") + if buffers.cutedsl_too_old(version): + want = ".".join(str(v) for v in buffers.CUTEDSL_MIN_VERSION) + raise NotImplementedError(f"{engine} requires nvidia-cutlass-dsl >= {want}; found {version[1]}") + if not facts.uniform_io: + raise NotImplementedError(f"{engine}: q/k/v dtypes must match") + if facts.io_dtype not in (cudnn.data_type.BFLOAT16, cudnn.data_type.HALF, None): + raise NotImplementedError(f"{engine}: q/k/v must be fp16/bf16, got {facts.io_dtype}") + if not facts.thd_layout: + raise NotImplementedError(f"{engine}: q/k/v must be THD [total_T, heads, dim]") + if facts.d_qk != 128 or facts.d_v != 128: + raise NotImplementedError(f"{engine}: head dims must be 128 (the recurrent state is 128x128), got K={facts.d_qk} V={facts.d_v}") + if facts.h_k not in (facts.h_q, facts.h_v): + raise NotImplementedError(f"{engine}: k heads ({facts.h_k}) must match q's ({facts.h_q}) or v's ({facts.h_v}; canonical GQA shares grouped k/v heads)") + if facts.h_v != facts.h_q and max(facts.h_q, facts.h_v) % min(facts.h_q, facts.h_v) != 0: + raise NotImplementedError(f"{engine}: q heads ({facts.h_q}) and v heads ({facts.h_v}) must be equal or one a multiple of the other") + if facts.g_dtype not in (cudnn.data_type.FLOAT, None): + raise NotImplementedError(f"{engine}: 'g' must be fp32, got {facts.g_dtype}") + if facts.cu_dtype not in (cudnn.data_type.INT32, cudnn.data_type.INT64, None): + raise NotImplementedError(f"{engine}: 'cu_seqlens' must be int32/int64, got {facts.cu_dtype}") + + +def dense_layout_message(plan_name, ports, offender) -> str: + """Name the port behind ``all_dense_layout``'s failing slot. Buffers pass + straight to the stride-plumbed kernels, so the one execute-time rule is a + stride-1 innermost dim; this walk only runs on the way to raising.""" + for slots in ports.values(): + for direction in (slots.inputs, slots.outputs): + for port, slot in direction.items(): + if slot == offender: + return f"{plan_name}: buffer for {port!r} must have a stride-1 innermost dim (buffers pass straight to the kernel)" + return f"{plan_name}: the buffer at variant-pack slot {offender} must have a stride-1 innermost dim" + + +class FrostLaPlan(CompiledPlan): + """A compiled LA executor, driven from the normalized variant pack: the + port-to-slot join is a property of the graph, so it happens once and is + kept; between executes only the buffer addresses move.""" + + takes_variant_pack = True + + def __init__(self, compiled): + self.compiled = compiled + self.ports = None + self.indices = None + + def get_workspace_size(self) -> int: + return self.compiled.workspace_bytes() + + def execute(self, graph, variant_pack, ctx) -> None: + ports = self.ports + if ports is None: + ports = self.ports = bind_ports(graph, variant_pack) + (slots,) = ports.values() + names = list(slots.inputs) + list(slots.outputs) + self.indices = list(slots.inputs.values()) + list(slots.outputs.values()) + self.compiled.bind(names) + ok, offender = variant_pack.all_dense_layout() + if not ok: + raise ValueError(dense_layout_message(self.compiled.plan_name, ports, offender)) + views = variant_pack.operands(self.indices) + workspace = Workspace.over(variant_pack, self.compiled.ws_bytes, type(self.compiled).__name__) + self.compiled.run(views, workspace, ctx.stream) diff --git a/python/cudnn/linear_attention/frost/gdn2_engine.py b/python/cudnn/linear_attention/frost/gdn2_engine.py index 39487e44c..5e3f8d1ed 100644 --- a/python/cudnn/linear_attention/frost/gdn2_engine.py +++ b/python/cudnn/linear_attention/frost/gdn2_engine.py @@ -10,7 +10,6 @@ from __future__ import annotations import math -from typing import Any from cudnn import behavior_note from cudnn.engines.base import BaseEngine, CompiledPlan @@ -18,7 +17,8 @@ from cudnn.frost.buffers import current_device_id from cudnn.frost.device import multiprocessor_count from cudnn.frost.workspace import WorkspaceLayout, carve_plan -from ..graph_analyzer import FrostLaPlan, frost_la_gate, require, analyze +from ..graph_analyzer import analyze +from .engine import FrostLaPlan, frost_la_gate def build_gdn2(graph): @@ -46,7 +46,7 @@ class Gdn2FrostEngine(BaseEngine): kernel with a forward checkpoint recompute when the graph has no ``state_checkpoints`` input.""" name = "gdn2_frost" - behavior_notes = (behavior_note.RUNTIME_COMPILATION,) # JIT-compiled at build_plans() + behavior_notes = (behavior_note.RUNTIME_COMPILATION,) def check_support(self, graph) -> None: import cudnn @@ -58,31 +58,40 @@ def check_support(self, graph) -> None: raise NotImplementedError(f"Gdn2FrostEngine: checkpoint_every_n_tokens must be a positive multiple of 16 on the GDN-2 node (got {ckpt})") if not facts.gates_at_ho: raise NotImplementedError(f"Gdn2FrostEngine: g/beta/w must carry HO = max(q, v) heads ({facts.h_o})") - if facts.is_bwd and facts.safe_gate: - raise NotImplementedError("Gdn2FrostEngine: safe_gate is a forward-node attribute") fp32 = cudnn.data_type.FLOAT if facts.io_dtype is not None: - require("Gdn2FrostEngine", "beta", facts.beta_dtype, facts.io_dtype) - require("Gdn2FrostEngine", "w", facts.w_dtype, facts.io_dtype) - require("Gdn2FrostEngine", "a_log", facts.a_log_dtype, fp32) - require("Gdn2FrostEngine", "dt_bias", facts.dt_bias_dtype, fp32) - if facts.is_bwd: - for port, got in (("dO", facts.do_dtype), ("state_checkpoints", facts.state_checkpoints_dtype)): + for port, got in (("beta", facts.beta_dtype), ("w", facts.w_dtype)): if got not in (facts.io_dtype, None): - raise NotImplementedError(f"Gdn2FrostEngine: '{port}' must match the io dtype") - require("Gdn2FrostEngine", "initial_state", facts.state_dtype, fp32) - require("Gdn2FrostEngine", "d_final_state", facts.d_final_state_dtype, fp32) - require("Gdn2FrostEngine", "d_initial_state", facts.d_initial_state_dtype, fp32) - require("Gdn2FrostEngine", "dG", facts.dg_dtype, fp32) - if facts.io_dtype is not None: - require("Gdn2FrostEngine", "dBeta", facts.dbeta_dtype, facts.io_dtype) - require("Gdn2FrostEngine", "dW", facts.dw_dtype, facts.io_dtype) + raise NotImplementedError(f"Gdn2FrostEngine: '{port}' must match the io dtype, got {got}") + for port, got in (("a_log", facts.a_log_dtype), ("dt_bias", facts.dt_bias_dtype)): + if got not in (fp32, None): + raise NotImplementedError(f"Gdn2FrostEngine: '{port}' must be fp32, got {got}") + if facts.is_bwd: + for port, got in ( + ("d_a_log", facts.d_a_log_dtype), + ("d_dt_bias", facts.d_dt_bias_dtype), + ("initial_state", facts.state_dtype), + ("d_final_state", facts.d_final_state_dtype), + ("d_initial_state", facts.d_initial_state_dtype), + ("dG", facts.dg_dtype), + ): + if got not in (fp32, None): + raise NotImplementedError(f"Gdn2FrostEngine: '{port}' must be fp32, got {got}") + for port, got in ( + ("dO", facts.do_dtype), + ("state_checkpoints", facts.state_checkpoints_dtype), + ("dBeta", facts.dbeta_dtype), + ("dW", facts.dw_dtype), + ): + if facts.io_dtype is not None and got not in (facts.io_dtype, None): + raise NotImplementedError(f"Gdn2FrostEngine: '{port}' must match the io dtype, got {got}") else: state_dtypes = (fp32, cudnn.data_type.BFLOAT16) - require("Gdn2FrostEngine", "initial_state", facts.state_dtype, state_dtypes) - require("Gdn2FrostEngine", "final_state", facts.final_state_dtype, state_dtypes) - if facts.io_dtype is not None: - require("Gdn2FrostEngine", "state_checkpoints", facts.state_checkpoints_out_dtype, facts.io_dtype) + for port, got in (("initial_state", facts.state_dtype), ("final_state", facts.final_state_dtype)): + if got not in state_dtypes + (None,): + raise NotImplementedError(f"Gdn2FrostEngine: '{port}' must be fp32/bf16, got {got}") + if facts.io_dtype is not None and facts.state_checkpoints_out_dtype not in (facts.io_dtype, None): + raise NotImplementedError("Gdn2FrostEngine: 'state_checkpoints' must match the io dtype") if not facts.state_pair_match: raise NotImplementedError("Gdn2FrostEngine: initial_state and final_state dtypes must match") @@ -94,16 +103,20 @@ class CompiledGdn2: """Compiled FROST GDN-2 plan: a callable over the resolved node buffers.""" def __init__(self, node, kernel_mod): - from .common.split_k import WORK_ITEM_FIELDS, build_split_table, chunk_scratch_rows, compute_ideal_chunks, max_work_items + from .common.split_k import WORK_ITEM_FIELDS, build_split_table, chunk_scratch_rows, compute_ideal_chunks, max_work_items, run_table self.node = node self.kernel = kernel_mod self.build_split_table = build_split_table + self.run_table = run_table + self.table = None + self.kcache = None self.plan_name = "Gdn2FrostEngine (GDN2)" scale = node.params.get("scale") self.scale = float(scale) if scale is not None else 1.0 / math.sqrt(node.inputs["q"].dim[-1]) self.use_qk_l2norm = bool(node.params.get("use_qk_l2norm", False)) self.safe_gate = bool(node.params.get("safe_gate", False)) + self.use_beta_sigmoid = bool(node.params.get("use_beta_sigmoid", False)) glb = node.params.get("gate_lower_bound") self.gate_lower_bound = float(glb) if glb is not None else kernel_mod.DEFAULT_GATE_LOWER_BOUND self.has_final_state = "final_state" in node.outputs @@ -119,7 +132,7 @@ def __init__(self, node, kernel_mod): HO = g.dim[1] B = node.inputs["cu_seqlens"].dim[0] - 1 layout = WorkspaceLayout() - self.off_sched = layout.add(8) # [ticket, done] for the dynamic scheduler + self.off_sched = layout.add(8) self.num_sm = multiprocessor_count(current_device_id()) self.n_tiles = B * HO self.n_heads_out = HO @@ -139,6 +152,7 @@ def __init__(self, node, kernel_mod): self.tensormap_bytes = tensormap_workspace_bytes(kernel_mod, B) self.off_tensormaps = layout.add(self.tensormap_bytes, align=128) + self.needs_table = self.split self.ws_bytes = layout.size regions = [ (self.off_sched, "int32", (2,)), @@ -156,21 +170,36 @@ def __init__(self, node, kernel_mod): def workspace_bytes(self) -> int: return self.ws_bytes - def __call__(self, node_buffers, *, workspace=None, stream=None) -> Any: - nb = node_buffers[self.node] - q = nb.inputs["q"] - k = nb.inputs["k"] - v = nb.inputs["v"] - g = nb.inputs["g"] - beta = nb.inputs["beta"] - w = nb.inputs["w"] - cu = nb.inputs["cu_seqlens"] - state0 = nb.inputs.get("initial_state") - o = nb.outputs["O"] - final_state = nb.outputs["final_state"] if self.has_final_state else None - state_checkpoints = nb.outputs["state_checkpoints"] if self.has_state_checkpoints else None - a_log = nb.inputs.get("a_log") - dt_bias = nb.inputs.get("dt_bias") + def bind(self, names) -> None: + pos = {name: i for i, name in enumerate(names)} + self.iq = pos["q"] + self.ik = pos["k"] + self.iv = pos["v"] + self.ig = pos["g"] + self.ibeta = pos["beta"] + self.iw = pos["w"] + self.icu = pos["cu_seqlens"] + self.is0 = pos.get("initial_state") + self.io_ = pos["O"] + self.ifs = pos.get("final_state") + self.ick = pos.get("state_checkpoints") + self.ia_log = pos.get("a_log") + self.idt_bias = pos.get("dt_bias") + + def run(self, views, workspace, stream) -> None: + q = views[self.iq] + k = views[self.ik] + v = views[self.iv] + g = views[self.ig] + beta = views[self.ibeta] + w = views[self.iw] + cu = views[self.icu] + state0 = views[self.is0] if self.is0 is not None else None + o = views[self.io_] + final_state = views[self.ifs] if self.ifs is not None else None + state_checkpoints = views[self.ick] if self.ick is not None else None + a_log = views[self.ia_log] if self.ia_log is not None else None + dt_bias = views[self.idt_bias] if self.idt_bias is not None else None stream = stream if stream is not None else 0 if self.split: @@ -178,32 +207,64 @@ def __call__(self, node_buffers, *, workspace=None, stream=None) -> Any: else: sched_ctr, work_items, work_count, tensormaps = workspace.carve(self.carve) item_scratch = chunk_scratch = None - self.build_split_table( - g, - cu, - work_items, - work_count, - ideal_chunks=self.ideal, - n_tiles=self.n_tiles, - num_sms=self.num_sm, - b_t=self.b_t, - chunk_scratch=chunk_scratch, - item_scratch=item_scratch, - log_gate=True, - safe_gate=self.safe_gate, - a_log=a_log, - dt_bias=dt_bias, - gate_lower_bound=self.gate_lower_bound if self.safe_gate else None, - sched_ctr=sched_ctr, - split=self.split, - stream=stream, - ) + + if self.kcache is not None and (self.table is not None or not self.needs_table): + if self.needs_table: + self.run_table(self.table, g, a_log, dt_bias, cu, chunk_scratch, item_scratch, work_items, work_count, sched_ctr, stream) + self.kernel.run_prefill( + self.kcache, + q, + k, + v, + g, + a_log if self.safe_gate else None, + dt_bias if self.safe_gate else None, + beta, + w, + cu, + state0, + o, + final_state, + state_checkpoints, + work_items, + work_count, + sched_ctr, + item_scratch, + tensormaps, + self.ckpt if self.has_state_checkpoints else 0, + self.scale, + stream, + ) + return + + if not self.needs_table: + self.table = None + else: + self.table = self.build_split_table( + g, + cu, + work_items, + work_count, + ideal_chunks=self.ideal, + n_tiles=self.n_tiles, + num_sms=self.num_sm, + b_t=self.b_t, + chunk_scratch=chunk_scratch, + item_scratch=item_scratch, + log_gate=True, + safe_gate=self.safe_gate, + a_log=a_log, + dt_bias=dt_bias, + gate_lower_bound=self.gate_lower_bound if self.safe_gate else None, + sched_ctr=sched_ctr, + split=self.split, + stream=stream, + ) ckpt_kwargs = {} if self.has_state_checkpoints: - # the kernel derives the per-sequence checkpoint entry offsets on device ckpt_kwargs = dict(checkpoint_every_n_tokens=self.ckpt, output_state_checkpoints=state_checkpoints) - self.kernel.chunk_gdn2_sm100( + self.kcache = self.kernel.chunk_gdn2_sm100( q, k, v, @@ -220,9 +281,11 @@ def __call__(self, node_buffers, *, workspace=None, stream=None) -> Any: gate_lower_bound=self.gate_lower_bound, a_log=a_log, dt_bias=dt_bias, + use_beta_sigmoid=self.use_beta_sigmoid, work_items=work_items, work_count=work_count, sched_ctr=sched_ctr, + work_item_scratch=item_scratch, tensormap_workspace=tensormaps, **ckpt_kwargs, stream=stream, @@ -236,20 +299,33 @@ class CompiledGdn2Bwd: plus GVA/GQA head scratch for dQ/dK/dV.""" def __init__(self, node, bwd_mod, regen_mod): - from .common.split_k import WORK_ITEM_FIELDS, build_split_table, chunk_scratch_rows, compute_ideal_chunks, max_work_items + from .common.split_k import WORK_ITEM_FIELDS, build_split_table, chunk_scratch_rows, compute_ideal_chunks, max_work_items, run_table self.node = node self.bwd = bwd_mod self.regen = regen_mod self.build_split_table = build_split_table + self.run_table = run_table + self.table = None + self.kcache = None + self.regen_cache = None self.plan_name = "Gdn2FrostEngine (GDN2_BWD)" from .common.downcast import downcast_state + from .common.gate_bwd import GATE_BWD_BLOCKS, channel_gate_bwd + from .common.head_reduce import head_group_reduce from .common.host import tensormap_workspace_bytes self.downcast_state = downcast_state + self.head_group_reduce = head_group_reduce + self.channel_gate_bwd = channel_gate_bwd scale = node.params.get("scale") self.scale = float(scale) if scale is not None else 1.0 / math.sqrt(node.inputs["q"].dim[-1]) self.use_qk_l2norm = bool(node.params.get("use_qk_l2norm", False)) + self.safe_gate = bool(node.params.get("safe_gate", False)) + self.use_beta_sigmoid = bool(node.params.get("use_beta_sigmoid", False)) + glb = node.params.get("gate_lower_bound") + self.gate_lower_bound = float(glb) if glb is not None else bwd_mod.DEFAULT_GATE_LOWER_BOUND + self.gate_bwd_blocks = GATE_BWD_BLOCKS self.has_state_checkpoints = "state_checkpoints" in node.inputs self.has_state0 = "initial_state" in node.inputs self.has_dstate0 = "d_initial_state" in node.outputs @@ -264,12 +340,11 @@ def __init__(self, node, bwd_mod, regen_mod): self.io_name = "float16" if node.inputs["q"].get_data_type().name == "HALF" else "bfloat16" self.n_heads_out, self.total = HO, total layout = WorkspaceLayout() - self.off_sched = layout.add(16) # one [ticket, done] ring each for the regen and bwd kernels + self.off_sched = layout.add(16) self.num_sm = multiprocessor_count(current_device_id()) self.bwd_dyn_sched = B * HO <= self.num_sm self.batch_invariant = bool(node.params.get("batch_invariant", False)) - # cuts never in batch-invariant mode: whole-sequence items keep each - # sequence's math independent of the batch composition + # cuts never in batch-invariant mode self.split = not self.batch_invariant self.n_tiles = B * HO if self.split: @@ -284,7 +359,6 @@ def __init__(self, node, bwd_mod, regen_mod): self.off_item_scratch = layout.add(self.work_item_rows * WORK_ITEM_FIELDS * 4) self.chunk_scratch_rows = chunk_scratch_rows(total, B, self.b_t) self.off_chunk_scratch = layout.add(self.chunk_scratch_rows * HO * 4) - # chunk-0 entering state, io dtype (downcast initial_state or zeros) self.off_state0_io = layout.add(B * HO * K * V * 2) if self.has_state0 else None if not self.has_state_checkpoints: self.state_checkpoints_rows = max(total // self.b_t + B, 1) @@ -301,8 +375,13 @@ def __init__(self, node, bwd_mod, regen_mod): self.off_dk_ho = layout.add(total * HO * K * 2) if self.fold_dv: self.off_dv_ho = layout.add(total * HO * V * 2) + if self.safe_gate: + self.off_gate_part_a = layout.add(self.gate_bwd_blocks * HO * K * 4) + self.off_gate_part_dt = layout.add(self.gate_bwd_blocks * HO * K * 4) self.bwd_tm_bytes = tensormap_workspace_bytes(bwd_mod, B) self.off_bwd_tensormaps = layout.add(self.bwd_tm_bytes, align=128) + self.order_in_regen = not self.has_state_checkpoints + self.needs_table = self.split self.ws_bytes = layout.size regions = [ ("sched_regen", self.off_sched, "int32", (2,)), @@ -326,58 +405,183 @@ def __init__(self, node, bwd_mod, regen_mod): regions.append(("dk_ho", self.off_dk_ho, self.io_name, (total, HO, K))) if self.fold_dv: regions.append(("dv_ho", self.off_dv_ho, self.io_name, (total, HO, V))) + if self.safe_gate: + regions.append(("gate_part_a", self.off_gate_part_a, "float32", (self.gate_bwd_blocks * HO * K,))) + regions.append(("gate_part_dt", self.off_gate_part_dt, "float32", (self.gate_bwd_blocks * HO * K,))) self.carve_names = [name for name, _off, _dt, _shape in regions] self.carve = carve_plan("Gdn2FrostEngine (GDN2_BWD)", [(off, dt, shape) for _name, off, dt, shape in regions]) def workspace_bytes(self) -> int: return self.ws_bytes - def __call__(self, node_buffers, *, workspace=None, stream=None) -> Any: - nb = node_buffers[self.node] - q = nb.inputs["q"] - k = nb.inputs["k"] - v = nb.inputs["v"] - g = nb.inputs["g"] - beta = nb.inputs["beta"] - w = nb.inputs["w"] - cu = nb.inputs["cu_seqlens"] - do = nb.inputs["dO"] - state_checkpoints = nb.inputs.get("state_checkpoints") - state0 = nb.inputs.get("initial_state") - dstate_in = nb.inputs.get("d_final_state") - dq = nb.outputs["dQ"] - dk = nb.outputs["dK"] - dv = nb.outputs["dV"] - dg = nb.outputs["dG"] - db = nb.outputs["dBeta"] - dw = nb.outputs["dW"] - dstate0 = nb.outputs.get("d_initial_state") + def bind(self, names) -> None: + pos = {name: i for i, name in enumerate(names)} + self.iq = pos["q"] + self.ik = pos["k"] + self.iv = pos["v"] + self.ig = pos["g"] + self.ibeta = pos["beta"] + self.iw = pos["w"] + self.icu = pos["cu_seqlens"] + self.ido = pos["dO"] + self.ick = pos.get("state_checkpoints") + self.is0 = pos.get("initial_state") + self.idfs = pos.get("d_final_state") + self.idq = pos["dQ"] + self.idk = pos["dK"] + self.idv = pos["dV"] + self.idg = pos["dG"] + self.idb = pos["dBeta"] + self.idw = pos["dW"] + self.ids0 = pos.get("d_initial_state") + self.ia_log = pos.get("a_log") + self.idt_bias = pos.get("dt_bias") + self.ida_log = pos.get("d_a_log") + self.iddt_bias = pos.get("d_dt_bias") + + def run(self, views, workspace, stream) -> None: + q = views[self.iq] + k = views[self.ik] + v = views[self.iv] + g = views[self.ig] + beta = views[self.ibeta] + w = views[self.iw] + cu = views[self.icu] + do = views[self.ido] + state_checkpoints = views[self.ick] if self.ick is not None else None + state0 = views[self.is0] if self.is0 is not None else None + dstate_in = views[self.idfs] if self.idfs is not None else None + dq = views[self.idq] + dk = views[self.idk] + dv = views[self.idv] + dg = views[self.idg] + db = views[self.idb] + dw = views[self.idw] + dstate0 = views[self.ids0] if self.ids0 is not None else None + a_log = views[self.ia_log] if self.ia_log is not None else None + dt_bias = views[self.idt_bias] if self.idt_bias is not None else None + d_a_log = views[self.ida_log] if self.ida_log is not None else None + d_dt_bias = views[self.iddt_bias] if self.iddt_bias is not None else None stream = stream if stream is not None else 0 - HO, total = self.n_heads_out, self.total - K, V = q.shape[-1], v.shape[-1] - B = cu.shape[0] - 1 region = dict(zip(self.carve_names, workspace.carve(self.carve))) sched_regen = region["sched_regen"] sched_bwd = region["sched_bwd"] work_items = region["work_items"] work_count = region["work_count"] - self.build_split_table( - g, - cu, - work_items, - work_count, - ideal_chunks=self.ideal, - n_tiles=self.n_tiles, - num_sms=self.num_sm, - b_t=self.b_t, - chunk_scratch=region.get("chunk_scratch"), - item_scratch=region.get("item_scratch"), - log_gate=True, - sched_ctr=region["sched_all"], - split=self.split, - stream=stream, - ) + + if self.kcache is not None and (self.table is not None or not self.needs_table): + if self.needs_table: + self.run_table( + self.table, + g, + a_log, + dt_bias, + cu, + region.get("chunk_scratch"), + region.get("item_scratch"), + work_items, + work_count, + region["sched_all"], + stream, + ) + state0_io = None + if state0 is not None: + state0_io = region["state0_io"] + self.downcast_state(state0, state0_io, stream=stream) + if self.has_state_checkpoints: + checkpoint_series = state_checkpoints + else: + checkpoint_series = region["state_checkpoints"] + self.regen.run_recompute( + self.regen_cache, + k, + v, + g, + a_log if self.safe_gate else None, + dt_bias if self.safe_gate else None, + beta, + w, + cu, + state0, + None, + checkpoint_series, + work_items, + work_count, + sched_regen, + region["sched_all"], + region.get("item_scratch"), + region["regen_tensormaps"], + self.b_t, + stream, + ) + dq_out = region["dq_ho"] if self.fold_dq else dq + dk_out = region["dk_ho"] if self.fold_dk else dk + dv_out = region["dv_ho"] if self.fold_dv else dv + self.bwd.run_bwd( + self.kcache, + q, + k, + v, + g, + beta, + w, + do, + checkpoint_series, + dq_out, + dk_out, + dv_out, + dg, + db, + dw, + cu, + state0_io, + dstate0 if self.has_dstate0 else None, + dstate_in, + work_items, + work_count, + sched_bwd if self.bwd_dyn_sched else None, + region["sched_all"] if not self.order_in_regen else None, + region.get("item_scratch") if not self.order_in_regen else None, + region["bwd_tensormaps"], + self.scale, + stream, + a_log=a_log if self.safe_gate else None, + dt_bias=dt_bias if self.safe_gate else None, + ) + if self.safe_gate: + self.channel_gate_bwd( + dg, g, a_log, dt_bias, d_a_log, d_dt_bias, region["gate_part_a"], region["gate_part_dt"], self.gate_lower_bound, stream=stream + ) + if self.fold_dq or self.fold_dk or self.fold_dv: + for src_ho, dst in ((dq_out, dq), (dk_out, dk), (dv_out, dv)): + if src_ho is not dst: + self.head_group_reduce(src_ho, dst, stream=stream) + return + + if not self.needs_table: + self.table = None + else: + self.table = self.build_split_table( + g, + cu, + work_items, + work_count, + ideal_chunks=self.ideal, + n_tiles=self.n_tiles, + num_sms=self.num_sm, + b_t=self.b_t, + chunk_scratch=region.get("chunk_scratch"), + item_scratch=region.get("item_scratch"), + log_gate=True, + safe_gate=self.safe_gate, + a_log=a_log, + dt_bias=dt_bias, + gate_lower_bound=self.gate_lower_bound if self.safe_gate else None, + sched_ctr=region["sched_all"], + split=self.split, + stream=stream, + ) state0_io = None if state0 is not None: @@ -387,7 +591,7 @@ def __call__(self, node_buffers, *, workspace=None, stream=None) -> Any: checkpoint_series = state_checkpoints else: checkpoint_series = region["state_checkpoints"] - self.regen.chunk_gdn2_recompute_sm100( + self.regen_cache = self.regen.chunk_gdn2_recompute_sm100( k, v, g, @@ -399,9 +603,17 @@ def __call__(self, node_buffers, *, workspace=None, stream=None) -> Any: checkpoint_every_n_tokens=self.b_t, output_state_checkpoints=checkpoint_series, use_qk_l2norm_in_kernel=self.use_qk_l2norm, + safe_gate=self.safe_gate, + gate_lower_bound=self.gate_lower_bound, + a_log=a_log, + dt_bias=dt_bias, + use_beta_sigmoid=self.use_beta_sigmoid, work_items=work_items, work_count=work_count, sched_ctr=sched_regen, + sched_all=region["sched_all"], + work_item_scratch=region.get("item_scratch"), + order_in_prologue=True, tensormap_workspace=region["regen_tensormaps"], stream=stream, ) @@ -414,7 +626,7 @@ def __call__(self, node_buffers, *, workspace=None, stream=None) -> Any: if self.fold_dv: dv_out = region["dv_ho"] - self.bwd.chunk_gdn2_bwd_sm100( + self.kcache = self.bwd.chunk_gdn2_bwd_sm100( q, k, v, @@ -435,16 +647,26 @@ def __call__(self, node_buffers, *, workspace=None, stream=None) -> Any: d_initial_state=dstate0 if self.has_dstate0 else None, d_final_state=dstate_in, use_qk_l2norm_in_kernel=self.use_qk_l2norm, + safe_gate=self.safe_gate, + gate_lower_bound=self.gate_lower_bound, + a_log=a_log, + dt_bias=dt_bias, + use_beta_sigmoid=self.use_beta_sigmoid, work_items=work_items, work_count=work_count, sched_ctr=sched_bwd if self.bwd_dyn_sched else None, + sched_all=region["sched_all"] if not self.order_in_regen else None, + work_item_scratch=region.get("item_scratch") if not self.order_in_regen else None, + order_in_prologue=not self.order_in_regen, tensormap_workspace=region["bwd_tensormaps"], stream=stream, ) + if self.safe_gate: + self.channel_gate_bwd( + dg, g, a_log, dt_bias, d_a_log, d_dt_bias, region["gate_part_a"], region["gate_part_dt"], self.gate_lower_bound, stream=stream + ) if dq_out is not dq or dk_out is not dk or dv_out is not dv: - from .common.head_reduce import head_group_reduce - for src_ho, dst in ((dq_out, dq), (dk_out, dk), (dv_out, dv)): if src_ho is not dst: - head_group_reduce(src_ho, dst, stream=stream) + self.head_group_reduce(src_ho, dst, stream=stream) return None diff --git a/python/cudnn/linear_attention/frost/gdn_engine.py b/python/cudnn/linear_attention/frost/gdn_engine.py index d908011f2..be3cef8f5 100644 --- a/python/cudnn/linear_attention/frost/gdn_engine.py +++ b/python/cudnn/linear_attention/frost/gdn_engine.py @@ -10,7 +10,6 @@ from __future__ import annotations import math -from typing import Any from cudnn import behavior_note from cudnn.engines.base import BaseEngine, CompiledPlan @@ -19,7 +18,8 @@ from cudnn.frost.buffers import current_device_id from cudnn.frost.device import multiprocessor_count from cudnn.frost.workspace import WorkspaceLayout, carve_plan -from ..graph_analyzer import FrostLaPlan, frost_la_gate, require, analyze +from ..graph_analyzer import analyze +from .engine import FrostLaPlan, frost_la_gate def build_gdn(graph): @@ -47,17 +47,13 @@ class GdnFrostEngine(BaseEngine): elsewhere so ranking falls back to ``GdnCuTileEngine``.""" name = "gdn_frost" - behavior_notes = (behavior_note.RUNTIME_COMPILATION,) # JIT-compiled at build_plans() + behavior_notes = (behavior_note.RUNTIME_COMPILATION,) def check_support(self, graph) -> None: import cudnn facts = graph._facts_for(analyze) frost_la_gate("GdnFrostEngine", facts, "GDN") - if facts.use_qk_l2norm: - raise NotImplementedError("GdnFrostEngine: use_qk_l2norm is not supported (the kernel takes q/k as given)") - if facts.safe_gate: - raise NotImplementedError("GdnFrostEngine: safe_gate is not supported (no gate-activation path; the cuTile GDN engine serves it)") ckpt = facts.checkpoint_every_n_tokens if ckpt and (facts.is_bwd or ckpt % 64 != 0): raise NotImplementedError(f"GdnFrostEngine: checkpoint_every_n_tokens must be a positive multiple of 64 on the GDN node (got {ckpt})") @@ -66,21 +62,36 @@ def check_support(self, graph) -> None: fp32 = cudnn.data_type.FLOAT io = (cudnn.data_type.BFLOAT16, cudnn.data_type.HALF) state_dtypes = (fp32, cudnn.data_type.BFLOAT16) - require("GdnFrostEngine", "beta", facts.beta_dtype, fp32) - require("GdnFrostEngine", "initial_state", facts.state_dtype, state_dtypes) - require("GdnFrostEngine", "final_state", facts.final_state_dtype, state_dtypes) + beta_want = facts.io_dtype if facts.use_beta_sigmoid else fp32 + if beta_want is not None and facts.beta_dtype not in (beta_want, None): + raise NotImplementedError(f"GdnFrostEngine: 'beta' must be {beta_want} (io-dtype logits under use_beta_sigmoid), got {facts.beta_dtype}") + for port, got in (("a_log", facts.a_log_dtype), ("dt_bias", facts.dt_bias_dtype)): + if got not in (fp32, None): + raise NotImplementedError(f"GdnFrostEngine: '{port}' must be fp32, got {got}") + for port, got in (("initial_state", facts.state_dtype), ("final_state", facts.final_state_dtype)): + if got not in state_dtypes + (None,): + raise NotImplementedError(f"GdnFrostEngine: '{port}' must be fp32/bf16, got {got}") if not facts.state_pair_match: raise NotImplementedError("GdnFrostEngine: initial_state and final_state dtypes must match") if facts.is_bwd: - require("GdnFrostEngine", "dO", facts.do_dtype, io) - if facts.io_dtype is not None: - require("GdnFrostEngine", "state_checkpoints", facts.state_checkpoints_dtype, facts.io_dtype) - require("GdnFrostEngine", "d_final_state", facts.d_final_state_dtype, fp32) - require("GdnFrostEngine", "d_initial_state", facts.d_initial_state_dtype, fp32) - require("GdnFrostEngine", "dG", facts.dg_dtype, fp32) - require("GdnFrostEngine", "dBeta", facts.dbeta_dtype, fp32) - elif facts.io_dtype is not None: - require("GdnFrostEngine", "state_checkpoints", facts.state_checkpoints_out_dtype, facts.io_dtype) + for port, got in ( + ("d_a_log", facts.d_a_log_dtype), + ("d_dt_bias", facts.d_dt_bias_dtype), + ("d_final_state", facts.d_final_state_dtype), + ("d_initial_state", facts.d_initial_state_dtype), + ("dG", facts.dg_dtype), + ): + if got not in (fp32, None): + raise NotImplementedError(f"GdnFrostEngine: '{port}' must be fp32, got {got}") + if facts.do_dtype not in io + (None,): + raise NotImplementedError(f"GdnFrostEngine: 'dO' must be fp16/bf16, got {facts.do_dtype}") + if facts.io_dtype is not None and facts.state_checkpoints_dtype not in (facts.io_dtype, None): + raise NotImplementedError("GdnFrostEngine: 'state_checkpoints' must match the io dtype") + dbeta_want = facts.io_dtype if facts.use_beta_sigmoid and facts.io_dtype is not None else fp32 + if facts.dbeta_dtype not in (dbeta_want, None): + raise NotImplementedError(f"GdnFrostEngine: 'dBeta' must be {dbeta_want}, got {facts.dbeta_dtype}") + elif facts.io_dtype is not None and facts.state_checkpoints_out_dtype not in (facts.io_dtype, None): + raise NotImplementedError("GdnFrostEngine: 'state_checkpoints' must match the io dtype") def build_plan(self, graph, plan, ctx=None) -> CompiledPlan: return FrostLaPlan(build_gdn(graph)) @@ -90,19 +101,30 @@ class CompiledGdn: """Compiled FROST GDN plan: a callable over the resolved node buffers.""" def __init__(self, node, kernel_mod): - from .common.split_k import WORK_ITEM_FIELDS, build_split_table, chunk_scratch_rows, compute_ideal_chunks, max_work_items + from .common.split_k import WORK_ITEM_FIELDS, build_split_table, chunk_scratch_rows, compute_ideal_chunks, max_work_items, run_table self.node = node self.kernel = kernel_mod self.build_split_table = build_split_table + self.run_table = run_table + self.table = None + self.kcache = None self.plan_name = "GdnFrostEngine (GDN)" + from .common.l2norm import l2norm_qk + + self.l2norm_qk = l2norm_qk scale = node.params.get("scale") self.scale = float(scale) if scale is not None else 1.0 / math.sqrt(node.inputs["q"].dim[-1]) + self.use_qk_l2norm = bool(node.params.get("use_qk_l2norm", False)) + self.safe_gate = bool(node.params.get("safe_gate", False)) + self.use_beta_sigmoid = bool(node.params.get("use_beta_sigmoid", False)) q, v, g = node.inputs["q"], node.inputs["v"], node.inputs["g"] self.b_t = kernel_mod.CFG.B_T total = q.dim[0] HO = g.dim[1] + HQ, HK = q.dim[1], node.inputs["k"].dim[1] + self.io_name = "float16" if q.get_data_type().name == "HALF" else "bfloat16" B = node.inputs["cu_seqlens"].dim[0] - 1 self.has_final_state = "final_state" in node.outputs self.ckpt = int(node.params.get("checkpoint_every_n_tokens", 0) or 0) @@ -116,7 +138,7 @@ def __init__(self, node, kernel_mod): self.tensormap_words = tensormap_workspace_bytes(kernel_mod, B) // 8 self.off_tensormaps = layout.add(self.tensormap_words * 8) - self.off_sched = layout.add(8) # [ticket, done] for the dynamic scheduler + self.off_sched = layout.add(8) self.num_sm = multiprocessor_count(current_device_id()) self.n_tiles = B * HO self.n_heads_out = HO @@ -132,6 +154,12 @@ def __init__(self, node, kernel_mod): self.off_item_scratch = layout.add(self.work_item_rows * WORK_ITEM_FIELDS * 4) self.chunk_scratch_rows = chunk_scratch_rows(total, B, self.b_t) self.off_chunk_scratch = layout.add(self.chunk_scratch_rows * HO * 4) + if self.use_qk_l2norm: + self.off_q_n = layout.add(total * HQ * 128 * 2) + self.off_k_n = layout.add(total * HK * 128 * 2) + self.off_inv_q = layout.add(total * HQ * 4) + self.off_inv_k = layout.add(total * HK * 4) + self.needs_table = self.split self.ws_bytes = layout.size regions = [ (self.off_tensormaps, "int64", (self.tensormap_words,)), @@ -144,47 +172,110 @@ def __init__(self, node, kernel_mod): (self.off_item_scratch, "int32", (self.work_item_rows, WORK_ITEM_FIELDS)), (self.off_chunk_scratch, "float32", (self.chunk_scratch_rows, self.n_heads_out)), ] + if self.use_qk_l2norm: + regions += [ + (self.off_q_n, self.io_name, (total, HQ, 128)), + (self.off_k_n, self.io_name, (total, HK, 128)), + (self.off_inv_q, "float32", (total, HQ)), + (self.off_inv_k, "float32", (total, HK)), + ] self.carve = carve_plan("GdnFrostEngine (GDN)", regions) def workspace_bytes(self) -> int: return self.ws_bytes - def __call__(self, node_buffers, *, workspace=None, stream=None) -> Any: - nb = node_buffers[self.node] - q = nb.inputs["q"] - k = nb.inputs["k"] - v = nb.inputs["v"] - g = nb.inputs["g"] - beta = nb.inputs["beta"] - cu = nb.inputs["cu_seqlens"] - state0 = nb.inputs.get("initial_state") - o = nb.outputs["O"] - final_state = nb.outputs["final_state"] if self.has_final_state else None - state_checkpoints = nb.outputs["state_checkpoints"] if self.has_state_checkpoints else None + def bind(self, names) -> None: + pos = {name: i for i, name in enumerate(names)} + self.iq = pos["q"] + self.ik = pos["k"] + self.iv = pos["v"] + self.ig = pos["g"] + self.ibeta = pos["beta"] + self.icu = pos["cu_seqlens"] + self.is0 = pos.get("initial_state") + self.io_ = pos["O"] + self.ifs = pos.get("final_state") + self.ick = pos.get("state_checkpoints") + self.ia_log = pos.get("a_log") + self.idt_bias = pos.get("dt_bias") + + def run(self, views, workspace, stream) -> None: + q = views[self.iq] + k = views[self.ik] + v = views[self.iv] + g = views[self.ig] + beta = views[self.ibeta] + cu = views[self.icu] + state0 = views[self.is0] if self.is0 is not None else None + o = views[self.io_] + final_state = views[self.ifs] if self.ifs is not None else None + state_checkpoints = views[self.ick] if self.ick is not None else None + a_log = views[self.ia_log] if self.ia_log is not None else None + dt_bias = views[self.idt_bias] if self.idt_bias is not None else None stream = stream if stream is not None else 0 + carved = workspace.carve(self.carve) if self.split: - tensormaps, sched_ctr, work_items, work_count, item_scratch, chunk_scratch = workspace.carve(self.carve) + tensormaps, sched_ctr, work_items, work_count, item_scratch, chunk_scratch, *l2n = carved else: - tensormaps, sched_ctr, work_items, work_count = workspace.carve(self.carve) + tensormaps, sched_ctr, work_items, work_count, *l2n = carved item_scratch = chunk_scratch = None - self.build_split_table( - g, - cu, - work_items, - work_count, - ideal_chunks=self.ideal, - n_tiles=self.n_tiles, - num_sms=self.num_sm, - b_t=self.b_t, - chunk_scratch=chunk_scratch, - item_scratch=item_scratch, - log_gate=True, - sched_ctr=sched_ctr, - split=self.split, - stream=stream, - ) + if self.use_qk_l2norm: + q_n, k_n, inv_q, inv_k = l2n + self.l2norm_qk(q, k, q_n, k_n, inv_q, inv_k, stream=stream) + q, k = q_n, k_n + + if self.kcache is not None and (self.table is not None or not self.needs_table): + if self.needs_table: + self.run_table(self.table, g, a_log, dt_bias, cu, chunk_scratch, item_scratch, work_items, work_count, sched_ctr, stream) + self.kernel.run_prefill( + self.kcache, + q, + k, + v, + g, + beta, + o, + cu, + state0, + final_state, + state_checkpoints, + work_items, + work_count, + sched_ctr, + item_scratch, + tensormaps, + self.ckpt, + self.scale, + stream, + a_log=a_log if self.safe_gate else None, + dt_bias=dt_bias if self.safe_gate else None, + ) + return + + if not self.needs_table: + self.table = None + else: + self.table = self.build_split_table( + g, + cu, + work_items, + work_count, + ideal_chunks=self.ideal, + n_tiles=self.n_tiles, + num_sms=self.num_sm, + b_t=self.b_t, + chunk_scratch=chunk_scratch, + item_scratch=item_scratch, + log_gate=True, + safe_gate=self.safe_gate, + a_log=a_log, + dt_bias=dt_bias, + sched_ctr=sched_ctr, + split=self.split, + stream=stream, + ) - self.kernel.chunk_gdn_sm100( + self.kcache = self.kernel.chunk_gdn_sm100( q, k, v, @@ -201,6 +292,11 @@ def __call__(self, node_buffers, *, workspace=None, stream=None) -> Any: checkpoint_every_n_tokens=self.ckpt, output_state_checkpoints=state_checkpoints, log_gate=True, + safe_gate=self.safe_gate, + a_log=a_log, + dt_bias=dt_bias, + use_beta_sigmoid=self.use_beta_sigmoid, + work_item_scratch=item_scratch, workspace=tensormaps, stream=stream, ) @@ -215,23 +311,38 @@ class CompiledGdnBwd: (checkpoint-only) kernel when the port is absent.""" def __init__(self, node, bwd_mod, regen_mod): - from .common.split_k import WORK_ITEM_FIELDS, build_split_table, chunk_scratch_rows, compute_ideal_chunks, max_work_items + from .common.split_k import WORK_ITEM_FIELDS, build_split_table, chunk_scratch_rows, compute_ideal_chunks, max_work_items, run_table self.node = node self.bwd = bwd_mod self.regen = regen_mod self.build_split_table = build_split_table + self.run_table = run_table + self.table = None + self.kcache = None + self.regen_cache = None self.plan_name = "GdnFrostEngine (GDN_BWD)" from .common.downcast import downcast_state + from .common.gate_bwd import scalar_gate_bwd, scalar_gate_blocks + from .common.head_reduce import head_group_reduce from .common.host import tensormap_workspace_bytes + from .common.l2norm import l2norm_qk, l2norm_qk_bwd self.downcast_state = downcast_state + self.head_group_reduce = head_group_reduce + self.scalar_gate_bwd = scalar_gate_bwd + self.l2norm_qk = l2norm_qk + self.l2norm_qk_bwd = l2norm_qk_bwd scale = node.params.get("scale") self.scale = float(scale) if scale is not None else 1.0 / math.sqrt(node.inputs["q"].dim[-1]) + self.use_qk_l2norm = bool(node.params.get("use_qk_l2norm", False)) + self.safe_gate = bool(node.params.get("safe_gate", False)) + self.use_beta_sigmoid = bool(node.params.get("use_beta_sigmoid", False)) q, v, g = node.inputs["q"], node.inputs["v"], node.inputs["g"] self.b_t = bwd_mod.CFG.B_T total = q.dim[0] + self.gate_bwd_blocks = scalar_gate_blocks(total) K, V = q.dim[-1], v.dim[-1] HQ, HV = q.dim[1], v.dim[1] HO = g.dim[1] @@ -243,14 +354,12 @@ def __init__(self, node, bwd_mod, regen_mod): self.num_sm = multiprocessor_count(current_device_id()) self.bwd_dyn_sched = B * HO <= self.num_sm self.batch_invariant = bool(node.params.get("batch_invariant", False)) - # cuts never in batch-invariant mode: whole-sequence items keep each - # sequence's math independent of the batch composition + # cuts never in batch-invariant mode self.split = not self.batch_invariant layout = WorkspaceLayout() - self.off_sched = layout.add(16) # one [ticket, done] ring each for the regen and bwd kernels + self.off_sched = layout.add(16) self.tensormap_words = tensormap_workspace_bytes(bwd_mod, B) // 8 self.off_tensormaps = layout.add(self.tensormap_words * 8) - # chunk-0 entering state, io dtype (downcast initial_state; absent = in-kernel zeros) self.off_state0_io = layout.add(B * HO * K * V * 2) if self.has_state0 else None if not self.has_state_checkpoints: self.state_checkpoints_rows = max(total // self.b_t + B, 1) @@ -267,6 +376,14 @@ def __init__(self, node, bwd_mod, regen_mod): self.off_dk_ho = layout.add(total * HO * K * 2) if self.fold_dv: self.off_dv_ho = layout.add(total * HO * V * 2) + if self.use_qk_l2norm: + self.off_q_n = layout.add(total * HQ * 128 * 2) + self.off_k_n = layout.add(total * HK * 128 * 2) + self.off_inv_q = layout.add(total * HQ * 4) + self.off_inv_k = layout.add(total * HK * 4) + if self.safe_gate: + self.off_gate_part_a = layout.add(self.gate_bwd_blocks * HO * 4) + self.off_gate_part_dt = layout.add(self.gate_bwd_blocks * HO * 4) self.n_tiles = B * HO if self.split: self.ideal = compute_ideal_chunks(total, HO, self.num_sm, self.b_t) @@ -280,6 +397,9 @@ def __init__(self, node, bwd_mod, regen_mod): self.off_item_scratch = layout.add(self.work_item_rows * WORK_ITEM_FIELDS * 4) self.chunk_scratch_rows = chunk_scratch_rows(total, B, self.b_t) self.off_chunk_scratch = layout.add(self.chunk_scratch_rows * HO * 4) + self.needs_table = self.split + self.regen_orders = not self.has_state_checkpoints + self.bwd_orders = self.has_state_checkpoints self.n_heads_out, self.total = HO, total self.ws_bytes = layout.size regions = [ @@ -304,61 +424,190 @@ def __init__(self, node, bwd_mod, regen_mod): regions.append(("dk_ho", self.off_dk_ho, self.io_name, (total, HO, K))) if self.fold_dv: regions.append(("dv_ho", self.off_dv_ho, self.io_name, (total, HO, V))) + if self.use_qk_l2norm: + regions.append(("q_n", self.off_q_n, self.io_name, (total, HQ, 128))) + regions.append(("k_n", self.off_k_n, self.io_name, (total, HK, 128))) + regions.append(("inv_q", self.off_inv_q, "float32", (total, HQ))) + regions.append(("inv_k", self.off_inv_k, "float32", (total, HK))) + if self.safe_gate: + regions.append(("gate_part_a", self.off_gate_part_a, "float32", (self.gate_bwd_blocks * HO,))) + regions.append(("gate_part_dt", self.off_gate_part_dt, "float32", (self.gate_bwd_blocks * HO,))) self.carve_names = [name for name, _off, _dt, _shape in regions] self.carve = carve_plan("GdnFrostEngine (GDN_BWD)", [(off, dt, shape) for _name, off, dt, shape in regions]) def workspace_bytes(self) -> int: return self.ws_bytes - def __call__(self, node_buffers, *, workspace=None, stream=None) -> Any: - nb = node_buffers[self.node] - q = nb.inputs["q"] - k = nb.inputs["k"] - v = nb.inputs["v"] - g = nb.inputs["g"] - beta = nb.inputs["beta"] - cu = nb.inputs["cu_seqlens"] - do = nb.inputs["dO"] - state_checkpoints = nb.inputs.get("state_checkpoints") - state0 = nb.inputs.get("initial_state") - dstate_in = nb.inputs.get("d_final_state") - dq = nb.outputs["dQ"] - dk = nb.outputs["dK"] - dv = nb.outputs["dV"] - dg = nb.outputs["dG"] - db = nb.outputs["dBeta"] - dstate0 = nb.outputs.get("d_initial_state") - HO, total = self.n_heads_out, self.total - K, V = q.shape[-1], v.shape[-1] - B = cu.shape[0] - 1 - + def bind(self, names) -> None: + pos = {name: i for i, name in enumerate(names)} + self.iq = pos["q"] + self.ik = pos["k"] + self.iv = pos["v"] + self.ig = pos["g"] + self.ibeta = pos["beta"] + self.icu = pos["cu_seqlens"] + self.ido = pos["dO"] + self.ick = pos.get("state_checkpoints") + self.is0 = pos.get("initial_state") + self.idfs = pos.get("d_final_state") + self.idq = pos["dQ"] + self.idk = pos["dK"] + self.idv = pos["dV"] + self.idg = pos["dG"] + self.idb = pos["dBeta"] + self.ids0 = pos.get("d_initial_state") + self.ia_log = pos.get("a_log") + self.idt_bias = pos.get("dt_bias") + self.ida_log = pos.get("d_a_log") + self.iddt_bias = pos.get("d_dt_bias") + + def run(self, views, workspace, stream) -> None: + q = views[self.iq] + k = views[self.ik] + v = views[self.iv] + g = views[self.ig] + beta = views[self.ibeta] + cu = views[self.icu] + do = views[self.ido] + state_checkpoints = views[self.ick] if self.ick is not None else None + state0 = views[self.is0] if self.is0 is not None else None + dstate_in = views[self.idfs] if self.idfs is not None else None + dq = views[self.idq] + dk = views[self.idk] + dv = views[self.idv] + dg = views[self.idg] + db = views[self.idb] + dstate0 = views[self.ids0] if self.ids0 is not None else None + a_log = views[self.ia_log] if self.ia_log is not None else None + dt_bias = views[self.idt_bias] if self.idt_bias is not None else None + d_a_log = views[self.ida_log] if self.ida_log is not None else None + d_dt_bias = views[self.iddt_bias] if self.iddt_bias is not None else None stream = stream if stream is not None else 0 + region = dict(zip(self.carve_names, workspace.carve(self.carve))) sched_regen = region["sched_regen"] sched_bwd = region["sched_bwd"] work_items = region["work_items"] work_count = region["work_count"] - self.build_split_table( - g, - cu, - work_items, - work_count, - ideal_chunks=self.ideal, - n_tiles=self.n_tiles, - num_sms=self.num_sm, - b_t=self.b_t, - chunk_scratch=region.get("chunk_scratch"), - item_scratch=region.get("item_scratch"), - log_gate=True, - sched_ctr=region["sched_all"], - split=self.split, - stream=stream, - ) + if self.use_qk_l2norm: + self.l2norm_qk(q, k, region["q_n"], region["k_n"], region["inv_q"], region["inv_k"], stream=stream) + q, k = region["q_n"], region["k_n"] + + if self.kcache is not None and (self.table is not None or not self.needs_table): + if self.needs_table: + self.run_table( + self.table, + g, + a_log, + dt_bias, + cu, + region.get("chunk_scratch"), + region.get("item_scratch"), + work_items, + work_count, + region["sched_all"], + stream, + ) + state0_io = None + if state0 is not None: + if state0.dtype == self.io_name and buffers.is_contiguous(tuple(state0.shape), state0.stride()) and state0.data_ptr() % 16 == 0: + state0_io = state0 + else: + state0_io = region["state0_io"] + self.downcast_state(state0, state0_io, stream=stream) + if self.has_state_checkpoints: + checkpoint_series = state_checkpoints + else: + checkpoint_series = region["state_checkpoints"] + self.regen.run_recompute( + self.regen_cache, + k, + v, + g, + beta, + cu, + state0, + None, + checkpoint_series, + work_items, + work_count, + sched_regen, + region["sched_all"] if self.regen_orders else None, + region.get("item_scratch") if self.regen_orders else None, + region["regen_tensormaps"], + self.b_t, + stream, + a_log=a_log if self.safe_gate else None, + dt_bias=dt_bias if self.safe_gate else None, + ) + dq_out = region["dq_ho"] if self.fold_dq else dq + dk_out = region["dk_ho"] if self.fold_dk else dk + dv_out = region["dv_ho"] if self.fold_dv else dv + self.bwd.run_bwd( + self.kcache, + q, + k, + v, + g, + beta, + do, + checkpoint_series, + dq_out, + dk_out, + dv_out, + dg, + db, + cu, + state0_io, + dstate0, + dstate_in, + work_items, + work_count, + sched_bwd if self.bwd_dyn_sched else None, + region["sched_all"] if self.bwd_orders else None, + region.get("item_scratch") if self.bwd_orders else None, + region["tensormaps"], + self.scale, + stream, + a_log=a_log if self.safe_gate else None, + dt_bias=dt_bias if self.safe_gate else None, + ) + if self.safe_gate: + self.scalar_gate_bwd(dg, g, a_log, dt_bias, d_a_log, d_dt_bias, region["gate_part_a"], region["gate_part_dt"], stream=stream) + if self.fold_dq or self.fold_dk or self.fold_dv: + for src_ho, dst in ((dq_out, dq), (dk_out, dk), (dv_out, dv)): + if src_ho is not dst: + self.head_group_reduce(src_ho, dst, stream=stream) + if self.use_qk_l2norm: + self.l2norm_qk_bwd(dq, dk, region["q_n"], region["k_n"], region["inv_q"], region["inv_k"], stream=stream) + return + + if not self.needs_table: + self.table = None + else: + self.table = self.build_split_table( + g, + cu, + work_items, + work_count, + ideal_chunks=self.ideal, + n_tiles=self.n_tiles, + num_sms=self.num_sm, + b_t=self.b_t, + chunk_scratch=region.get("chunk_scratch"), + item_scratch=region.get("item_scratch"), + log_gate=True, + safe_gate=self.safe_gate, + a_log=a_log, + dt_bias=dt_bias, + sched_ctr=region["sched_all"], + split=self.split, + stream=stream, + ) state0_io = None if state0 is not None: if state0.dtype == self.io_name and buffers.is_contiguous(tuple(state0.shape), state0.stride()) and state0.data_ptr() % 16 == 0: - # already the staged form: io dtype, compact, aligned — bind directly state0_io = state0 else: state0_io = region["state0_io"] @@ -367,7 +616,7 @@ def __call__(self, node_buffers, *, workspace=None, stream=None) -> Any: checkpoint_series = state_checkpoints else: checkpoint_series = region["state_checkpoints"] - self.regen.chunk_gdn_recompute_sm100( + self.regen_cache = self.regen.chunk_gdn_recompute_sm100( k, v, g, @@ -377,9 +626,16 @@ def __call__(self, node_buffers, *, workspace=None, stream=None) -> Any: None, checkpoint_every_n_tokens=self.b_t, output_state_checkpoints=checkpoint_series, + safe_gate=self.safe_gate, + a_log=a_log, + dt_bias=dt_bias, + use_beta_sigmoid=self.use_beta_sigmoid, work_items=work_items, work_count=work_count, sched_ctr=sched_regen, + sched_all=region["sched_all"] if self.regen_orders else None, + work_item_scratch=region.get("item_scratch") if self.regen_orders else None, + order_in_prologue=self.regen_orders, log_gate=True, workspace=region["regen_tensormaps"], stream=stream, @@ -392,7 +648,7 @@ def __call__(self, node_buffers, *, workspace=None, stream=None) -> Any: dk_out = region["dk_ho"] if self.fold_dv: dv_out = region["dv_ho"] - self.bwd.chunk_gdn_bwd_sm100( + self.kcache = self.bwd.chunk_gdn_bwd_sm100( q, k, v, @@ -410,17 +666,26 @@ def __call__(self, node_buffers, *, workspace=None, stream=None) -> Any: initial_state=state0_io, d_initial_state=dstate0, d_final_state=dstate_in, + safe_gate=self.safe_gate, + a_log=a_log, + dt_bias=dt_bias, + use_beta_sigmoid=self.use_beta_sigmoid, work_items=work_items, work_count=work_count, sched_ctr=sched_bwd if self.bwd_dyn_sched else None, + sched_all=region["sched_all"] if self.bwd_orders else None, + work_item_scratch=region.get("item_scratch") if self.bwd_orders else None, + order_in_prologue=self.bwd_orders, log_gate=True, workspace=region["tensormaps"], stream=stream, ) + if self.safe_gate: + self.scalar_gate_bwd(dg, g, a_log, dt_bias, d_a_log, d_dt_bias, region["gate_part_a"], region["gate_part_dt"], stream=stream) if dq_out is not dq or dk_out is not dk or dv_out is not dv: - from .common.head_reduce import head_group_reduce - for src_ho, dst in ((dq_out, dq), (dk_out, dk), (dv_out, dv)): if src_ho is not dst: - head_group_reduce(src_ho, dst, stream=stream) + self.head_group_reduce(src_ho, dst, stream=stream) + if self.use_qk_l2norm: + self.l2norm_qk_bwd(dq, dk, region["q_n"], region["k_n"], region["inv_q"], region["inv_k"], stream=stream) return None diff --git a/python/cudnn/linear_attention/frost/kda_engine.py b/python/cudnn/linear_attention/frost/kda_engine.py index 471c9c5cd..28ebbec8a 100644 --- a/python/cudnn/linear_attention/frost/kda_engine.py +++ b/python/cudnn/linear_attention/frost/kda_engine.py @@ -11,7 +11,6 @@ from __future__ import annotations import math -from typing import Any from cudnn import behavior_note from cudnn.engines.base import BaseEngine, CompiledPlan @@ -19,7 +18,8 @@ from cudnn.frost.buffers import current_device_id from cudnn.frost.device import multiprocessor_count from cudnn.frost.workspace import WorkspaceLayout, carve_plan -from ..graph_analyzer import FrostLaPlan, frost_la_gate, require, analyze +from ..graph_analyzer import analyze +from .engine import FrostLaPlan, frost_la_gate def build_kda(graph): @@ -47,7 +47,7 @@ class KdaFrostEngine(BaseEngine): forward and KDA_BWD (with a forward checkpoint recompute when ``state_checkpoints`` is absent).""" name = "kda_frost" - behavior_notes = (behavior_note.RUNTIME_COMPILATION,) # JIT-compiled at build_plans() + behavior_notes = (behavior_note.RUNTIME_COMPILATION,) def check_support(self, graph) -> None: import cudnn @@ -59,32 +59,37 @@ def check_support(self, graph) -> None: raise NotImplementedError(f"KdaFrostEngine: checkpoint_every_n_tokens must be a positive multiple of 16 on the KDA node (got {ckpt})") if not facts.gates_at_ho: raise NotImplementedError(f"KdaFrostEngine: g/beta must carry HO = max(q, v) heads ({facts.h_o})") - if facts.is_bwd and (facts.use_beta_sigmoid or facts.safe_gate): - raise NotImplementedError("KdaFrostEngine: use_beta_sigmoid/safe_gate are forward-node attributes") fp32 = cudnn.data_type.FLOAT - if not facts.is_bwd and facts.use_beta_sigmoid: - # in-kernel sigmoid: Beta arrives as io-dtype logits - if facts.beta_dtype not in (facts.io_dtype, None): - raise NotImplementedError("KdaFrostEngine: use_beta_sigmoid takes io-dtype beta logits") - else: - require("KdaFrostEngine", "beta", facts.beta_dtype, fp32) - require("KdaFrostEngine", "a_log", facts.a_log_dtype, fp32) - require("KdaFrostEngine", "dt_bias", facts.dt_bias_dtype, fp32) + beta_want = facts.io_dtype if facts.use_beta_sigmoid else fp32 + if beta_want is not None and facts.beta_dtype not in (beta_want, None): + raise NotImplementedError(f"KdaFrostEngine: 'beta' must be {beta_want} (io-dtype logits under use_beta_sigmoid), got {facts.beta_dtype}") + for port, got in (("a_log", facts.a_log_dtype), ("dt_bias", facts.dt_bias_dtype)): + if got not in (fp32, None): + raise NotImplementedError(f"KdaFrostEngine: '{port}' must be fp32, got {got}") if facts.is_bwd: + for port, got in ( + ("d_a_log", facts.d_a_log_dtype), + ("d_dt_bias", facts.d_dt_bias_dtype), + ("initial_state", facts.state_dtype), + ("d_final_state", facts.d_final_state_dtype), + ("d_initial_state", facts.d_initial_state_dtype), + ("dG", facts.dg_dtype), + ): + if got not in (fp32, None): + raise NotImplementedError(f"KdaFrostEngine: '{port}' must be fp32, got {got}") for port, got in (("dO", facts.do_dtype), ("state_checkpoints", facts.state_checkpoints_dtype)): if got not in (facts.io_dtype, None): raise NotImplementedError(f"KdaFrostEngine: '{port}' must match the io dtype") - require("KdaFrostEngine", "initial_state", facts.state_dtype, fp32) - require("KdaFrostEngine", "d_final_state", facts.d_final_state_dtype, fp32) - require("KdaFrostEngine", "d_initial_state", facts.d_initial_state_dtype, fp32) - require("KdaFrostEngine", "dG", facts.dg_dtype, fp32) - require("KdaFrostEngine", "dBeta", facts.dbeta_dtype, fp32) + dbeta_want = facts.io_dtype if facts.use_beta_sigmoid and facts.io_dtype is not None else fp32 + if facts.dbeta_dtype not in (dbeta_want, None): + raise NotImplementedError(f"KdaFrostEngine: 'dBeta' must be {dbeta_want}, got {facts.dbeta_dtype}") else: state_dtypes = (fp32, cudnn.data_type.BFLOAT16) - require("KdaFrostEngine", "initial_state", facts.state_dtype, state_dtypes) - if facts.io_dtype is not None: - require("KdaFrostEngine", "state_checkpoints", facts.state_checkpoints_out_dtype, facts.io_dtype) - require("KdaFrostEngine", "final_state", facts.final_state_dtype, state_dtypes) + for port, got in (("initial_state", facts.state_dtype), ("final_state", facts.final_state_dtype)): + if got not in state_dtypes + (None,): + raise NotImplementedError(f"KdaFrostEngine: '{port}' must be fp32/bf16, got {got}") + if facts.io_dtype is not None and facts.state_checkpoints_out_dtype not in (facts.io_dtype, None): + raise NotImplementedError("KdaFrostEngine: 'state_checkpoints' must match the io dtype") if not facts.state_pair_match: raise NotImplementedError("KdaFrostEngine: initial_state and final_state dtypes must match") @@ -96,11 +101,14 @@ class CompiledKda: """Compiled FROST KDA plan: a callable over the resolved node buffers.""" def __init__(self, node, kernel_mod): - from .common.split_k import WORK_ITEM_FIELDS, build_split_table, chunk_scratch_rows, compute_ideal_chunks, max_work_items + from .common.split_k import WORK_ITEM_FIELDS, build_split_table, chunk_scratch_rows, compute_ideal_chunks, max_work_items, run_table self.node = node self.kernel = kernel_mod self.build_split_table = build_split_table + self.run_table = run_table + self.table = None + self.kcache = None self.plan_name = "KdaFrostEngine (KDA)" scale = node.params.get("scale") self.scale = float(scale) if scale is not None else 1.0 / math.sqrt(node.inputs["q"].dim[-1]) @@ -122,7 +130,7 @@ def __init__(self, node, kernel_mod): HO = g.dim[1] B = node.inputs["cu_seqlens"].dim[0] - 1 layout = WorkspaceLayout() - self.off_sched = layout.add(8) # [ticket, done] for the dynamic scheduler + self.off_sched = layout.add(8) self.num_sm = multiprocessor_count(current_device_id()) self.n_tiles = B * HO self.n_heads_out = HO @@ -142,6 +150,7 @@ def __init__(self, node, kernel_mod): self.tensormap_bytes = tensormap_workspace_bytes(kernel_mod, B) self.off_tensormaps = layout.add(self.tensormap_bytes, align=128) + self.needs_table = self.split self.ws_bytes = layout.size regions = [ (self.off_sched, "int32", (2,)), @@ -159,20 +168,34 @@ def __init__(self, node, kernel_mod): def workspace_bytes(self) -> int: return self.ws_bytes - def __call__(self, node_buffers, *, workspace=None, stream=None) -> Any: - nb = node_buffers[self.node] - q = nb.inputs["q"] - k = nb.inputs["k"] - v = nb.inputs["v"] - g = nb.inputs["g"] - beta = nb.inputs["beta"] - cu = nb.inputs["cu_seqlens"] - state0 = nb.inputs.get("initial_state") - o = nb.outputs["O"] - final_state = nb.outputs["final_state"] if self.has_final_state else None - state_checkpoints = nb.outputs["state_checkpoints"] if self.has_state_checkpoints else None - a_log = nb.inputs.get("a_log") - dt_bias = nb.inputs.get("dt_bias") + def bind(self, names) -> None: + pos = {name: i for i, name in enumerate(names)} + self.iq = pos["q"] + self.ik = pos["k"] + self.iv = pos["v"] + self.ig = pos["g"] + self.ibeta = pos["beta"] + self.icu = pos["cu_seqlens"] + self.is0 = pos.get("initial_state") + self.io_ = pos["O"] + self.ifs = pos.get("final_state") + self.ick = pos.get("state_checkpoints") + self.ia_log = pos.get("a_log") + self.idt_bias = pos.get("dt_bias") + + def run(self, views, workspace, stream) -> None: + q = views[self.iq] + k = views[self.ik] + v = views[self.iv] + g = views[self.ig] + beta = views[self.ibeta] + cu = views[self.icu] + state0 = views[self.is0] if self.is0 is not None else None + o = views[self.io_] + final_state = views[self.ifs] if self.ifs is not None else None + state_checkpoints = views[self.ick] if self.ick is not None else None + a_log = views[self.ia_log] if self.ia_log is not None else None + dt_bias = views[self.idt_bias] if self.idt_bias is not None else None stream = stream if stream is not None else 0 if self.split: @@ -180,32 +203,63 @@ def __call__(self, node_buffers, *, workspace=None, stream=None) -> Any: else: sched_ctr, work_items, work_count, tensormaps = workspace.carve(self.carve) item_scratch = chunk_scratch = None - self.build_split_table( - g, - cu, - work_items, - work_count, - ideal_chunks=self.ideal, - n_tiles=self.n_tiles, - num_sms=self.num_sm, - b_t=self.b_t, - chunk_scratch=chunk_scratch, - item_scratch=item_scratch, - log_gate=True, - safe_gate=self.safe_gate, - a_log=a_log, - dt_bias=dt_bias, - gate_lower_bound=self.gate_lower_bound if self.safe_gate else None, - sched_ctr=sched_ctr, - split=self.split, - stream=stream, - ) + + if self.kcache is not None and (self.table is not None or not self.needs_table): + if self.needs_table: + self.run_table(self.table, g, a_log, dt_bias, cu, chunk_scratch, item_scratch, work_items, work_count, sched_ctr, stream) + self.kernel.run_prefill( + self.kcache, + q, + k, + v, + g, + a_log if self.safe_gate else None, + dt_bias if self.safe_gate else None, + beta, + cu, + state0, + o, + final_state, + state_checkpoints, + work_items, + work_count, + sched_ctr, + item_scratch, + tensormaps, + self.ckpt if self.has_state_checkpoints else 0, + self.scale, + stream, + ) + return + + if not self.needs_table: + self.table = None + else: + self.table = self.build_split_table( + g, + cu, + work_items, + work_count, + ideal_chunks=self.ideal, + n_tiles=self.n_tiles, + num_sms=self.num_sm, + b_t=self.b_t, + chunk_scratch=chunk_scratch, + item_scratch=item_scratch, + log_gate=True, + safe_gate=self.safe_gate, + a_log=a_log, + dt_bias=dt_bias, + gate_lower_bound=self.gate_lower_bound if self.safe_gate else None, + sched_ctr=sched_ctr, + split=self.split, + stream=stream, + ) ckpt_kwargs = {} if self.has_state_checkpoints: - # the kernel derives the per-sequence checkpoint entry offsets on device ckpt_kwargs = dict(checkpoint_every_n_tokens=self.ckpt, output_state_checkpoints=state_checkpoints) - self.kernel.chunk_kda_sm100( + self.kcache = self.kernel.chunk_kda_sm100( q, k, v, @@ -225,6 +279,7 @@ def __call__(self, node_buffers, *, workspace=None, stream=None) -> Any: work_items=work_items, work_count=work_count, sched_ctr=sched_ctr, + work_item_scratch=item_scratch, tensormap_workspace=tensormaps, **ckpt_kwargs, stream=stream, @@ -240,20 +295,33 @@ class CompiledKdaBwd: counts.""" def __init__(self, node, bwd_mod, regen_mod): - from .common.split_k import WORK_ITEM_FIELDS, build_split_table, chunk_scratch_rows, compute_ideal_chunks, max_work_items + from .common.split_k import WORK_ITEM_FIELDS, build_split_table, chunk_scratch_rows, compute_ideal_chunks, max_work_items, run_table self.node = node self.bwd = bwd_mod self.regen = regen_mod self.build_split_table = build_split_table + self.run_table = run_table + self.table = None + self.kcache = None + self.regen_cache = None self.plan_name = "KdaFrostEngine (KDA_BWD)" from .common.downcast import downcast_state + from .common.gate_bwd import GATE_BWD_BLOCKS, channel_gate_bwd + from .common.head_reduce import head_group_reduce from .common.host import tensormap_workspace_bytes self.downcast_state = downcast_state + self.head_group_reduce = head_group_reduce + self.channel_gate_bwd = channel_gate_bwd scale = node.params.get("scale") self.scale = float(scale) if scale is not None else 1.0 / math.sqrt(node.inputs["q"].dim[-1]) self.use_qk_l2norm = bool(node.params.get("use_qk_l2norm", False)) + self.safe_gate = bool(node.params.get("safe_gate", False)) + self.use_beta_sigmoid = bool(node.params.get("use_beta_sigmoid", False)) + glb = node.params.get("gate_lower_bound") + self.gate_lower_bound = float(glb) if glb is not None else bwd_mod.DEFAULT_GATE_LOWER_BOUND + self.gate_bwd_blocks = GATE_BWD_BLOCKS self.has_state_checkpoints = "state_checkpoints" in node.inputs self.has_state0 = "initial_state" in node.inputs self.has_dstate0 = "d_initial_state" in node.outputs @@ -268,12 +336,11 @@ def __init__(self, node, bwd_mod, regen_mod): self.io_name = "float16" if node.inputs["q"].get_data_type().name == "HALF" else "bfloat16" self.n_heads_out, self.total = HO, total layout = WorkspaceLayout() - self.off_sched = layout.add(16) # one [ticket, done] ring each for the regen and bwd kernels + self.off_sched = layout.add(16) self.num_sm = multiprocessor_count(current_device_id()) self.bwd_dyn_sched = B * HO <= self.num_sm self.batch_invariant = bool(node.params.get("batch_invariant", False)) - # cuts never in batch-invariant mode: whole-sequence items keep each - # sequence's math independent of the batch composition + # cuts never in batch-invariant mode self.split = not self.batch_invariant self.n_tiles = B * HO if self.split: @@ -288,7 +355,6 @@ def __init__(self, node, bwd_mod, regen_mod): self.off_item_scratch = layout.add(self.work_item_rows * WORK_ITEM_FIELDS * 4) self.chunk_scratch_rows = chunk_scratch_rows(total, B, self.b_t) self.off_chunk_scratch = layout.add(self.chunk_scratch_rows * HO * 4) - # chunk-0 entering state, io dtype (downcast initial_state; absent = in-kernel zeros) self.off_state0_io = layout.add(B * HO * K * V * 2) if self.has_state0 else None if not self.has_state_checkpoints: self.state_checkpoints_rows = max(total // self.b_t + B, 1) @@ -305,8 +371,12 @@ def __init__(self, node, bwd_mod, regen_mod): self.off_dk_ho = layout.add(total * HO * K * 2) if self.fold_dv: self.off_dv_ho = layout.add(total * HO * V * 2) + if self.safe_gate: + self.off_gate_part_a = layout.add(self.gate_bwd_blocks * HO * K * 4) + self.off_gate_part_dt = layout.add(self.gate_bwd_blocks * HO * K * 4) self.bwd_tm_bytes = tensormap_workspace_bytes(bwd_mod, B) self.off_bwd_tensormaps = layout.add(self.bwd_tm_bytes, align=128) + self.needs_table = self.split self.ws_bytes = layout.size regions = [ ("sched_regen", self.off_sched, "int32", (2,)), @@ -330,56 +400,176 @@ def __init__(self, node, bwd_mod, regen_mod): regions.append(("dk_ho", self.off_dk_ho, self.io_name, (total, HO, K))) if self.fold_dv: regions.append(("dv_ho", self.off_dv_ho, self.io_name, (total, HO, V))) + if self.safe_gate: + regions.append(("gate_part_a", self.off_gate_part_a, "float32", (self.gate_bwd_blocks * HO * K,))) + regions.append(("gate_part_dt", self.off_gate_part_dt, "float32", (self.gate_bwd_blocks * HO * K,))) self.carve_names = [name for name, _off, _dt, _shape in regions] self.carve = carve_plan("KdaFrostEngine (KDA_BWD)", [(off, dt, shape) for _name, off, dt, shape in regions]) def workspace_bytes(self) -> int: return self.ws_bytes - def __call__(self, node_buffers, *, workspace=None, stream=None) -> Any: - nb = node_buffers[self.node] - q = nb.inputs["q"] - k = nb.inputs["k"] - v = nb.inputs["v"] - g = nb.inputs["g"] - beta = nb.inputs["beta"] - cu = nb.inputs["cu_seqlens"] - do = nb.inputs["dO"] - state_checkpoints = nb.inputs.get("state_checkpoints") - state0 = nb.inputs.get("initial_state") - dstate_in = nb.inputs.get("d_final_state") - dq = nb.outputs["dQ"] - dk = nb.outputs["dK"] - dv = nb.outputs["dV"] - dg = nb.outputs["dG"] - db = nb.outputs["dBeta"] - dstate0 = nb.outputs.get("d_initial_state") + def bind(self, names) -> None: + pos = {name: i for i, name in enumerate(names)} + self.iq = pos["q"] + self.ik = pos["k"] + self.iv = pos["v"] + self.ig = pos["g"] + self.ibeta = pos["beta"] + self.icu = pos["cu_seqlens"] + self.ido = pos["dO"] + self.ick = pos.get("state_checkpoints") + self.is0 = pos.get("initial_state") + self.idfs = pos.get("d_final_state") + self.idq = pos["dQ"] + self.idk = pos["dK"] + self.idv = pos["dV"] + self.idg = pos["dG"] + self.idb = pos["dBeta"] + self.ids0 = pos.get("d_initial_state") + self.ia_log = pos.get("a_log") + self.idt_bias = pos.get("dt_bias") + self.ida_log = pos.get("d_a_log") + self.iddt_bias = pos.get("d_dt_bias") + + def run(self, views, workspace, stream) -> None: + q = views[self.iq] + k = views[self.ik] + v = views[self.iv] + g = views[self.ig] + beta = views[self.ibeta] + cu = views[self.icu] + do = views[self.ido] + state_checkpoints = views[self.ick] if self.ick is not None else None + state0 = views[self.is0] if self.is0 is not None else None + dstate_in = views[self.idfs] if self.idfs is not None else None + dq = views[self.idq] + dk = views[self.idk] + dv = views[self.idv] + dg = views[self.idg] + db = views[self.idb] + dstate0 = views[self.ids0] if self.ids0 is not None else None + a_log = views[self.ia_log] if self.ia_log is not None else None + dt_bias = views[self.idt_bias] if self.idt_bias is not None else None + d_a_log = views[self.ida_log] if self.ida_log is not None else None + d_dt_bias = views[self.iddt_bias] if self.iddt_bias is not None else None stream = stream if stream is not None else 0 - HO, total = self.n_heads_out, self.total - K, V = q.shape[-1], v.shape[-1] - B = cu.shape[0] - 1 region = dict(zip(self.carve_names, workspace.carve(self.carve))) sched_regen = region["sched_regen"] sched_bwd = region["sched_bwd"] work_items = region["work_items"] work_count = region["work_count"] - self.build_split_table( - g, - cu, - work_items, - work_count, - ideal_chunks=self.ideal, - n_tiles=self.n_tiles, - num_sms=self.num_sm, - b_t=self.b_t, - chunk_scratch=region.get("chunk_scratch"), - item_scratch=region.get("item_scratch"), - log_gate=True, - sched_ctr=region["sched_all"], - split=self.split, - stream=stream, - ) + + if self.kcache is not None and (self.table is not None or not self.needs_table): + if self.needs_table: + self.run_table( + self.table, + g, + a_log, + dt_bias, + cu, + region.get("chunk_scratch"), + region.get("item_scratch"), + work_items, + work_count, + region["sched_all"], + stream, + ) + state0_io = None + if state0 is not None: + state0_io = region["state0_io"] + self.downcast_state(state0, state0_io, stream=stream) + if self.has_state_checkpoints: + checkpoint_series = state_checkpoints + else: + checkpoint_series = region["state_checkpoints"] + self.regen.run_recompute( + self.regen_cache, + k, + v, + g, + a_log if self.safe_gate else None, + dt_bias if self.safe_gate else None, + beta, + cu, + state0, + None, + checkpoint_series, + work_items, + work_count, + sched_regen, + region["sched_all"], + region.get("item_scratch"), + region["regen_tensormaps"], + self.b_t, + stream, + ) + dq_out = region["dq_ho"] if self.fold_dq else dq + dk_out = region["dk_ho"] if self.fold_dk else dk + dv_out = region["dv_ho"] if self.fold_dv else dv + self.bwd.run_bwd( + self.kcache, + q, + k, + v, + g, + beta, + do, + checkpoint_series, + dq_out, + dk_out, + dv_out, + dg, + db, + cu, + state0_io, + dstate0 if self.has_dstate0 else None, + dstate_in, + work_items, + work_count, + sched_bwd if self.bwd_dyn_sched else None, + region["sched_all"] if self.has_state_checkpoints else None, + region.get("item_scratch") if self.has_state_checkpoints else None, + region["bwd_tensormaps"], + self.scale, + stream, + a_log=a_log if self.safe_gate else None, + dt_bias=dt_bias if self.safe_gate else None, + ) + if self.safe_gate: + self.channel_gate_bwd( + dg, g, a_log, dt_bias, d_a_log, d_dt_bias, region["gate_part_a"], region["gate_part_dt"], self.gate_lower_bound, stream=stream + ) + if self.fold_dq or self.fold_dk or self.fold_dv: + for src_ho, dst in ((dq_out, dq), (dk_out, dk), (dv_out, dv)): + if src_ho is not dst: + self.head_group_reduce(src_ho, dst, stream=stream) + return + + if not self.needs_table: + self.table = None + else: + self.table = self.build_split_table( + g, + cu, + work_items, + work_count, + ideal_chunks=self.ideal, + n_tiles=self.n_tiles, + num_sms=self.num_sm, + b_t=self.b_t, + chunk_scratch=region.get("chunk_scratch"), + item_scratch=region.get("item_scratch"), + log_gate=True, + safe_gate=self.safe_gate, + a_log=a_log, + dt_bias=dt_bias, + gate_lower_bound=self.gate_lower_bound if self.safe_gate else None, + sched_ctr=region["sched_all"], + split=self.split, + stream=stream, + ) state0_io = None if state0 is not None: @@ -389,7 +579,7 @@ def __call__(self, node_buffers, *, workspace=None, stream=None) -> Any: checkpoint_series = state_checkpoints else: checkpoint_series = region["state_checkpoints"] - self.regen.chunk_kda_recompute_sm100( + self.regen_cache = self.regen.chunk_kda_recompute_sm100( k, v, g, @@ -400,9 +590,17 @@ def __call__(self, node_buffers, *, workspace=None, stream=None) -> Any: checkpoint_every_n_tokens=self.b_t, output_state_checkpoints=checkpoint_series, use_qk_l2norm_in_kernel=self.use_qk_l2norm, + safe_gate=self.safe_gate, + gate_lower_bound=self.gate_lower_bound, + a_log=a_log, + dt_bias=dt_bias, + use_beta_sigmoid=self.use_beta_sigmoid, work_items=work_items, work_count=work_count, sched_ctr=sched_regen, + sched_all=region["sched_all"], + work_item_scratch=region.get("item_scratch"), + order_in_prologue=True, tensormap_workspace=region["regen_tensormaps"], stream=stream, ) @@ -415,7 +613,7 @@ def __call__(self, node_buffers, *, workspace=None, stream=None) -> Any: if self.fold_dv: dv_out = region["dv_ho"] - self.bwd.chunk_kda_bwd_sm100( + self.kcache = self.bwd.chunk_kda_bwd_sm100( q, k, v, @@ -434,16 +632,26 @@ def __call__(self, node_buffers, *, workspace=None, stream=None) -> Any: d_initial_state=dstate0 if self.has_dstate0 else None, d_final_state=dstate_in, use_qk_l2norm_in_kernel=self.use_qk_l2norm, + safe_gate=self.safe_gate, + gate_lower_bound=self.gate_lower_bound, + a_log=a_log, + dt_bias=dt_bias, + use_beta_sigmoid=self.use_beta_sigmoid, work_items=work_items, work_count=work_count, sched_ctr=sched_bwd if self.bwd_dyn_sched else None, + sched_all=region["sched_all"] if self.has_state_checkpoints else None, + work_item_scratch=region.get("item_scratch") if self.has_state_checkpoints else None, + order_in_prologue=self.has_state_checkpoints, tensormap_workspace=region["bwd_tensormaps"], stream=stream, ) + if self.safe_gate: + self.channel_gate_bwd( + dg, g, a_log, dt_bias, d_a_log, d_dt_bias, region["gate_part_a"], region["gate_part_dt"], self.gate_lower_bound, stream=stream + ) if dq_out is not dq or dk_out is not dk or dv_out is not dv: - from .common.head_reduce import head_group_reduce - for src_ho, dst in ((dq_out, dq), (dk_out, dk), (dv_out, dv)): if src_ho is not dst: - head_group_reduce(src_ho, dst, stream=stream) + self.head_group_reduce(src_ho, dst, stream=stream) return None diff --git a/python/cudnn/linear_attention/frost/kernel/gdn2_bprop_f16.py b/python/cudnn/linear_attention/frost/kernel/gdn2_bprop_f16.py index 5ab67d122..7e1c0409b 100644 --- a/python/cudnn/linear_attention/frost/kernel/gdn2_bprop_f16.py +++ b/python/cudnn/linear_attention/frost/kernel/gdn2_bprop_f16.py @@ -60,12 +60,13 @@ import cutlass.cute as cute from cutlass.cute.runtime import from_dlpack -from ..common.split_k import decode_work_item +from ..common.split_k import ORDER_CAPACITY, ORDER_ELEMS, ORDER_THREADS, decode_work_item, order_body from ..common.host import get_dtype from cudnn.frost.buffers import current_device_id, data_ptr from cudnn.frost.device import multiprocessor_count from ..common.thd import TENSOR_MAP_QWORDS, emit_copy_desc, emit_checkpoint_seq_descs, emit_seq_descs from .gdn2_bprop_config import CFG + from cudnn.frost.tile_dsl.barrier import ( advance, MBarrier, @@ -77,18 +78,23 @@ from cudnn.frost.tile_dsl.swizzle import swizzle_lin_S, swizzle_xor_128b from cudnn.frost.tile_dsl.tma import tma_load_tile, tma_store_commit, tma_store_tile, tma_store_wait, tma_tensormap_acquire from cudnn.frost.tile_dsl.pointwise import ( - opaque_f32_zero, f16x2_to_f32, - fmul2, + fadd2, ffma2, + fmul2, + fp32_to_fp16, movmatrix_16b, mul_f16x2, - fp32_to_fp16, + opaque_f32_zero, sub_f16x2, ) LOG2_E: float = 1.4426950408889634 + +DEFAULT_GATE_LOWER_BOUND: float = -5.0 + + L2_NORM_EPS: float = 1.0e-12 @@ -842,14 +848,10 @@ def super_mma_warp( tinv_lo1, tinv_hi1 = f16x2_to_f32(tinv_p1, dtype=cfg.io_dtype) tinv_lo2, tinv_hi2 = f16x2_to_f32(tinv_p2, dtype=cfg.io_dtype) tinv_lo3, tinv_hi3 = f16x2_to_f32(tinv_p3, dtype=cfg.io_dtype) - tinv_acc[0] = tinv_lo0 + upd_acc[0] - tinv_acc[1] = tinv_hi0 + upd_acc[1] - tinv_acc[2] = tinv_lo1 + upd_acc[2] - tinv_acc[3] = tinv_hi1 + upd_acc[3] - tinv_acc[4] = tinv_lo2 + upd_acc[4] - tinv_acc[5] = tinv_hi2 + upd_acc[5] - tinv_acc[6] = tinv_lo3 + upd_acc[6] - tinv_acc[7] = tinv_hi3 + upd_acc[7] + tinv_acc[0], tinv_acc[1] = fadd2(tinv_lo0, tinv_hi0, upd_acc[0], upd_acc[1]) + tinv_acc[2], tinv_acc[3] = fadd2(tinv_lo1, tinv_hi1, upd_acc[2], upd_acc[3]) + tinv_acc[4], tinv_acc[5] = fadd2(tinv_lo2, tinv_hi2, upd_acc[4], upd_acc[5]) + tinv_acc[6], tinv_acc[7] = fadd2(tinv_lo3, tinv_hi3, upd_acc[6], upd_acc[7]) nvvm.stmatrix( sIntermediate_ptr + 1 * (cfg.b_t * cfg.b_t) + stsm_idx, @@ -1795,6 +1797,18 @@ def tmaldg_warp( tile_idx = next_tile +@cute.jit +def gate_scale(cfg, raw_gate: cutlass.Float32) -> cutlass.Float32: + """Map raw gate to the log2-domain decay increment.""" + + if cutlass.const_expr(cfg.safe_gate): + half = cutlass.Float32(0.5) + sigmoid = cute.math.tanh(raw_gate * half, approx=True) * half + half + return cfg.gate_scale_log2 * sigmoid + # Default ABI: Gate arrives in natural-log space + return raw_gate * cutlass.Float32(LOG2_E) + + @cute.jit def compute0_warp_group( cfg, @@ -1808,6 +1822,8 @@ def compute0_warp_group( tmem_hold, warp_idx, scale, + mA_log, + mDt_bias, sK_inv_raw, sGate_raw, sK_raw, @@ -1839,6 +1855,9 @@ def compute0_warp_group( value_dim = tmem_sp * cfg.threads_per_warp + lane state_copy_addr = (tmem_row + tmem_sp * cfg.threads_per_warp) << 16 state_index = PipelineState.start(phase=0) + cg0_prefix_dim = cg0_warp * cfg.threads_per_warp + lane + cg0_a_log_exp = cutlass.Float32(1.0) + cg0_dt_bias_value = cutlass.Float32(0.0) gbase = cutlass.Int32(0) sched_state = PipelineState.start(phase=0) tile_idx = cutlass.Int32(bidx) @@ -1847,6 +1866,10 @@ def compute0_warp_group( while tile_idx < total_tiles: batch_idx, head_idx, batch_start, batch_end, seqlen_b, num_chunks_b, wstart, wend, cstart, cend = decode_work_item(cfg, tile_idx, mWorkItems) sk_nt = cend - wstart + if cutlass.const_expr(cfg.safe_gate): + if sk_nt > 0: + cg0_a_log_exp = cute.math.exp2(mA_log[head_idx].to(cutlass.Float32) * LOG2_E, fastmath=True) + cg0_dt_bias_value = mDt_bias[head_idx, cg0_prefix_dim].to(cutlass.Float32) for rev_idx in cutlass.range(sk_nt, unroll=1): chunk_idx = cend - cutlass.Int32(1) - rev_idx gc = gbase + rev_idx @@ -1884,14 +1907,38 @@ def compute0_warp_group( prefix_idx = f32_segment * (cfg.b_t * 32) + row * 32 + swizzle_xor_128b(row, f32_segment_dim, elem_bytes=4) gate_raw[row] = (sGate_ptr + prefix_idx).load() g_prefix_regs = cutlass.Array(cutlass.Float32, cfg.b_t, alignment=16) - for row in cutlass.range_constexpr(cfg.b_t): - gate = gate_raw[row] - token_idx = chunk_idx * cutlass.Int32(cfg.b_t) + cutlass.Int32(row) - if token_idx < seqlen_b: - gate = gate * cutlass.Float32(LOG2_E) - else: - gate = cutlass.Float32(0.0) - g_prefix_regs[row] = gate + if cutlass.const_expr(cfg.safe_gate): + valid_rows = seqlen_b - chunk_idx * cutlass.Int32(cfg.b_t) + valid_mask = cutlass.vector.create_mask([cfg.b_t], [valid_rows]) + for row_pair in cutlass.range_constexpr(cfg.b_t // 2): + row0 = row_pair * 2 + row1 = row0 + 1 + gate0 = cg0_a_log_exp * (gate_raw[row0] + cg0_dt_bias_value) + gate1 = cg0_a_log_exp * (gate_raw[row1] + cg0_dt_bias_value) + gate0 = gate_scale( + cfg, + gate0, + ) + gate1 = gate_scale( + cfg, + gate1, + ) + gate_pair = cutlass.Vector.from_elements((gate0, gate1), cutlass.Float32) + gate_pair = cutlass.vector.where(valid_mask[row0 : row1 + 1], gate_pair, 0.0) + g_prefix_regs[row0] = gate_pair[0] + g_prefix_regs[row1] = gate_pair[1] + else: + for row in cutlass.range_constexpr(cfg.b_t): + gate = gate_raw[row] + token_idx = chunk_idx * cutlass.Int32(cfg.b_t) + cutlass.Int32(row) + if token_idx < seqlen_b: + gate = gate_scale( + cfg, + gate, + ) + else: + gate = cutlass.Float32(0.0) + g_prefix_regs[row] = gate prefix_acc = cutlass.Float32(0.0) for row_pair in cutlass.range_constexpr(cfg.b_t // 2): @@ -1991,7 +2038,11 @@ def compute0_warp_group( k_val = raw_k_frag_f32[dim_offset] raw_q_regs[reg_base + dim_offset] = q_val raw_k_regs[reg_base + dim_offset] = k_val - raw_beta_regs[reg_base + dim_offset] = raw_beta_frag_f32[dim_offset] + beta_val = raw_beta_frag_f32[dim_offset] + if cutlass.const_expr(cfg.beta_sigmoid): + half = cutlass.Float32(0.5) + beta_val = (cute.math.tanh(beta_val * half, approx=True) * half + half).to(cfg.io_dtype).to(cutlass.Float32) + raw_beta_regs[reg_base + dim_offset] = beta_val if cutlass.const_expr(cfg.l2norm): if cutlass.const_expr(dim_offset % 2 == 0): qk0_lo, qk0_hi = ffma2(q_val, k_val, q_val, k_val, qk0_lo, qk0_hi) @@ -2640,8 +2691,9 @@ def compute2_warp_group( hval1 = (sDstate_raw.data_ptr() + dstate_addr1).load().to(cutlass.Float32) hacc[(2 * j) % 8] = hacc[(2 * j) % 8] + hval0 * state_pair[0].to(cutlass.Float32) hacc[(2 * j + 1) % 8] = hacc[(2 * j + 1) % 8] + hval1 * state_pair[1].to(cutlass.Float32) - part_a = (hacc[0] + hacc[4]) + (hacc[1] + hacc[5]) - part_b = (hacc[2] + hacc[6]) + (hacc[3] + hacc[7]) + pa0, pb0 = fadd2(hacc[0], hacc[2], hacc[4], hacc[6]) + pa1, pb1 = fadd2(hacc[1], hacc[3], hacc[5], hacc[7]) + part_a, part_b = fadd2(pa0, pb0, pa1, pb1) dgate_last_val = dgate_last_val + (part_a + part_b) bars.mb_dstate_smem_cg2_done.arrive() bars.mb_state_inp_cg2_done[gc % 2].arrive() @@ -2704,9 +2756,11 @@ def compute2_warp_group( if cutlass.const_expr(cfg.l2norm): k_v = k_v * sNorm_raw[norm_base + cfg.b_t + t] db_regs[t] = k_v * dk_decay - dgate_regs[t] = (sBetaP_ptr + f16_seg * (cfg.b_t * 64) + t * 64 + swizzle_xor_128b(t, f16_dim, elem_bytes=2)).load().to( - cutlass.Float32 - ) * dk_decay + beta_v = (sBetaP_ptr + f16_seg * (cfg.b_t * 64) + t * 64 + swizzle_xor_128b(t, f16_dim, elem_bytes=2)).load().to(cutlass.Float32) + if cutlass.const_expr(cfg.beta_sigmoid): + half = cutlass.Float32(0.5) + beta_v = (cute.math.tanh(beta_v * half, approx=True) * half + half).to(cfg.io_dtype).to(cutlass.Float32) + dgate_regs[t] = beta_v * dk_decay dk_n[t] = dk_n[t] + dgate_regs[t] nvvm.tcgen05_wait("load") @@ -2723,11 +2777,14 @@ def compute2_warp_group( if cutlass.const_expr(cfg.l2norm): q_v = q_v * sNorm_raw[norm_base + t] k_v = k_v * sNorm_raw[norm_base + cfg.b_t + t] - dgate_regs[t] = ( - q_v * dq_n[t] - + (sBetaP_ptr + f16_seg * (cfg.b_t * 64) + t * 64 + swizzle_xor_128b(t, f16_dim, elem_bytes=2)).load().to(cutlass.Float32) * db_regs[t] - - k_v * (dk_n[t] - dgate_regs[t]) - ) + beta_v = (sBetaP_ptr + f16_seg * (cfg.b_t * 64) + t * 64 + swizzle_xor_128b(t, f16_dim, elem_bytes=2)).load().to(cutlass.Float32) + if cutlass.const_expr(cfg.beta_sigmoid): + half = cutlass.Float32(0.5) + beta_v = (cute.math.tanh(beta_v * half, approx=True) * half + half).to(cfg.io_dtype).to(cutlass.Float32) + dgate_regs[t] = q_v * dq_n[t] + beta_v * db_regs[t] - k_v * (dk_n[t] - dgate_regs[t]) + if cutlass.const_expr(cfg.beta_sigmoid): + # after dgate, which consumes db_regs pre-chain-rule + db_regs[t] = db_regs[t] * (beta_v - beta_v * beta_v) dgate_regs[cfg.b_t - 1] = dgate_regs[cfg.b_t - 1] + ((dgate_last_acc[0] + dgate_last_acc[1]) + (dgate_last_acc[2] + dgate_last_acc[3])) nvvm.fence_proxy("async.shared", space="cta") bars.mb_beta_done[raw_stage].arrive() @@ -2738,26 +2795,35 @@ def compute2_warp_group( dots = cutlass.Array(cutlass.Float32, cfg.b_t, alignment=16) for half in cutlass.range_constexpr(2): p_words = nvvm.tcgen05_ld("32x32b", nvvm.make_tmem_ptr(row_addr + qk_col + half * (cfg.b_t // 4), cutlass.Float32), num=cfg.b_t // 4) - for tt in cutlass.range_constexpr(cfg.b_t // 2): + for tt2 in cutlass.range_constexpr(cfg.b_t // 4): + tt = 2 * tt2 t = half * (cfg.b_t // 2) + tt - p_pair = cutlass.Vector.from_elements((p_words[tt // 2],), cutlass.Float32).bitcast(cfg.io_dtype) - dots[t] = grad[t] * p_pair[tt % 2].to(cutlass.Float32) * sNorm_raw[norm_base + inv_off + t] + p_pair = cutlass.Vector.from_elements((p_words[tt2],), cutlass.Float32).bitcast(cfg.io_dtype) + gp_lo, gp_hi = fmul2(grad[t], grad[t + 1], p_pair[0].to(cutlass.Float32), p_pair[1].to(cutlass.Float32)) + dots[t], dots[t + 1] = fmul2(gp_lo, gp_hi, sNorm_raw[norm_base + inv_off + t], sNorm_raw[norm_base + inv_off + t + 1]) for off in cutlass.range_constexpr(5): step = cutlass.const_expr(1 << off) - for t in cutlass.range_constexpr(cfg.b_t): - dots[t] = dots[t] + cutlass.Float32(nvvm.shfl_sync(0xFFFFFFFF, dots[t], step, 31, kind=nvvm.Shfl.BFLY)) + for t2 in cutlass.range_constexpr(cfg.b_t // 2): + t = 2 * t2 + bfly_lo = cutlass.Float32(nvvm.shfl_sync(0xFFFFFFFF, dots[t], step, 31, kind=nvvm.Shfl.BFLY)) + bfly_hi = cutlass.Float32(nvvm.shfl_sync(0xFFFFFFFF, dots[t + 1], step, 31, kind=nvvm.Shfl.BFLY)) + dots[t], dots[t + 1] = fadd2(dots[t], dots[t + 1], bfly_lo, bfly_hi) if lane == 0: for t in cutlass.range_constexpr(cfg.b_t): sRed1_raw[wg1_sp * cfg.b_t + t] = dots[t] nvvm.barrier_cta_sync(cfg.cg2_sync_barrier_id, thread_count=cfg.cg2_threads) for half in cutlass.range_constexpr(2): a_words = nvvm.tcgen05_ld("32x32b", nvvm.make_tmem_ptr(row_addr + qk_col + half * (cfg.b_t // 4), cutlass.Float32), num=cfg.b_t // 4) - for tt in cutlass.range_constexpr(cfg.b_t // 2): - t = half * (cfg.b_t // 2) + tt - a_pair = cutlass.Vector.from_elements((a_words[tt // 2],), cutlass.Float32).bitcast(cfg.io_dtype) - total_dot = sRed1_raw[t] + sRed1_raw[cfg.b_t + t] + sRed1_raw[2 * cfg.b_t + t] + sRed1_raw[3 * cfg.b_t + t] - norm_t = sNorm_raw[norm_base + inv_off + t] - grad[t] = (grad[t] - a_pair[tt % 2].to(cutlass.Float32) * norm_t * total_dot) * norm_t + for tt2 in cutlass.range_constexpr(cfg.b_t // 4): + t = half * (cfg.b_t // 2) + 2 * tt2 + a_pair = cutlass.Vector.from_elements((a_words[tt2],), cutlass.Float32).bitcast(cfg.io_dtype) + dot_lo, dot_hi = fadd2(sRed1_raw[t], sRed1_raw[t + 1], sRed1_raw[cfg.b_t + t], sRed1_raw[cfg.b_t + t + 1]) + dot_lo, dot_hi = fadd2(dot_lo, dot_hi, sRed1_raw[2 * cfg.b_t + t], sRed1_raw[2 * cfg.b_t + t + 1]) + dot_lo, dot_hi = fadd2(dot_lo, dot_hi, sRed1_raw[3 * cfg.b_t + t], sRed1_raw[3 * cfg.b_t + t + 1]) + norm_lo = sNorm_raw[norm_base + inv_off + t] + norm_hi = sNorm_raw[norm_base + inv_off + t + 1] + grad[t] = (grad[t] - a_pair[0].to(cutlass.Float32) * norm_lo * dot_lo) * norm_lo + grad[t + 1] = (grad[t + 1] - a_pair[1].to(cutlass.Float32) * norm_hi * dot_hi) * norm_hi nvvm.barrier_cta_sync(cfg.cg2_sync_barrier_id, thread_count=cfg.cg2_threads) nvvm.tcgen05_wait("load") @@ -2815,23 +2881,24 @@ def compute2_warp_group( # --------------------------------------------------------------------------- -@cute.kernel -def build_all_descs_kernel( - base_q: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_k: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_v: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_gate: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_do: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_beta: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_w: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_dq: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_dk: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_dv: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_dgate: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_dwo: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_dbo: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_checkpoint: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_initial_state: cutlass.GridConstant[cuda.tensor_map.TensorMap], +@cute.jit +def build_descs_body( + widx, + base_q, + base_k, + base_v, + base_gate, + base_do, + base_beta, + base_w, + base_dq, + base_dk, + base_dv, + base_dgate, + base_dwo, + base_dbo, + base_checkpoint, + base_initial_state, desc_ws: cute.Tensor, cu_seqlens: cute.Tensor, q: cute.Tensor, @@ -2866,10 +2933,9 @@ def build_all_descs_kernel( checkpoint_rs: cutlass.Int32, checkpoint_every_n: cutlass.Int32, ) -> None: - """Single-launch builder for the per-batch TMA-descriptor arrays (one - warp per array).""" - tidx, _, _ = cute.arch.thread_idx() - widx = cutlass.Int32(tidx) // cutlass.Int32(32) + """Per-batch descriptor-array build, one warp per array. Runs inside the + prologue kernel after its order pass; warps past the array count fall + through the widx guards.""" arr_words = n_batch * cutlass.Int32(TENSOR_MAP_QWORDS) sub0 = cute.make_tensor(desc_ws.iterator, cute.make_layout((arr_words,), stride=(1,))) sub1 = cute.make_tensor(desc_ws.iterator + arr_words, cute.make_layout((arr_words,), stride=(1,))) @@ -2950,10 +3016,156 @@ def build_all_descs_kernel( nvvm.fence_proxy_release(nvvm.MemScope.GPU, from_proxy=nvvm.Proxy.GENERIC, to_proxy=nvvm.Proxy.TENSORMAP) +@cute.kernel +def prologue_kernel( + run_order: cutlass.Constexpr[bool], + order_gen: cutlass.Constexpr[bool], + has_sched: cutlass.Constexpr[bool], + b_t: cutlass.Constexpr[int], + base_q: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_k: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_v: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_gate: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_do: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_beta: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_w: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_dq: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_dk: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_dv: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_dgate: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_dwo: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_dbo: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_checkpoint: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_initial_state: cutlass.GridConstant[cuda.tensor_map.TensorMap], + desc_ws: cute.Tensor, + cu_seqlens: cute.Tensor, + q: cute.Tensor, + k: cute.Tensor, + v: cute.Tensor, + gate: cute.Tensor, + do: cute.Tensor, + beta: cute.Tensor, + w: cute.Tensor, + dq: cute.Tensor, + dk: cute.Tensor, + dv: cute.Tensor, + dgate: cute.Tensor, + dwo: cute.Tensor, + dbo: cute.Tensor, + state_checkpoints: cute.Tensor, + initial_state: cute.Tensor | None, + mStaging: cute.Tensor | None, + mCount: cute.Tensor, + mWorkItems: cute.Tensor, + mSched: cute.Tensor | None, + n_batch: cutlass.Int32, + q_rs: cutlass.Int32, + k_rs: cutlass.Int32, + v_rs: cutlass.Int32, + g_rs: cutlass.Int32, + do_rs: cutlass.Int32, + beta_rs: cutlass.Int32, + w_rs: cutlass.Int32, + dq_rs: cutlass.Int32, + dk_rs: cutlass.Int32, + dv_rs: cutlass.Int32, + dgate_rs: cutlass.Int32, + dwo_rs: cutlass.Int32, + dbo_rs: cutlass.Int32, + checkpoint_rs: cutlass.Int32, + checkpoint_every_n: cutlass.Int32, +) -> None: + """Single-CTA prologue. Under ``run_order`` this kernel is the first + work-item-table consumer, so it LPT-orders the table and zeroes both + consumers' sched rings via :func:`order_body`; it then builds the + per-batch TMA-descriptor arrays via :func:`build_descs_body`, one warp + per array (the extra warps only take part in the order phase).""" + tidx, _, _ = cute.arch.thread_idx() + tidx = cutlass.Int32(tidx) + widx = tidx // cutlass.Int32(32) + if cutlass.const_expr(run_order): + sKey = cutlass.Array(cutlass.Int32, ORDER_CAPACITY, space=cutlass.AddressSpace.smem, alignment=16) + sIdx = cutlass.Array(cutlass.Int32, ORDER_CAPACITY, space=cutlass.AddressSpace.smem, alignment=16) + sSpread = cutlass.Array(cutlass.Int32, 2, space=cutlass.AddressSpace.smem, alignment=8) + n_heads_out = cutlass.Int32(gate.shape[1]) + order_body( + order_gen, + has_sched, + b_t, + ORDER_THREADS, + ORDER_ELEMS, + tidx, + n_heads_out, + n_heads_out * n_batch, + cu_seqlens, + mStaging, + mCount, + mWorkItems, + mSched, + sKey, + sIdx, + sSpread, + ) + build_descs_body( + widx, + base_q, + base_k, + base_v, + base_gate, + base_do, + base_beta, + base_w, + base_dq, + base_dk, + base_dv, + base_dgate, + base_dwo, + base_dbo, + base_checkpoint, + base_initial_state, + desc_ws, + cu_seqlens, + q, + k, + v, + gate, + do, + beta, + w, + dq, + dk, + dv, + dgate, + dwo, + dbo, + state_checkpoints, + initial_state, + n_batch, + q_rs, + k_rs, + v_rs, + g_rs, + do_rs, + beta_rs, + w_rs, + dq_rs, + dk_rs, + dv_rs, + dgate_rs, + dwo_rs, + dbo_rs, + checkpoint_rs, + checkpoint_every_n, + ) + + @cute.jit -def build_descs( +def prologue( io_dtype: cutlass.Constexpr, b_t: cutlass.Constexpr[int], + run_order: cutlass.Constexpr[bool], + order_gen: cutlass.Constexpr[bool], + has_sched: cutlass.Constexpr[bool], q: cute.Tensor, k: cute.Tensor, v: cute.Tensor, @@ -2970,11 +3182,16 @@ def build_descs( state_checkpoints: cute.Tensor, initial_state: cute.Tensor | None, cu_seqlens: cute.Tensor, + work_item_staging: cute.Tensor | None, + work_count: cute.Tensor, + work_items: cute.Tensor, + sched_all: cute.Tensor | None, tensormap_workspace: cute.Tensor, stream: cuda_driver.CUstream, ): - """Build the 15 per-(batch, head) TMA-descriptor arrays into - ``tensormap_workspace``.""" + """One-launch prologue: LPT-order the work items (when this kernel is + the table's first consumer) and build the 15 per-(batch, head) + TMA-descriptor arrays into ``tensormap_workspace``.""" h_q = q.shape[1] h_k = k.shape[1] h_v = v.shape[1] @@ -3034,8 +3251,11 @@ def build_descs( ) base_initial_state = cuda.create_tensor_map_tiled_from_view(initial_state_view, box_dims=(64, d_k, 1, 1), stride_order=(0, 1, 2, 3), swizzle=swz) - n_warps = 15 if initial_state is not None else 14 - build_all_descs_kernel( + prologue_kernel( + run_order, + order_gen, + has_sched, + b_t, base_q, base_k, base_v, @@ -3068,6 +3288,10 @@ def build_descs( dbo, state_checkpoints, initial_state, + work_item_staging, + work_count, + work_items, + sched_all, cutlass.Int32(batch_size), cutlass.Int32(q.stride[0]), cutlass.Int32(k.stride[0]), @@ -3084,7 +3308,7 @@ def build_descs( cutlass.Int32(dbo.stride[0]), cutlass.Int32(state_checkpoints.stride[0]), cutlass.Int32(b_t), - ).launch(grid=(1, 1, 1), block=(32 * n_warps, 1, 1), stream=stream) + ).launch(grid=(1, 1, 1), block=(ORDER_THREADS, 1, 1), stream=stream) @cute.jit @@ -3092,6 +3316,8 @@ def host( cfg: cutlass.Constexpr, state_checkpoints: cute.Tensor, mState_init: cute.Tensor | None, + a_log: cute.Tensor | None, + dt_bias: cute.Tensor | None, dgate: cute.Tensor, dbeta: cute.Tensor, dw: cute.Tensor, @@ -3115,6 +3341,8 @@ def host( tensormap_workspace, n_desc, cu_seqlens, + a_log, + dt_bias, dgate, dbeta, dw, @@ -3138,6 +3366,8 @@ def kernel( tensormap_workspace: cute.Tensor, n_desc: cutlass.Int32, cu_seqlens: cute.Tensor, + mA_log: cute.Tensor | None, + mDt_bias: cute.Tensor | None, mDgate: cute.Tensor, mDb: cute.Tensor, mDw_out: cute.Tensor, @@ -3547,6 +3777,8 @@ def kernel( tmem_hold, warp_idx, scale, + mA_log, + mDt_bias, sK_inv_raw, sGate_raw, sK_raw, @@ -3622,6 +3854,9 @@ class Gdn2BwdCfg: use_dstate_in: bool use_dstate0: bool l2norm: bool + safe_gate: bool + gate_scale_log2: float + beta_sigmoid: bool use_initial_state: bool q_ratio: int k_ratio: int @@ -3713,6 +3948,9 @@ def build_cfg( use_dstate_in: bool, use_dstate0: bool, l2norm: bool, + safe_gate: bool, + gate_scale_log2: float, + beta_sigmoid: bool, use_initial_state: bool, q_ratio: int, k_ratio: int, @@ -3728,6 +3966,9 @@ def build_cfg( use_dstate_in=use_dstate_in, use_dstate0=use_dstate0, l2norm=l2norm, + safe_gate=safe_gate, + gate_scale_log2=gate_scale_log2, + beta_sigmoid=beta_sigmoid, use_initial_state=use_initial_state, q_ratio=q_ratio, k_ratio=k_ratio, @@ -3803,8 +4044,14 @@ def get_compiled_cache( use_dstate_in: bool, use_dstate0: bool, l2norm: bool, + safe_gate: bool, + gate_lower_bound: float, + beta_sigmoid: bool, use_initial_state: bool, dyn_sched: bool, + order_in_prologue: bool, + order_gen: bool, + has_sched: bool, ): return {} @@ -3831,9 +4078,17 @@ def chunk_gdn2_bwd_sm100( d_initial_state=None, d_final_state=None, use_qk_l2norm_in_kernel: bool = False, + safe_gate: bool = False, + gate_lower_bound: float = DEFAULT_GATE_LOWER_BOUND, + a_log=None, + dt_bias=None, + use_beta_sigmoid: bool = False, work_items=None, work_count=None, sched_ctr=None, + sched_all=None, + work_item_scratch=None, + order_in_prologue: bool = False, tensormap_workspace, stream, ) -> None: @@ -3845,8 +4100,11 @@ def chunk_gdn2_bwd_sm100( q: ``(total_tokens, HQ, DK)`` float16/bfloat16 k: ``(total_tokens, HK, DK)`` float16/bfloat16 v: ``(total_tokens, HV, DV)`` float16/bfloat16 - gate: ``(total_tokens, HO, DK)`` fp32 natural-log per-channel decay - beta: ``(total_tokens, HO, DK)`` io dtype post-sigmoid per-key erase + gate: ``(total_tokens, HO, DK)`` fp32 per-channel decay. Natural-log + unless ``safe_gate``, which applies the safe-gate transform + ``lower_bound * sigmoid(exp(a_log) * (gate + dt_bias))`` + beta: ``(total_tokens, HO, DK)`` io dtype per-key erase. Post-sigmoid, + or logits when ``use_beta_sigmoid`` w: ``(total_tokens, HO, DV)`` io dtype post-sigmoid per-value write do: ``(total_tokens, HO, DV)`` io dtype state_checkpoints: ``(total_checkpoints, HO, DK, DV)`` io dtype (KV, v contiguous - the GDN @@ -3854,8 +4112,10 @@ def chunk_gdn2_bwd_sm100( slot: sequence-local entry ``c - 1`` is the state ENTERING chunk c >= 1 of sequence b; chunk 0 seeds from ``initial_state`` dq/dk/dv: io dtype at ``HO = max(HQ, HV)`` heads, pre-allocated - dgate: ``(total_tokens, HO, DK)`` fp32 (dL/d ln alpha), pre-allocated - dbeta: ``(total_tokens, HO, DK)`` io dtype, pre-allocated + dgate: ``(total_tokens, HO, DK)`` fp32 (dL/d ln alpha; ``safe_gate`` + leaves it in the transformed gate space), pre-allocated + dbeta: ``(total_tokens, HO, DK)`` io dtype (post-sigmoid space, or + wrt the raw logits under ``beta_sigmoid``), pre-allocated dw: ``(total_tokens, HO, DV)`` io dtype, pre-allocated cu_seqlens: ``(num_seqs + 1,)`` int32 scale: attention scale factor @@ -3865,6 +4125,10 @@ def chunk_gdn2_bwd_sm100( d_final_state: fp32 ``(num_seqs, HO, DK, DV)`` IN (dL/d final state) use_qk_l2norm_in_kernel: q/k arrive raw; the kernel normalizes for the recompute math and chains the L2-norm backward into dq/dk + safe_gate: interpret ``gate`` through the safe-gate transform + a_log: ``(HO,)`` float32, safe-gate per-head log-amplitude (None = 0) + dt_bias: ``(HO, DK)`` float32, safe-gate channel bias (None = 0) + use_beta_sigmoid: ``beta`` holds logits; sigmoid in-kernel work_items/work_count: split-K table (``common/split_k.py``, REQUIRED; an uncut table row is the whole (b, h) sequence); each item computes chunks ``[wstart, cend)`` backward and writes @@ -3885,6 +4149,9 @@ def chunk_gdn2_bwd_sm100( if work_items is None or work_count is None: raise ValueError("work_items/work_count are required (the split-table stage builds them for every launch)") dyn_sched = sched_ctr is not None + order_gen = work_item_scratch is None + if order_in_prologue and sched_all is None: + raise ValueError("order_in_prologue requires sched_all (the prologue zeroes both consumers' sched rings)") for name, t in (("state_checkpoints", state_checkpoints), ("beta", beta), ("w", w), ("dbeta", dbeta), ("dw", dw)) + ( (("initial_state", initial_state),) if use_initial_state else () ): @@ -3894,6 +4161,13 @@ def chunk_gdn2_bwd_sm100( if HO % hh != 0: raise ValueError(f"{name}={hh} must divide {HO}") B = cu_seqlens.shape[0] - 1 + gate_scale_log2 = gate_lower_bound * LOG2_E + + if safe_gate and (a_log is None or dt_bias is None): + raise ValueError("safe_gate requires a_log and dt_bias") + if not safe_gate: + a_log = None + dt_bias = None cu_stream = cuda_driver.CUstream(int(stream)) cache = get_compiled_cache( @@ -3905,8 +4179,14 @@ def chunk_gdn2_bwd_sm100( use_dstate_in, use_dstate0, use_qk_l2norm_in_kernel, + safe_gate, + gate_lower_bound, + use_beta_sigmoid, use_initial_state, dyn_sched, + order_in_prologue, + order_gen, + sched_all is not None, ) if "compiled" not in cache: @@ -3916,6 +4196,9 @@ def chunk_gdn2_bwd_sm100( use_dstate_in=use_dstate_in, use_dstate0=use_dstate0, l2norm=use_qk_l2norm_in_kernel, + safe_gate=safe_gate, + gate_scale_log2=gate_scale_log2, + beta_sigmoid=use_beta_sigmoid, use_initial_state=use_initial_state, q_ratio=HO // HQ, k_ratio=HO // HK, @@ -3944,6 +4227,8 @@ def chunk_gdn2_bwd_sm100( initial_state_cute = ( from_dlpack(initial_state, assumed_align=16).mark_layout_dynamic(leading_dim=len(initial_state.shape) - 1) if use_initial_state else None ) + a_log_cute = from_dlpack(a_log, assumed_align=4) if a_log is not None else None + dt_bias_cute = from_dlpack(dt_bias, assumed_align=16) if dt_bias is not None else None dgate_cute = from_dlpack(dgate, assumed_align=16).mark_layout_dynamic(leading_dim=len(dgate.shape) - 1) dbeta_cute = from_dlpack(dbeta, assumed_align=16).mark_layout_dynamic(leading_dim=len(dbeta.shape) - 1) dw_cute = from_dlpack(dw, assumed_align=16).mark_layout_dynamic(leading_dim=len(dw.shape) - 1) @@ -3952,6 +4237,8 @@ def chunk_gdn2_bwd_sm100( cfg, state_checkpoints_cute, initial_state_cute, + a_log_cute, + dt_bias_cute, dgate_cute, dbeta_cute, dw_cute, @@ -3967,58 +4254,175 @@ def chunk_gdn2_bwd_sm100( options="--enable-tvm-ffi --opt-level 2", ) - # ---- per-(batch, head) descriptor arrays: rebuild on input change ------------ - # desc build runs every execute by contract (cu contents are data; - # buffer pointers may change) - capture-safe, single tiny launch - if "build_descs" not in cache: + if "prologue" not in cache: io_dtype = get_dtype(q.dtype) - - q_bd = from_dlpack(q, assumed_align=16).mark_layout_dynamic(leading_dim=2) - k_bd = from_dlpack(k, assumed_align=16).mark_layout_dynamic(leading_dim=2) - v_bd = from_dlpack(v, assumed_align=16).mark_layout_dynamic(leading_dim=2) - gate_bd = from_dlpack(gate, assumed_align=16).mark_layout_dynamic(leading_dim=2) - do_bd = from_dlpack(do, assumed_align=16).mark_layout_dynamic(leading_dim=2) - beta_bd = from_dlpack(beta, assumed_align=16).mark_layout_dynamic(leading_dim=2) - w_bd = from_dlpack(w, assumed_align=16).mark_layout_dynamic(leading_dim=2) - dq_bd = from_dlpack(dq, assumed_align=16).mark_layout_dynamic(leading_dim=2) - dk_bd = from_dlpack(dk, assumed_align=16).mark_layout_dynamic(leading_dim=2) - dv_bd = from_dlpack(dv, assumed_align=16).mark_layout_dynamic(leading_dim=2) - dgate_bd = from_dlpack(dgate, assumed_align=16).mark_layout_dynamic(leading_dim=2) - dwo_bd = from_dlpack(dw, assumed_align=16).mark_layout_dynamic(leading_dim=2) - dbo_bd = from_dlpack(dbeta, assumed_align=16).mark_layout_dynamic(leading_dim=2) - state_checkpoints_bd = from_dlpack(state_checkpoints, assumed_align=16).mark_layout_dynamic(leading_dim=3) - initial_state_bd = from_dlpack(initial_state, assumed_align=16).mark_layout_dynamic(leading_dim=3) if use_initial_state else None - - cu_bd = from_dlpack(cu_seqlens, assumed_align=8).mark_layout_dynamic() - ws_bd = from_dlpack(tensormap_workspace, assumed_align=128).mark_layout_dynamic() - cache["build_descs"] = cute.compile( - build_descs, + q_pl = from_dlpack(q, assumed_align=16).mark_layout_dynamic(leading_dim=2) + k_pl = from_dlpack(k, assumed_align=16).mark_layout_dynamic(leading_dim=2) + v_pl = from_dlpack(v, assumed_align=16).mark_layout_dynamic(leading_dim=2) + gate_pl = from_dlpack(gate, assumed_align=16).mark_layout_dynamic(leading_dim=2) + do_pl = from_dlpack(do, assumed_align=16).mark_layout_dynamic(leading_dim=2) + beta_pl = from_dlpack(beta, assumed_align=16).mark_layout_dynamic(leading_dim=2) + w_pl = from_dlpack(w, assumed_align=16).mark_layout_dynamic(leading_dim=2) + dq_pl = from_dlpack(dq, assumed_align=16).mark_layout_dynamic(leading_dim=2) + dk_pl = from_dlpack(dk, assumed_align=16).mark_layout_dynamic(leading_dim=2) + dv_pl = from_dlpack(dv, assumed_align=16).mark_layout_dynamic(leading_dim=2) + dgate_pl = from_dlpack(dgate, assumed_align=16).mark_layout_dynamic(leading_dim=2) + dwo_pl = from_dlpack(dw, assumed_align=16).mark_layout_dynamic(leading_dim=2) + dbo_pl = from_dlpack(dbeta, assumed_align=16).mark_layout_dynamic(leading_dim=2) + state_checkpoints_pl = from_dlpack(state_checkpoints, assumed_align=16).mark_layout_dynamic(leading_dim=3) + initial_state_pl = from_dlpack(initial_state, assumed_align=16).mark_layout_dynamic(leading_dim=3) if use_initial_state else None + cu_pl = from_dlpack(cu_seqlens, assumed_align=8).mark_layout_dynamic() + ws_pl = from_dlpack(tensormap_workspace, assumed_align=128).mark_layout_dynamic() + staging_pl = None + if not order_gen: + staging_pl = from_dlpack(work_item_scratch, assumed_align=16) + staging_pl.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1), divisibility=1) + work_items_pl = from_dlpack(work_items, assumed_align=16) + work_items_pl.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1), divisibility=1) + work_count_pl = from_dlpack(work_count, assumed_align=4).mark_layout_dynamic() + sched_all_pl = None + if sched_all is not None: + sched_all_pl = from_dlpack(sched_all, assumed_align=4).mark_layout_dynamic() + cache["prologue"] = cute.compile( + prologue, io_dtype, CFG.B_T, - q_bd, - k_bd, - v_bd, - gate_bd, - do_bd, - beta_bd, - w_bd, - dq_bd, - dk_bd, - dv_bd, - dgate_bd, - dwo_bd, - dbo_bd, - state_checkpoints_bd, - initial_state_bd, - cu_bd, - ws_bd, + order_in_prologue, + order_gen, + sched_all is not None, + q_pl, + k_pl, + v_pl, + gate_pl, + do_pl, + beta_pl, + w_pl, + dq_pl, + dk_pl, + dv_pl, + dgate_pl, + dwo_pl, + dbo_pl, + state_checkpoints_pl, + initial_state_pl, + cu_pl, + staging_pl, + work_count_pl, + work_items_pl, + sched_all_pl, + ws_pl, cu_stream, options="--enable-tvm-ffi", ) - cache["build_descs"](q, k, v, gate, do, beta, w, dq, dk, dv, dgate, dw, dbeta, state_checkpoints, initial_state, cu_seqlens, tensormap_workspace, cu_stream) + cache["prologue"]( + q, + k, + v, + gate, + do, + beta, + w, + dq, + dk, + dv, + dgate, + dw, + dbeta, + state_checkpoints, + initial_state, + cu_seqlens, + work_item_scratch if not order_gen else None, + work_count, + work_items, + sched_all, + tensormap_workspace, + cu_stream, + ) + cache["compiled"]( + state_checkpoints, + initial_state, + a_log, + dt_bias, + dgate, + dbeta, + dw, + cu_seqlens, + d_initial_state, + d_final_state, + work_items, + work_count, + sched_ctr, + tensormap_workspace, + scale, + cu_stream, + ) + return cache + + +def run_bwd( + cache, + q, + k, + v, + gate, + beta, + w, + do, + state_checkpoints, + dq, + dk, + dv, + dgate, + dbeta, + dw, + cu_seqlens, + initial_state, + d_initial_state, + d_final_state, + work_items, + work_count, + sched_ctr, + sched_all, + work_item_scratch, + tensormap_workspace, + scale, + stream, + a_log=None, + dt_bias=None, +) -> None: + """Replay the compiled plan: the prologue launch, then the main launch. + The caller owns the contract, which the plan validated at build, so + nothing here raises.""" + cu_stream = cuda_driver.CUstream(int(stream)) + cache["prologue"]( + q, + k, + v, + gate, + do, + beta, + w, + dq, + dk, + dv, + dgate, + dw, + dbeta, + state_checkpoints, + initial_state, + cu_seqlens, + work_item_scratch, + work_count, + work_items, + sched_all, + tensormap_workspace, + cu_stream, + ) cache["compiled"]( state_checkpoints, initial_state, + a_log, + dt_bias, dgate, dbeta, dw, diff --git a/python/cudnn/linear_attention/frost/kernel/gdn2_prefill_f16.py b/python/cudnn/linear_attention/frost/kernel/gdn2_prefill_f16.py index 2e97e2fc8..05e7215c2 100644 --- a/python/cudnn/linear_attention/frost/kernel/gdn2_prefill_f16.py +++ b/python/cudnn/linear_attention/frost/kernel/gdn2_prefill_f16.py @@ -93,12 +93,13 @@ import cutlass.cute as cute from cutlass.cute.runtime import from_dlpack -from ..common.split_k import decode_work_item +from ..common.split_k import ORDER_CAPACITY, ORDER_ELEMS, ORDER_THREADS, decode_work_item, order_body from ..common.host import get_dtype from cudnn.frost.buffers import current_device_id, data_ptr from cudnn.frost.device import multiprocessor_count from ..common.thd import TENSOR_MAP_QWORDS, emit_checkpoint_seq_descs, emit_seq_descs from .gdn2_prefill_config import CFG + from cudnn.frost.tile_dsl.barrier import ( advance, MBarrier, @@ -602,14 +603,10 @@ def super_mma_warp( tinv_lo1, tinv_hi1 = f16x2_to_f32(tinv_p1, dtype=cfg.io_dtype) tinv_lo2, tinv_hi2 = f16x2_to_f32(tinv_p2, dtype=cfg.io_dtype) tinv_lo3, tinv_hi3 = f16x2_to_f32(tinv_p3, dtype=cfg.io_dtype) - tinv_acc[0] = tinv_lo0 + upd_acc[0] - tinv_acc[1] = tinv_hi0 + upd_acc[1] - tinv_acc[2] = tinv_lo1 + upd_acc[2] - tinv_acc[3] = tinv_hi1 + upd_acc[3] - tinv_acc[4] = tinv_lo2 + upd_acc[4] - tinv_acc[5] = tinv_hi2 + upd_acc[5] - tinv_acc[6] = tinv_lo3 + upd_acc[6] - tinv_acc[7] = tinv_hi3 + upd_acc[7] + tinv_acc[0], tinv_acc[1] = fadd2(tinv_lo0, tinv_hi0, upd_acc[0], upd_acc[1]) + tinv_acc[2], tinv_acc[3] = fadd2(tinv_lo1, tinv_hi1, upd_acc[2], upd_acc[3]) + tinv_acc[4], tinv_acc[5] = fadd2(tinv_lo2, tinv_hi2, upd_acc[4], upd_acc[5]) + tinv_acc[6], tinv_acc[7] = fadd2(tinv_lo3, tinv_hi3, upd_acc[6], upd_acc[7]) bars.mb_intermediate_done[intermediate_stage].wait(((global_chunk // cfg.smem_intermediate_stages) + 1) % 2) nvvm.stmatrix( @@ -1296,6 +1293,9 @@ def compute0_warp_group( raw_q_regs[reg_base + dim_offset] = q_val raw_k_regs[reg_base + dim_offset] = k_val beta_val = raw_beta_vec_f32[dim_offset] + if cutlass.const_expr(cfg.beta_sigmoid): + half = cutlass.Float32(0.5) + beta_val = (cute.math.tanh(beta_val * half, approx=True) * half + half).to(cfg.io_dtype).to(cutlass.Float32) raw_beta_regs[reg_base + dim_offset] = beta_val if cutlass.const_expr(cfg.l2norm): if cutlass.const_expr(dim_offset % 2 == 0): @@ -2078,6 +2078,7 @@ def host( stream, ) -> None: num_sequences = cu_seqlens.shape[0] - 1 + grid_shape = (cfg.max_active_clusters, 1, 1) kernel( cfg, @@ -2450,6 +2451,7 @@ class Gdn2Cfg: l2norm: bool safe_gate: bool gate_scale_log2: float + beta_sigmoid: bool q_ratio: int k_ratio: int v_ratio: int @@ -2535,6 +2537,7 @@ def build_cfg( l2norm: bool, safe_gate: bool, gate_scale_log2: float, + beta_sigmoid: bool, q_ratio: int, k_ratio: int, v_ratio: int, @@ -2555,6 +2558,7 @@ def build_cfg( l2norm=l2norm, safe_gate=safe_gate, gate_scale_log2=gate_scale_log2, + beta_sigmoid=beta_sigmoid, q_ratio=q_ratio, k_ratio=k_ratio, v_ratio=v_ratio, @@ -2606,16 +2610,17 @@ def build_cfg( TENSORMAP_STATIC_SLOTS = 0 -@cute.kernel -def build_all_descs_kernel( - base_q: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_k: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_v: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_gate: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_beta: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_w: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_o: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_checkpoint: cutlass.GridConstant[cuda.tensor_map.TensorMap], +@cute.jit +def build_descs_body( + widx, + base_q, + base_k, + base_v, + base_gate, + base_beta, + base_w, + base_o, + base_checkpoint, desc_ws: cute.Tensor, cu_seqlens: cute.Tensor, q: cute.Tensor, @@ -2637,12 +2642,12 @@ def build_all_descs_kernel( checkpoint_row_stride: cutlass.Int32, checkpoint_every_n: cutlass.Int32, ) -> None: - """Single-launch builder for the per-BATCH descriptor arrays (one warp - per array; warp ``i`` emits array ``i`` and release-fences its slots). - Heads are load coordinates, so only the sequence base and token extent - are patched per slot.""" - tidx, _, _ = cute.arch.thread_idx() - widx = cutlass.Int32(tidx) // cutlass.Int32(32) + """Per-BATCH descriptor-array build, one warp per array (warp ``i`` + emits array ``i`` and release-fences its slots; heads are load + coordinates, so only the sequence base and token extent are patched per + slot): the body of the standalone builder kernel, also run by the fused + main kernel's CTA 0 prologue (warps past the array count fall through + the widx guards).""" arr_words = n_batch * cutlass.Int32(TENSOR_MAP_QWORDS) sub0 = cute.make_tensor(desc_ws.iterator, cute.make_layout((arr_words,), stride=(1,))) sub1 = cute.make_tensor(desc_ws.iterator + arr_words, cute.make_layout((arr_words,), stride=(1,))) @@ -2688,10 +2693,112 @@ def build_all_descs_kernel( nvvm.fence_proxy_release(nvvm.MemScope.GPU, from_proxy=nvvm.Proxy.GENERIC, to_proxy=nvvm.Proxy.TENSORMAP) +@cute.kernel +def prologue_kernel( + order_gen: cutlass.Constexpr[bool], + has_sched: cutlass.Constexpr[bool], + b_t: cutlass.Constexpr[int], + base_q: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_k: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_v: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_gate: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_beta: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_w: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_o: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_checkpoint: cutlass.GridConstant[cuda.tensor_map.TensorMap], + desc_ws: cute.Tensor, + cu_seqlens: cute.Tensor, + q: cute.Tensor, + k: cute.Tensor, + v: cute.Tensor, + gate: cute.Tensor, + beta: cute.Tensor, + w: cute.Tensor, + o: cute.Tensor, + state_checkpoints: cute.Tensor | None, + mStaging: cute.Tensor | None, + mCount: cute.Tensor, + mWorkItems: cute.Tensor, + mSched: cute.Tensor | None, + n_batch: cutlass.Int32, + q_row_stride: cutlass.Int32, + k_row_stride: cutlass.Int32, + v_row_stride: cutlass.Int32, + gate_row_stride: cutlass.Int32, + beta_row_stride: cutlass.Int32, + w_row_stride: cutlass.Int32, + o_row_stride: cutlass.Int32, + checkpoint_row_stride: cutlass.Int32, + checkpoint_every_n: cutlass.Int32, +) -> None: + """Single-CTA prologue: LPT-order the work-item table and zero the sched + rings via :func:`order_body`, then build the per-batch TMA-descriptor + arrays via :func:`build_descs_body`, one warp per array (the extra warps + only take part in the order phase).""" + tidx, _, _ = cute.arch.thread_idx() + tidx = cutlass.Int32(tidx) + widx = tidx // cutlass.Int32(32) + sKey = cutlass.Array(cutlass.Int32, ORDER_CAPACITY, space=cutlass.AddressSpace.smem, alignment=16) + sIdx = cutlass.Array(cutlass.Int32, ORDER_CAPACITY, space=cutlass.AddressSpace.smem, alignment=16) + sSpread = cutlass.Array(cutlass.Int32, 2, space=cutlass.AddressSpace.smem, alignment=8) + n_heads_out = cutlass.Int32(gate.shape[1]) + order_body( + order_gen, + has_sched, + b_t, + ORDER_THREADS, + ORDER_ELEMS, + tidx, + n_heads_out, + n_heads_out * n_batch, + cu_seqlens, + mStaging, + mCount, + mWorkItems, + mSched, + sKey, + sIdx, + sSpread, + ) + build_descs_body( + widx, + base_q, + base_k, + base_v, + base_gate, + base_beta, + base_w, + base_o, + base_checkpoint, + desc_ws, + cu_seqlens, + q, + k, + v, + gate, + beta, + w, + o, + state_checkpoints, + n_batch, + q_row_stride, + k_row_stride, + v_row_stride, + gate_row_stride, + beta_row_stride, + w_row_stride, + o_row_stride, + checkpoint_row_stride, + checkpoint_every_n, + ) + + @cute.jit -def build_descs( +def prologue( io_dtype: cutlass.Constexpr, b_t: cutlass.Constexpr[int], + order_gen: cutlass.Constexpr[bool], + has_sched: cutlass.Constexpr[bool], q: cute.Tensor, k: cute.Tensor, v: cute.Tensor, @@ -2701,12 +2808,17 @@ def build_descs( o: cute.Tensor, state_checkpoints: cute.Tensor | None, cu_seqlens: cute.Tensor, + work_item_staging: cute.Tensor | None, + work_count: cute.Tensor, + work_items: cute.Tensor, + sched_ctr: cute.Tensor | None, tensormap_workspace: cute.Tensor, checkpoint_every_n: cutlass.Int32, stream: cuda_driver.CUstream, ): - """Build the 8 per-(batch, head) TMA-descriptor arrays (q, k, v, gate, - beta, w, o, state_checkpoints) into ``tensormap_workspace``. + """One-launch prologue: LPT-order the work items and build the 8 + per-(batch, head) TMA-descriptor arrays (q, k, v, gate, beta, w, o, + state_checkpoints) into ``tensormap_workspace``. Launched on every execute: the descriptors fold cu_seqlens contents into GLOBAL_ADDRESS and GLOBAL_DIM, which the host cannot read without a D2H sync. @@ -2756,8 +2868,10 @@ def build_descs( ), ) base_checkpoint = cuda.create_tensor_map_tiled_from_view(checkpoint_view, box_dims=(tma_granu_elems, d_k, 1, 1), stride_order=(0, 1, 2, 3), swizzle=swz) - n_warps = 8 if state_checkpoints is not None else 7 - build_all_descs_kernel( + prologue_kernel( + order_gen, + has_sched, + b_t, base_q, base_k, base_v, @@ -2776,6 +2890,10 @@ def build_descs( w, o, state_checkpoints, + work_item_staging, + work_count, + work_items, + sched_ctr, cutlass.Int32(batch_size), cutlass.Int32(q.stride[0]), cutlass.Int32(k.stride[0]), @@ -2786,7 +2904,7 @@ def build_descs( cutlass.Int32(o.stride[0]), cutlass.Int32(state_checkpoints.stride[0] if state_checkpoints is not None else 0), checkpoint_every_n, - ).launch(grid=(1, 1, 1), block=(32 * n_warps, 1, 1), stream=stream) + ).launch(grid=(1, 1, 1), block=(ORDER_THREADS, 1, 1), stream=stream) # ---- Torch adapter / host-side compilation --------------------------------------- @@ -2806,7 +2924,9 @@ def get_compiled_cache( l2norm: bool, safe_gate: bool, gate_lower_bound: float, + beta_sigmoid: bool, dyn_sched: bool, + order_gen: bool, ): """Return a mutable dict that lazily stores the compiled kernel.""" return {} @@ -2821,6 +2941,7 @@ def compile( l2norm: bool, safe_gate: bool, gate_scale_log2: float, + beta_sigmoid: bool, q_ratio: int, k_ratio: int, v_ratio: int, @@ -2858,6 +2979,7 @@ def compile( l2norm=l2norm, safe_gate=safe_gate, gate_scale_log2=gate_scale_log2, + beta_sigmoid=beta_sigmoid, q_ratio=q_ratio, k_ratio=k_ratio, v_ratio=v_ratio, @@ -2911,9 +3033,11 @@ def chunk_gdn2_sm100( gate_lower_bound: float = DEFAULT_GATE_LOWER_BOUND, a_log=None, dt_bias=None, + use_beta_sigmoid: bool = False, work_items=None, work_count=None, sched_ctr=None, + work_item_scratch=None, *, tensormap_workspace, stream, @@ -2931,7 +3055,8 @@ def chunk_gdn2_sm100( gate: ``(total_tokens, HO, DK)`` float32. Natural-log decay unless ``safe_gate``, which applies the safe-gate transform ``lower_bound * sigmoid(exp(a_log) * (gate + dt_bias))``. - beta: ``(total_tokens, HO, DK)`` io dtype, channel-wise erase gate + beta: ``(total_tokens, HO, DK)`` io dtype, channel-wise erase gate. + Post-sigmoid, or logits when ``use_beta_sigmoid`` w: ``(total_tokens, HO, DV)`` io dtype, channel-wise write gate output: ``(total_tokens, HO, DV)`` float16/bfloat16, pre-allocated cu_seqlens: ``(num_seqs + 1,)`` int32 @@ -2950,6 +3075,7 @@ def chunk_gdn2_sm100( safe_gate: interpret ``gate`` through the safe-gate transform a_log: ``(HO,)`` float32, safe-gate per-head log-amplitude (None = 0) dt_bias: ``(HO, DK)`` float32, safe-gate channel bias (None = 0) + use_beta_sigmoid: ``beta`` holds logits; sigmoid in-kernel work_items: ``(max_items, 8)`` int32 work-item table from ``common/split_k.py`` (REQUIRED; an uncut table row is the whole (b, h) sequence). Each item computes chunks ``[cstart, wend)`` @@ -2973,6 +3099,7 @@ def chunk_gdn2_sm100( if work_items is None or work_count is None: raise ValueError("work_items/work_count are required (the split-table stage builds them for every launch)") dyn_sched = sched_ctr is not None + order_gen = work_item_scratch is None if initial_state is not None: state_dtype_src = initial_state.dtype @@ -3009,7 +3136,9 @@ def chunk_gdn2_sm100( use_qk_l2norm_in_kernel, safe_gate, gate_lower_bound, + use_beta_sigmoid, dyn_sched, + order_gen, ) if "compiled" not in cache: @@ -3053,6 +3182,7 @@ def chunk_gdn2_sm100( use_qk_l2norm_in_kernel, safe_gate, gate_scale_log2, + use_beta_sigmoid, q_ratio, k_ratio, v_ratio, @@ -3080,37 +3210,56 @@ def chunk_gdn2_sm100( stream=cu_stream, ) - compiled = cache["compiled"] state_checkpoints_for_descs = output_state_checkpoints if enable_checkpoints else None - # desc build runs every execute by contract (cu contents are data; - # buffer pointers may change) — capture-safe, single tiny launch - if cache.get("build_descs_has_state_checkpoints") != (state_checkpoints_for_descs is not None): - cache.pop("build_descs", None) - cache["build_descs_has_state_checkpoints"] = state_checkpoints_for_descs is not None - if "build_descs" not in cache: + if "prologue" not in cache: io_dtype = get_dtype(q.dtype) - - cu_bd = from_dlpack(cu_seqlens, assumed_align=8).mark_layout_dynamic() - ws_bd = from_dlpack(tensormap_workspace, assumed_align=128).mark_layout_dynamic() - cache["build_descs"] = cute.compile( - build_descs, + q_pl = from_dlpack(q, assumed_align=16).mark_layout_dynamic(leading_dim=2) + k_pl = from_dlpack(k, assumed_align=16).mark_layout_dynamic(leading_dim=2) + v_pl = from_dlpack(v, assumed_align=16).mark_layout_dynamic(leading_dim=2) + gate_pl = from_dlpack(gate, assumed_align=16).mark_layout_dynamic(leading_dim=2) + beta_pl = from_dlpack(beta, assumed_align=16).mark_layout_dynamic(leading_dim=2) + w_pl = from_dlpack(w, assumed_align=16).mark_layout_dynamic(leading_dim=2) + o_pl = from_dlpack(output, assumed_align=16).mark_layout_dynamic(leading_dim=2) + cu_pl = from_dlpack(cu_seqlens, assumed_align=8).mark_layout_dynamic() + ws_pl = from_dlpack(tensormap_workspace, assumed_align=128).mark_layout_dynamic() + state_checkpoints_pl = None + if state_checkpoints_for_descs is not None: + state_checkpoints_pl = from_dlpack(state_checkpoints_for_descs, assumed_align=16).mark_layout_dynamic(leading_dim=3) + staging_pl = None + if not order_gen: + staging_pl = from_dlpack(work_item_scratch, assumed_align=16) + staging_pl.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1), divisibility=1) + work_items_pl = from_dlpack(work_items, assumed_align=16) + work_items_pl.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1), divisibility=1) + work_count_pl = from_dlpack(work_count, assumed_align=4).mark_layout_dynamic() + sched_pl = None + if dyn_sched: + sched_pl = from_dlpack(sched_ctr, assumed_align=4).mark_layout_dynamic() + cache["prologue"] = cute.compile( + prologue, io_dtype, CFG.B_T, - from_dlpack(q, assumed_align=16).mark_layout_dynamic(leading_dim=2), - from_dlpack(k, assumed_align=16).mark_layout_dynamic(leading_dim=2), - from_dlpack(v, assumed_align=16).mark_layout_dynamic(leading_dim=2), - from_dlpack(gate, assumed_align=16).mark_layout_dynamic(leading_dim=2), - from_dlpack(beta, assumed_align=16).mark_layout_dynamic(leading_dim=2), - from_dlpack(w, assumed_align=16).mark_layout_dynamic(leading_dim=2), - from_dlpack(output, assumed_align=16).mark_layout_dynamic(leading_dim=2), - None if state_checkpoints_for_descs is None else from_dlpack(state_checkpoints_for_descs, assumed_align=16).mark_layout_dynamic(leading_dim=3), - cu_bd, - ws_bd, + order_gen, + dyn_sched, + q_pl, + k_pl, + v_pl, + gate_pl, + beta_pl, + w_pl, + o_pl, + state_checkpoints_pl, + cu_pl, + staging_pl, + work_count_pl, + work_items_pl, + sched_pl, + ws_pl, cutlass.Int32(checkpoint_every_n_tokens), cu_stream, options="--enable-tvm-ffi", ) - cache["build_descs"]( + cache["prologue"]( q, k, v, @@ -3120,11 +3269,15 @@ def chunk_gdn2_sm100( output, state_checkpoints_for_descs, cu_seqlens, + work_item_scratch if not order_gen else None, + work_count, + work_items, + sched_ctr, tensormap_workspace, checkpoint_every_n_tokens, cu_stream, ) - compiled( + cache["compiled"]( q, k, v, @@ -3145,3 +3298,73 @@ def chunk_gdn2_sm100( scale, cu_stream, ) + return cache + + +def run_prefill( + cache, + q, + k, + v, + gate, + a_log, + dt_bias, + beta, + w, + cu_seqlens, + initial_state, + output, + output_state, + output_state_checkpoints, + work_items, + work_count, + sched_ctr, + work_item_scratch, + tensormap_workspace, + checkpoint_every_n_tokens, + scale, + stream, +) -> None: + """Replay the compiled plan: the prologue launch, then the main launch. + The caller owns the contract, which the plan validated at build, so + nothing here raises.""" + cu_stream = cuda_driver.CUstream(int(stream)) + cache["prologue"]( + q, + k, + v, + gate, + beta, + w, + output, + output_state_checkpoints, + cu_seqlens, + work_item_scratch, + work_count, + work_items, + sched_ctr, + tensormap_workspace, + checkpoint_every_n_tokens, + cu_stream, + ) + cache["compiled"]( + q, + k, + v, + gate, + a_log, + dt_bias, + beta, + w, + cu_seqlens, + initial_state, + output, + output_state, + work_items, + work_count, + sched_ctr, + tensormap_workspace, + checkpoint_every_n_tokens, + scale, + cu_stream, + ) diff --git a/python/cudnn/linear_attention/frost/kernel/gdn2_recompute_f16.py b/python/cudnn/linear_attention/frost/kernel/gdn2_recompute_f16.py index 399602fe8..9f1990d01 100644 --- a/python/cudnn/linear_attention/frost/kernel/gdn2_recompute_f16.py +++ b/python/cudnn/linear_attention/frost/kernel/gdn2_recompute_f16.py @@ -90,12 +90,13 @@ import cutlass.cute as cute from cutlass.cute.runtime import from_dlpack -from ..common.split_k import decode_work_item +from ..common.split_k import ORDER_CAPACITY, ORDER_ELEMS, ORDER_THREADS, decode_work_item, order_body from ..common.host import get_dtype from cudnn.frost.buffers import current_device_id, data_ptr from cudnn.frost.device import multiprocessor_count from ..common.thd import TENSOR_MAP_QWORDS, emit_checkpoint_seq_descs, emit_seq_descs from .gdn2_recompute_config import CFG + from cudnn.frost.tile_dsl.barrier import ( advance, MBarrier, @@ -572,14 +573,10 @@ def super_mma_warp( tinv_lo1, tinv_hi1 = f16x2_to_f32(tinv_p1, dtype=cfg.io_dtype) tinv_lo2, tinv_hi2 = f16x2_to_f32(tinv_p2, dtype=cfg.io_dtype) tinv_lo3, tinv_hi3 = f16x2_to_f32(tinv_p3, dtype=cfg.io_dtype) - tinv_acc[0] = tinv_lo0 + upd_acc[0] - tinv_acc[1] = tinv_hi0 + upd_acc[1] - tinv_acc[2] = tinv_lo1 + upd_acc[2] - tinv_acc[3] = tinv_hi1 + upd_acc[3] - tinv_acc[4] = tinv_lo2 + upd_acc[4] - tinv_acc[5] = tinv_hi2 + upd_acc[5] - tinv_acc[6] = tinv_lo3 + upd_acc[6] - tinv_acc[7] = tinv_hi3 + upd_acc[7] + tinv_acc[0], tinv_acc[1] = fadd2(tinv_lo0, tinv_hi0, upd_acc[0], upd_acc[1]) + tinv_acc[2], tinv_acc[3] = fadd2(tinv_lo1, tinv_hi1, upd_acc[2], upd_acc[3]) + tinv_acc[4], tinv_acc[5] = fadd2(tinv_lo2, tinv_hi2, upd_acc[4], upd_acc[5]) + tinv_acc[6], tinv_acc[7] = fadd2(tinv_lo3, tinv_hi3, upd_acc[6], upd_acc[7]) bars.mb_t_inv_done[intermediate_stage].wait(((global_chunk // cfg.smem_intermediate_stages) + 1) % 2) nvvm.stmatrix( @@ -1046,10 +1043,8 @@ def compute0_warp_group( # ---- optional K L2-norm + K_inv staging ------------------------------ if cutlass.const_expr(cfg.l2norm): - kk0_lo = opaque_f32_zero() - kk0_hi = opaque_f32_zero() - kk1_lo = opaque_f32_zero() - kk1_hi = opaque_f32_zero() + kk_lo = opaque_f32_zero() + kk_hi = opaque_f32_zero() for dim_half in cutlass.range_constexpr(2): dim_base = dim_half * (cfg.d_k // 2) + lane_in_row_group * 8 reg_base = dim_half * 8 @@ -1064,16 +1059,20 @@ def compute0_warp_group( k_val = raw_k_frag_f32[dim_offset] raw_k_regs[reg_base + dim_offset] = k_val beta_val = raw_beta_frag_f32[dim_offset] + if cutlass.const_expr(cfg.beta_sigmoid): + half = cutlass.Float32(0.5) + beta_val = (cute.math.tanh(beta_val * half, approx=True) * half + half).to(cfg.io_dtype).to(cutlass.Float32) raw_beta_regs[reg_base + dim_offset] = beta_val - if cutlass.const_expr(cfg.l2norm): - if cutlass.const_expr(dim_offset % 2 == 0): - kk0_lo, kk0_hi = ffma2(k_val, k_val, k_val, k_val, kk0_lo, kk0_hi) - else: - kk1_lo, kk1_hi = ffma2(k_val, k_val, k_val, k_val, kk1_lo, kk1_hi) + if cutlass.const_expr(cfg.l2norm): + # even dims in the lo lane, odd in the hi lane: same terms, same order + for dim_pair in cutlass.range_constexpr(4): + k_even = raw_k_frag_f32[2 * dim_pair] + k_odd = raw_k_frag_f32[2 * dim_pair + 1] + kk_lo, kk_hi = ffma2(k_even, k_odd, k_even, k_odd, kk_lo, kk_hi) k_inv_norm = opaque_one if cutlass.const_expr(cfg.l2norm): - k_sum_sq = kk0_hi + kk1_hi + k_sum_sq = kk_lo + kk_hi k_sum_sq = k_sum_sq + cutlass.Float32(nvvm.shfl_sync(0xFFFFFFFF, k_sum_sq, 4, 31, kind=nvvm.Shfl.BFLY)) k_sum_sq = k_sum_sq + cutlass.Float32(nvvm.shfl_sync(0xFFFFFFFF, k_sum_sq, 2, 31, kind=nvvm.Shfl.BFLY)) k_sum_sq = k_sum_sq + cutlass.Float32(nvvm.shfl_sync(0xFFFFFFFF, k_sum_sq, 1, 31, kind=nvvm.Shfl.BFLY)) @@ -1701,6 +1700,7 @@ def host( stream, ) -> None: num_sequences = cu_seqlens.shape[0] - 1 + grid_shape = (cfg.max_active_clusters, 1, 1) kernel( cfg, @@ -2019,6 +2019,7 @@ class Gdn2RecomputeCfg: l2norm: bool safe_gate: bool gate_scale_log2: float + beta_sigmoid: bool k_ratio: int v_ratio: int n_heads_out: int @@ -2096,6 +2097,7 @@ def build_cfg( l2norm: bool, safe_gate: bool, gate_scale_log2: float, + beta_sigmoid: bool, k_ratio: int, v_ratio: int, n_heads_out: int, @@ -2115,6 +2117,7 @@ def build_cfg( l2norm=l2norm, safe_gate=safe_gate, gate_scale_log2=gate_scale_log2, + beta_sigmoid=beta_sigmoid, k_ratio=k_ratio, v_ratio=v_ratio, n_heads_out=n_heads_out, @@ -2160,14 +2163,15 @@ def build_cfg( TENSORMAP_STATIC_SLOTS = 0 -@cute.kernel -def build_all_descs_kernel( - base_k: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_v: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_gate: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_beta: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_w: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_checkpoint: cutlass.GridConstant[cuda.tensor_map.TensorMap], +@cute.jit +def build_descs_body( + widx, + base_k, + base_v, + base_gate, + base_beta, + base_w, + base_checkpoint, tensormap_workspace: cute.Tensor, cu_seqlens: cute.Tensor, k: cute.Tensor, @@ -2185,9 +2189,9 @@ def build_all_descs_kernel( checkpoint_row_stride: cutlass.Int32, checkpoint_every_n: cutlass.Int32, ) -> None: - """Single-launch builder for the per-batch TMA-descriptor arrays.""" - tidx, _, _ = cute.arch.thread_idx() - widx = cutlass.Int32(tidx) // cutlass.Int32(32) + """Per-batch descriptor-array build, one warp per array. Runs inside the + prologue kernel after its order pass; warps past the array count fall + through the widx guards.""" arr_words = n_batch * cutlass.Int32(TENSOR_MAP_QWORDS) sub0 = cute.make_tensor(tensormap_workspace.iterator, cute.make_layout((arr_words,), stride=(1,))) sub1 = cute.make_tensor(tensormap_workspace.iterator + arr_words, cute.make_layout((arr_words,), stride=(1,))) @@ -2223,10 +2227,104 @@ def build_all_descs_kernel( nvvm.fence_proxy_release(nvvm.MemScope.GPU, from_proxy=nvvm.Proxy.GENERIC, to_proxy=nvvm.Proxy.TENSORMAP) +@cute.kernel +def prologue_kernel( + run_order: cutlass.Constexpr[bool], + order_gen: cutlass.Constexpr[bool], + has_sched: cutlass.Constexpr[bool], + b_t: cutlass.Constexpr[int], + base_k: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_v: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_gate: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_beta: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_w: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_checkpoint: cutlass.GridConstant[cuda.tensor_map.TensorMap], + tensormap_workspace: cute.Tensor, + cu_seqlens: cute.Tensor, + k: cute.Tensor, + v: cute.Tensor, + gate: cute.Tensor, + beta: cute.Tensor, + w: cute.Tensor, + state_checkpoints: cute.Tensor | None, + mStaging: cute.Tensor | None, + mCount: cute.Tensor, + mWorkItems: cute.Tensor, + mSched: cute.Tensor | None, + n_batch: cutlass.Int32, + k_row_stride: cutlass.Int32, + v_row_stride: cutlass.Int32, + gate_row_stride: cutlass.Int32, + beta_row_stride: cutlass.Int32, + w_row_stride: cutlass.Int32, + checkpoint_row_stride: cutlass.Int32, + checkpoint_every_n: cutlass.Int32, +) -> None: + """Single-CTA prologue. Under ``run_order`` this kernel is the first + work-item-table consumer, so it LPT-orders the table and zeroes both + consumers' sched rings via :func:`order_body`; it then builds the + per-batch TMA-descriptor arrays via :func:`build_descs_body`, one warp + per array (the extra warps only take part in the order phase).""" + tidx, _, _ = cute.arch.thread_idx() + tidx = cutlass.Int32(tidx) + widx = tidx // cutlass.Int32(32) + if cutlass.const_expr(run_order): + sKey = cutlass.Array(cutlass.Int32, ORDER_CAPACITY, space=cutlass.AddressSpace.smem, alignment=16) + sIdx = cutlass.Array(cutlass.Int32, ORDER_CAPACITY, space=cutlass.AddressSpace.smem, alignment=16) + sSpread = cutlass.Array(cutlass.Int32, 2, space=cutlass.AddressSpace.smem, alignment=8) + n_heads_out = cutlass.Int32(gate.shape[1]) + order_body( + order_gen, + has_sched, + b_t, + ORDER_THREADS, + ORDER_ELEMS, + tidx, + n_heads_out, + n_heads_out * n_batch, + cu_seqlens, + mStaging, + mCount, + mWorkItems, + mSched, + sKey, + sIdx, + sSpread, + ) + build_descs_body( + widx, + base_k, + base_v, + base_gate, + base_beta, + base_w, + base_checkpoint, + tensormap_workspace, + cu_seqlens, + k, + v, + gate, + beta, + w, + state_checkpoints, + n_batch, + k_row_stride, + v_row_stride, + gate_row_stride, + beta_row_stride, + w_row_stride, + checkpoint_row_stride, + checkpoint_every_n, + ) + + @cute.jit -def build_descs( +def prologue( io_dtype: cutlass.Constexpr, b_t: cutlass.Constexpr[int], + run_order: cutlass.Constexpr[bool], + order_gen: cutlass.Constexpr[bool], + has_sched: cutlass.Constexpr[bool], k: cute.Tensor, v: cute.Tensor, gate: cute.Tensor, @@ -2234,12 +2332,18 @@ def build_descs( w: cute.Tensor, state_checkpoints: cute.Tensor | None, cu_seqlens: cute.Tensor, + work_item_staging: cute.Tensor | None, + work_count: cute.Tensor, + work_items: cute.Tensor, + sched_all: cute.Tensor | None, tensormap_workspace: cute.Tensor, checkpoint_every_n: cutlass.Int32, stream: cuda_driver.CUstream, ): - """Build the 6 per-batch TMA-descriptor arrays (k, v, gate, beta, w, state_checkpoints) - into ``tensormap_workspace``.""" + """One-launch prologue: LPT-order the work items (when this kernel is + the table's first consumer) and build the 6 per-batch TMA-descriptor + arrays (k, v, gate, beta, w, state_checkpoints) into + ``tensormap_workspace``.""" h_k = k.shape[1] h_v = v.shape[1] ho = gate.shape[1] @@ -2273,8 +2377,11 @@ def build_descs( ), ) base_checkpoint = cuda.create_tensor_map_tiled_from_view(checkpoint_view, box_dims=(tma_box_elems, d_k, 1, 1), stride_order=(0, 1, 2, 3), swizzle=swz) - n_warps = 6 if state_checkpoints is not None else 5 - build_all_descs_kernel( + prologue_kernel( + run_order, + order_gen, + has_sched, + b_t, base_k, base_v, base_gate, @@ -2289,6 +2396,10 @@ def build_descs( beta, w, state_checkpoints, + work_item_staging, + work_count, + work_items, + sched_all, cutlass.Int32(batch_size), cutlass.Int32(k.stride[0]), cutlass.Int32(v.stride[0]), @@ -2297,7 +2408,7 @@ def build_descs( cutlass.Int32(w.stride[0]), cutlass.Int32(state_checkpoints.stride[0] if state_checkpoints is not None else 0), checkpoint_every_n, - ).launch(grid=(1, 1, 1), block=(32 * n_warps, 1, 1), stream=stream) + ).launch(grid=(1, 1, 1), block=(ORDER_THREADS, 1, 1), stream=stream) # ---- Torch adapter / host-side compilation --------------------------------------- @@ -2317,7 +2428,11 @@ def get_compiled_cache( l2norm: bool, safe_gate: bool, gate_lower_bound: float, + beta_sigmoid: bool, dyn_sched: bool, + order_in_prologue: bool, + order_gen: bool, + has_sched: bool, ): """Return a mutable dict that lazily stores the compiled kernel.""" return {} @@ -2332,6 +2447,7 @@ def compile( l2norm: bool, safe_gate: bool, gate_scale_log2: float, + beta_sigmoid: bool, k_ratio: int, v_ratio: int, n_heads_out: int, @@ -2365,6 +2481,7 @@ def compile( l2norm=l2norm, safe_gate=safe_gate, gate_scale_log2=gate_scale_log2, + beta_sigmoid=beta_sigmoid, k_ratio=k_ratio, v_ratio=v_ratio, n_heads_out=n_heads_out, @@ -2411,9 +2528,13 @@ def chunk_gdn2_recompute_sm100( gate_lower_bound: float = DEFAULT_GATE_LOWER_BOUND, a_log=None, dt_bias=None, + use_beta_sigmoid: bool = False, work_items=None, work_count=None, sched_ctr=None, + sched_all=None, + work_item_scratch=None, + order_in_prologue: bool = False, *, tensormap_workspace, stream, @@ -2430,7 +2551,8 @@ def chunk_gdn2_recompute_sm100( gate: ``(total_tokens, HO, DK)`` float32. Natural-log decay unless ``safe_gate``, which applies the safe-gate transform ``lower_bound * sigmoid(exp(a_log) * (gate + dt_bias))``. - beta: ``(total_tokens, HO, DK)`` io dtype, channel-wise erase gate + beta: ``(total_tokens, HO, DK)`` io dtype, channel-wise erase gate. + Post-sigmoid, or logits when ``use_beta_sigmoid`` w: ``(total_tokens, HO, DV)`` io dtype, channel-wise write gate cu_seqlens: ``(num_seqs + 1,)`` int32 initial_state: ``(num_seqs, HO, DK, DV)`` float32/bfloat16, or None @@ -2447,6 +2569,7 @@ def chunk_gdn2_recompute_sm100( safe_gate: interpret ``gate`` through the safe-gate transform a_log: ``(HO,)`` float32, safe-gate per-head log-amplitude (None = 0) dt_bias: ``(HO, DK)`` float32, safe-gate channel bias (None = 0) + use_beta_sigmoid: ``beta`` holds logits; sigmoid in-kernel work_items: ``(max_items, 8)`` int32 work-item table from ``common/split_k.py`` (REQUIRED; an uncut table row is the whole (b, h) sequence). Each item computes chunks ``[cstart, wend)`` @@ -2469,6 +2592,9 @@ def chunk_gdn2_recompute_sm100( if work_items is None or work_count is None: raise ValueError("work_items/work_count are required (the split-table stage builds them for every launch)") dyn_sched = sched_ctr is not None + order_gen = work_item_scratch is None + if order_in_prologue and sched_all is None: + raise ValueError("order_in_prologue requires sched_all (the prologue zeroes both consumers' sched rings)") if initial_state is not None: state_dtype_src = initial_state.dtype @@ -2504,7 +2630,11 @@ def chunk_gdn2_recompute_sm100( use_qk_l2norm_in_kernel, safe_gate, gate_lower_bound, + use_beta_sigmoid, dyn_sched, + order_in_prologue, + order_gen, + sched_all is not None, ) if "compiled" not in cache: @@ -2546,6 +2676,7 @@ def chunk_gdn2_recompute_sm100( use_qk_l2norm_in_kernel, safe_gate, gate_scale_log2, + use_beta_sigmoid, k_ratio, v_ratio, HO, @@ -2569,35 +2700,53 @@ def chunk_gdn2_recompute_sm100( stream=cu_stream, ) - compiled = cache["compiled"] state_checkpoints_for_descs = output_state_checkpoints if enable_checkpoints else None - # desc build runs every execute by contract (cu contents are data; - # buffer pointers may change) - capture-safe, single tiny launch - if cache.get("build_descs_has_state_checkpoints") != (state_checkpoints_for_descs is not None): - cache.pop("build_descs", None) - cache["build_descs_has_state_checkpoints"] = state_checkpoints_for_descs is not None - if "build_descs" not in cache: + if "prologue" not in cache: io_dtype = get_dtype(k.dtype) - - cu_bd = from_dlpack(cu_seqlens, assumed_align=8).mark_layout_dynamic() - ws_bd = from_dlpack(tensormap_workspace, assumed_align=128).mark_layout_dynamic() - cache["build_descs"] = cute.compile( - build_descs, + k_pl = from_dlpack(k, assumed_align=16).mark_layout_dynamic(leading_dim=2) + v_pl = from_dlpack(v, assumed_align=16).mark_layout_dynamic(leading_dim=2) + gate_pl = from_dlpack(gate, assumed_align=16).mark_layout_dynamic(leading_dim=2) + beta_pl = from_dlpack(beta, assumed_align=16).mark_layout_dynamic(leading_dim=2) + w_pl = from_dlpack(w, assumed_align=16).mark_layout_dynamic(leading_dim=2) + cu_pl = from_dlpack(cu_seqlens, assumed_align=8).mark_layout_dynamic() + ws_pl = from_dlpack(tensormap_workspace, assumed_align=128).mark_layout_dynamic() + state_checkpoints_pl = None + if state_checkpoints_for_descs is not None: + state_checkpoints_pl = from_dlpack(state_checkpoints_for_descs, assumed_align=16).mark_layout_dynamic(leading_dim=3) + staging_pl = None + if not order_gen: + staging_pl = from_dlpack(work_item_scratch, assumed_align=16) + staging_pl.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1), divisibility=1) + work_items_pl = from_dlpack(work_items, assumed_align=16) + work_items_pl.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1), divisibility=1) + work_count_pl = from_dlpack(work_count, assumed_align=4).mark_layout_dynamic() + sched_all_pl = None + if sched_all is not None: + sched_all_pl = from_dlpack(sched_all, assumed_align=4).mark_layout_dynamic() + cache["prologue"] = cute.compile( + prologue, io_dtype, CFG.B_T, - from_dlpack(k, assumed_align=16).mark_layout_dynamic(leading_dim=2), - from_dlpack(v, assumed_align=16).mark_layout_dynamic(leading_dim=2), - from_dlpack(gate, assumed_align=16).mark_layout_dynamic(leading_dim=2), - from_dlpack(beta, assumed_align=16).mark_layout_dynamic(leading_dim=2), - from_dlpack(w, assumed_align=16).mark_layout_dynamic(leading_dim=2), - None if state_checkpoints_for_descs is None else from_dlpack(state_checkpoints_for_descs, assumed_align=16).mark_layout_dynamic(leading_dim=3), - cu_bd, - ws_bd, + order_in_prologue, + order_gen, + sched_all is not None, + k_pl, + v_pl, + gate_pl, + beta_pl, + w_pl, + state_checkpoints_pl, + cu_pl, + staging_pl, + work_count_pl, + work_items_pl, + sched_all_pl, + ws_pl, cutlass.Int32(checkpoint_every_n_tokens), cu_stream, options="--enable-tvm-ffi", ) - cache["build_descs"]( + cache["prologue"]( k, v, gate, @@ -2605,11 +2754,15 @@ def chunk_gdn2_recompute_sm100( w, state_checkpoints_for_descs, cu_seqlens, + work_item_scratch if not order_gen else None, + work_count, + work_items, + sched_all, tensormap_workspace, checkpoint_every_n_tokens, cu_stream, ) - compiled( + cache["compiled"]( k, v, gate, @@ -2627,3 +2780,66 @@ def chunk_gdn2_recompute_sm100( checkpoint_every_n_tokens, cu_stream, ) + return cache + + +def run_recompute( + cache, + k, + v, + gate, + a_log, + dt_bias, + beta, + w, + cu_seqlens, + initial_state, + output_state, + output_state_checkpoints, + work_items, + work_count, + sched_ctr, + sched_all, + work_item_scratch, + tensormap_workspace, + checkpoint_every_n_tokens, + stream, +) -> None: + """Replay the compiled plan: the prologue launch, then the main launch. + The caller owns the contract, which the plan validated at build, so + nothing here raises.""" + cu_stream = cuda_driver.CUstream(int(stream)) + cache["prologue"]( + k, + v, + gate, + beta, + w, + output_state_checkpoints, + cu_seqlens, + work_item_scratch, + work_count, + work_items, + sched_all, + tensormap_workspace, + checkpoint_every_n_tokens, + cu_stream, + ) + cache["compiled"]( + k, + v, + gate, + a_log, + dt_bias, + beta, + w, + cu_seqlens, + initial_state, + output_state, + work_items, + work_count, + sched_ctr, + tensormap_workspace, + checkpoint_every_n_tokens, + cu_stream, + ) diff --git a/python/cudnn/linear_attention/frost/kernel/gdn_bprop_f16.py b/python/cudnn/linear_attention/frost/kernel/gdn_bprop_f16.py index 27fbda843..930d75ff1 100644 --- a/python/cudnn/linear_attention/frost/kernel/gdn_bprop_f16.py +++ b/python/cudnn/linear_attention/frost/kernel/gdn_bprop_f16.py @@ -70,7 +70,7 @@ dO 16384 1 state (checkpoint entry c-1) 32768 1 <-- io-dtype [DK,DV], TMA-loaded T_inv 8192 1 <-- inverse OUTPUT (upper tri = kernel-start zeros) - KK (pristine M_kk) 8192 1 <-- KK epi's only store; inverse input + dGate/dBeta + KK (strict-masked M_kk) 8192 1 <-- KK epi's only store; inverse input + dGate/dBeta A staging / sDa 8192 1 <-- ALIAS: A then the masked dA dM staging (sDm) 8192 1 <-- Step 8 -> dK dM-terms dstate_entry (sDstate) 32768 1 <-- f16 restage, dK-inter's A @@ -117,7 +117,8 @@ from cutlass.cute.runtime import from_dlpack from ..common.thd import emit_copy_desc, emit_checkpoint_seq_descs, emit_seq_descs, TENSOR_MAP_QWORDS -from ..common.split_k import decode_work_item +from ..common.split_k import ORDER_CAPACITY, ORDER_ELEMS, ORDER_THREADS, decode_work_item, order_body +from ..common.elementwise import softplus from ..common.host import get_dtype from cudnn.frost.buffers import current_device_id, data_ptr from cudnn.frost.device import multiprocessor_count @@ -733,6 +734,8 @@ def gate_beta_warp( mWorkItems, tidx, mGate, + mA_log, + mDt_bias, mBeta, mDgate, mDbeta, @@ -751,6 +754,8 @@ def gate_beta_warp( beta_store_index = PipelineState.start(phase=0) lidx = tidx % cfg.threads_per_warp n_cols = cfg.b_t // cfg.threads_per_warp + a_l2 = cutlass.Float32(0.0) + bias = cutlass.Float32(0.0) sched_state = PipelineState.start(phase=0) tile_idx = cutlass.Int32(bidx) FIRST_STATE_CHUNK = 0 if cfg.use_initial_state else 1 @@ -758,6 +763,11 @@ def gate_beta_warp( while tile_idx < total_tiles: batch_idx, head_idx, batch_start, batch_end, seqlen_b, num_chunks_b, wstart, wend, cstart, cend = decode_work_item(cfg, tile_idx, mWorkItems) num_item_chunks = cend - wstart + if cutlass.const_expr(cfg.safe_gate): + if num_item_chunks > 0: + # per-head transform constants, fixed for the whole tile + a_l2 = -cute.math.exp2(mA_log[head_idx].to(cutlass.Float32) * cutlass.Float32(RCP_LN2), fastmath=True) * cutlass.Float32(RCP_LN2) + bias = mDt_bias[head_idx].to(cutlass.Float32) # dGate/dBeta ownership: mask stores past the item's write range write_end = batch_start + wend * cfg.b_t write_end = write_end if write_end < batch_end else batch_end @@ -781,7 +791,11 @@ def gate_beta_warp( oob_neutral = cutlass.Float32(0.0) if cutlass.const_expr(cfg.log_gate) else cutlass.Float32(1.0) gate_vals[col] = gGate[pos] if pos_valid[col] else oob_neutral - if cutlass.const_expr(cfg.log_gate): + if cutlass.const_expr(cfg.safe_gate): + for col in cutlass.range_constexpr(n_cols): + contrib = a_l2 * softplus(gate_vals[col] + bias) + gate_vals[col] = contrib if pos_valid[col] else cutlass.Float32(0.0) + elif cutlass.const_expr(cfg.log_gate): for col in cutlass.range_constexpr(n_cols): gate_vals[col] = gate_vals[col] * cutlass.Float32(RCP_LN2) else: @@ -812,13 +826,25 @@ def gate_beta_warp( # ---- Beta load: GMEM -> SMEM (per-element cp.async) ------------------ beta_idx = beta_index.idx beta_index = advance(beta_index, cfg.smem_beta_stages) - for col in cutlass.range_constexpr(n_cols): - pos = lidx + col * cfg.threads_per_warp - src = gBeta.iterator + gBeta.layout((pos,)) - dst = sBeta.iterator + sBeta.layout((pos, 0, beta_idx)) - cp_size = cutlass.Int32(4) * cutlass.Int32(pos_valid[col]) - nvvm.cp_async_shared_global(dst, src, 4, nvvm.LoadCacheModifier.CA, cp_size=cp_size) - nvvm.cp_async_mbarrier_arrive(bars.mb_beta_ready[beta_idx].smem_ptr, noinc=True) + if cutlass.const_expr(cfg.beta_sigmoid): + # io-dtype logits -> sigmoid (tanh identity) -> fp32 SMEM + for col in cutlass.range_constexpr(n_cols): + pos = lidx + col * cfg.threads_per_warp + beta_value = cutlass.Float32(0.0) + if pos_valid[col]: + beta_value = gBeta[pos].to(cutlass.Float32) + half = cutlass.Float32(0.5) + beta_value = (cute.math.tanh(beta_value * half, approx=True) * half + half).to(mBeta.element_type).to(cutlass.Float32) + sBeta[pos, 0, beta_idx] = beta_value + bars.mb_beta_ready[beta_idx].arrive() + else: + for col in cutlass.range_constexpr(n_cols): + pos = lidx + col * cfg.threads_per_warp + src = gBeta.iterator + gBeta.layout((pos,)) + dst = sBeta.iterator + sBeta.layout((pos, 0, beta_idx)) + cp_size = cutlass.Int32(4) * cutlass.Int32(pos_valid[col]) + nvvm.cp_async_shared_global(dst, src, 4, nvvm.LoadCacheModifier.CA, cp_size=cp_size) + nvvm.cp_async_mbarrier_arrive(bars.mb_beta_ready[beta_idx].smem_ptr, noinc=True) for rev_idx in cutlass.range(num_item_chunks): # ---- prefetch the NEXT chunk's Gate/Beta ----------------------------- @@ -839,7 +865,11 @@ def gate_beta_warp( oob_neutral = cutlass.Float32(0.0) if cutlass.const_expr(cfg.log_gate) else cutlass.Float32(1.0) gate_vals[col] = gGate[pos] if pos_valid[col] else oob_neutral - if cutlass.const_expr(cfg.log_gate): + if cutlass.const_expr(cfg.safe_gate): + for col in cutlass.range_constexpr(n_cols): + contrib = a_l2 * softplus(gate_vals[col] + bias) + gate_vals[col] = contrib if pos_valid[col] else cutlass.Float32(0.0) + elif cutlass.const_expr(cfg.log_gate): for col in cutlass.range_constexpr(n_cols): gate_vals[col] = gate_vals[col] * cutlass.Float32(RCP_LN2) else: @@ -869,13 +899,25 @@ def gate_beta_warp( beta_idx = beta_index.idx beta_index = advance(beta_index, cfg.smem_beta_stages) - for col in cutlass.range_constexpr(n_cols): - pos = lidx + col * cfg.threads_per_warp - src = gBeta.iterator + gBeta.layout((pos,)) - dst = sBeta.iterator + sBeta.layout((pos, 0, beta_idx)) - cp_size = cutlass.Int32(4) * cutlass.Int32(pos_valid[col]) - nvvm.cp_async_shared_global(dst, src, 4, nvvm.LoadCacheModifier.CA, cp_size=cp_size) - nvvm.cp_async_mbarrier_arrive(bars.mb_beta_ready[beta_idx].smem_ptr, noinc=True) + if cutlass.const_expr(cfg.beta_sigmoid): + # io-dtype logits -> sigmoid (tanh identity) -> fp32 SMEM + for col in cutlass.range_constexpr(n_cols): + pos = lidx + col * cfg.threads_per_warp + beta_value = cutlass.Float32(0.0) + if pos_valid[col]: + beta_value = gBeta[pos].to(cutlass.Float32) + half = cutlass.Float32(0.5) + beta_value = (cute.math.tanh(beta_value * half, approx=True) * half + half).to(mBeta.element_type).to(cutlass.Float32) + sBeta[pos, 0, beta_idx] = beta_value + bars.mb_beta_ready[beta_idx].arrive() + else: + for col in cutlass.range_constexpr(n_cols): + pos = lidx + col * cfg.threads_per_warp + src = gBeta.iterator + gBeta.layout((pos,)) + dst = sBeta.iterator + sBeta.layout((pos, 0, beta_idx)) + cp_size = cutlass.Int32(4) * cutlass.Int32(pos_valid[col]) + nvvm.cp_async_shared_global(dst, src, 4, nvvm.LoadCacheModifier.CA, cp_size=cp_size) + nvvm.cp_async_mbarrier_arrive(bars.mb_beta_ready[beta_idx].smem_ptr, noinc=True) # ---- store-ready wait + in-place store back -------------------------- st_offset = batch_start + (cend - 1 - rev_idx) * cfg.b_t @@ -907,7 +949,10 @@ def gate_beta_warp( for col in cutlass.range_constexpr(n_cols): pos = lidx + col * cfg.threads_per_warp if cute.elem_less(st_offset + pos, write_end): - gBeta_st[pos] = sBeta[pos, 0, b_st_idx] + if cutlass.const_expr(cfg.beta_sigmoid): + gBeta_st[pos] = sBeta[pos, 0, b_st_idx].to(mDbeta.element_type) + else: + gBeta_st[pos] = sBeta[pos, 0, b_st_idx] tile_idx, sched_state = sched_next_tile(cfg, bars, sSched, sched_state, tile_idx, num_ctas) @@ -1536,11 +1581,11 @@ def mma_warp( bars.mb_dstate_acc_ready[dstate_idx].arrive(cta_group=1) # ---- dK attn += Q^T(S) @ dA ------------------------------------------ - bars.mb_da_ready[0].wait(da_ready_index.phase) - da_ready_index = advance(da_ready_index, 1) if have_dstate: bars.mb_dk_scale_acc_done[0].wait(dk_scale_index.phase) dk_scale_index = advance(dk_scale_index, 1) + bars.mb_da_ready[0].wait(da_ready_index.phase) + da_ready_index = advance(da_ready_index, 1) desc_q_mnmaj_dk_attn = d_q_trans0 desc_da_t = d_da_trans0 @@ -1986,7 +2031,7 @@ def compute0_warp_group( crow = warp_id * 16 + lane_id // 4 + ((k // 2) % 2) * 8 gBeta.append(sBeta[crow, 0, beta_idx]) - # ---- KK epi: M_kk[i,j] = W_kk[i,j] * T[i,j] * Beta[i] --------------- + # ---- KK epi: M_kk[i,j] = W_kk[i,j] * T_strict[i,j] * Beta[i] -------- tinv_idx = tinv_index.idx tinv_index = advance(tinv_index, cfg.smem_t_inv_stages) bars.mb_kk_acc_ready[0].wait(cg0_kk_ready.phase) @@ -1997,7 +2042,7 @@ def compute0_warp_group( kk_vec = nvvm.tcgen05_ld("16x256b", nvvm.make_tmem_ptr((tmem_warp_row << 16) + tmem_kk_col, cutlass.Float32), num=8) kk_pack = [] for k in cutlass.range_constexpr(num_vals // 2): - p0, p1 = fmul2(kk_vec[2 * k], kk_vec[2 * k + 1], decay_t[2 * k], decay_t[2 * k + 1]) + p0, p1 = fmul2(kk_vec[2 * k], kk_vec[2 * k + 1], decay_t_strict[2 * k], decay_t_strict[2 * k + 1]) v0, v1 = fmul2(p0, p1, gBeta[2 * k], gBeta[2 * k + 1]) kk_pack.append(fp32_to_fp16(v0, v1, dtype=cfg.io_dtype)) for c in cutlass.range_constexpr(ACC_N_FRAGS): @@ -2263,7 +2308,7 @@ def compute0_warp_group( nvvm.tcgen05_wait("load") bars.mb_dk_attn_acc_done[0].arrive() - # ---- dBeta/dGate M-terms: E = strict ⊙ dM_core ⊙ M_kk(sKK). ---------- + # ---- dBeta/dGate M-terms: E = dM_core ⊙ M_kk(sKK, strict-masked). ---- kk_frag = [] for c in cutlass.range_constexpr(ACC_N_FRAGS): kk_frag += list( @@ -2288,13 +2333,9 @@ def compute0_warp_group( binv_j = binv_row[j % 2] p_lo, p_hi = fmul2(dm_vec[2 * j], dm_vec[2 * j + 1], klo, khi) e_lo, e_hi = fmul2(p_lo, p_hi, binv_j, binv_j) - crow = warp_id * 16 + lane_id // 4 + (j % 2) * 8 - ccol = (lane_id % 4) * 2 + (j // 2) * 8 - e_val_lo = e_lo if crow > ccol else acc_zero - e_val_hi = e_hi if crow > ccol + 1 else acc_zero - row_acc[(j % 2) * 4 + (j // 2) % 4] += e_val_lo + e_val_hi + row_acc[(j % 2) * 4 + (j // 2) % 4] += e_lo + e_hi c0 = cutlass.const_expr((j // 2) * 2) - col_part[c0], col_part[c0 + 1] = fadd2(col_part[c0], col_part[c0 + 1], e_val_lo, e_val_hi) + col_part[c0], col_part[c0 + 1] = fadd2(col_part[c0], col_part[c0 + 1], e_lo, e_hi) row_part = [ (row_acc[0] + row_acc[1]) + (row_acc[2] + row_acc[3]), (row_acc[4] + row_acc[5]) + (row_acc[6] + row_acc[7]), @@ -2307,7 +2348,11 @@ def compute0_warp_group( if lane_id % 4 == 0: for rp in cutlass.range_constexpr(2): crow_r = warp_id * 16 + lane_id // 4 + rp * 8 - sBeta[crow_r, 0, beta_idx] = sBeta[crow_r, 0, beta_idx] - row_part[rp] * binv_row[rp] + db = sBeta[crow_r, 0, beta_idx] - row_part[rp] * binv_row[rp] + if cutlass.const_expr(cfg.beta_sigmoid): + b = gBeta[2 * rp] + db = db * (b - b * b) + sBeta[crow_r, 0, beta_idx] = db # ---- part reductions ------------------------------------------------- if chunk_idx + DSTATE_IN0 >= FIRST_STATE_CHUNK: @@ -3134,17 +3179,18 @@ def compute1_warp_group( dv_index = advance(dv_index, cfg.smem_dv_stages) -@cute.kernel -def build_all_descs_kernel( - base_q: cutlass.GridConstant[tma.TensorMap], - base_k: cutlass.GridConstant[tma.TensorMap], - base_v: cutlass.GridConstant[tma.TensorMap], - base_do: cutlass.GridConstant[tma.TensorMap], - base_checkpoint: cutlass.GridConstant[tma.TensorMap], - base_dq: cutlass.GridConstant[tma.TensorMap], - base_dk: cutlass.GridConstant[tma.TensorMap], - base_dv: cutlass.GridConstant[tma.TensorMap], - base_initial_state: cutlass.GridConstant[tma.TensorMap], +@cute.jit +def build_descs_body( + widx, + base_q, + base_k, + base_v, + base_do, + base_checkpoint, + base_dq, + base_dk, + base_dv, + base_initial_state, desc_ws: cute.Tensor, cu_seqlens: cute.Tensor, q: cute.Tensor, @@ -3167,9 +3213,9 @@ def build_all_descs_kernel( dk_rs: cutlass.Int32, dv_rs: cutlass.Int32, ) -> None: - """Single-launch builder for the per-BATCH descriptor arrays (one warp per array).""" - tidx, _, _ = cute.arch.thread_idx() - widx = cutlass.Int32(tidx) // cutlass.Int32(32) + """Per-batch descriptor-array build, one warp per array. Runs inside the + prologue kernel after its order pass; warps past the array count fall + through the widx guards.""" arr_words = n_batch * cutlass.Int32(TENSOR_MAP_QWORDS) sub0 = cute.make_tensor(desc_ws.iterator, cute.make_layout((arr_words,), stride=(1,))) sub1 = cute.make_tensor(desc_ws.iterator + arr_words, cute.make_layout((arr_words,), stride=(1,))) @@ -3220,10 +3266,118 @@ def build_all_descs_kernel( nvvm.fence_proxy_release(nvvm.MemScope.GPU, from_proxy=nvvm.Proxy.GENERIC, to_proxy=nvvm.Proxy.TENSORMAP) +@cute.kernel +def prologue_kernel( + run_order: cutlass.Constexpr[bool], + order_gen: cutlass.Constexpr[bool], + b_t: cutlass.Constexpr[int], + base_q: cutlass.GridConstant[tma.TensorMap], + base_k: cutlass.GridConstant[tma.TensorMap], + base_v: cutlass.GridConstant[tma.TensorMap], + base_do: cutlass.GridConstant[tma.TensorMap], + base_checkpoint: cutlass.GridConstant[tma.TensorMap], + base_dq: cutlass.GridConstant[tma.TensorMap], + base_dk: cutlass.GridConstant[tma.TensorMap], + base_dv: cutlass.GridConstant[tma.TensorMap], + base_initial_state: cutlass.GridConstant[tma.TensorMap], + desc_ws: cute.Tensor, + cu_seqlens: cute.Tensor, + q: cute.Tensor, + k: cute.Tensor, + v: cute.Tensor, + do_: cute.Tensor, + state_checkpoints: cute.Tensor, + dq: cute.Tensor, + dk: cute.Tensor, + dv: cute.Tensor, + state0: Optional[cute.Tensor], + mStaging: Optional[cute.Tensor], + mCount: cute.Tensor, + mWorkItems: cute.Tensor, + mSched: Optional[cute.Tensor], + n_batch: cutlass.Int32, + q_rs: cutlass.Int32, + k_rs: cutlass.Int32, + v_rs: cutlass.Int32, + do_rs: cutlass.Int32, + checkpoint_rs: cutlass.Int32, + checkpoint_every_n: cutlass.Int32, + dq_rs: cutlass.Int32, + dk_rs: cutlass.Int32, + dv_rs: cutlass.Int32, +) -> None: + """Single-CTA prologue. Under ``run_order`` this kernel is the first + work-item-table consumer, so it LPT-orders the table and zeroes both + consumers' sched rings via :func:`order_body`; it then builds the + per-batch TMA-descriptor arrays via :func:`build_descs_body`, one warp + per array (the extra warps only take part in the order phase).""" + tidx, _, _ = cute.arch.thread_idx() + tidx = cutlass.Int32(tidx) + widx = tidx // cutlass.Int32(32) + if cutlass.const_expr(run_order): + sKey = cutlass.Array(cutlass.Int32, ORDER_CAPACITY, space=cutlass.AddressSpace.smem, alignment=16) + sIdx = cutlass.Array(cutlass.Int32, ORDER_CAPACITY, space=cutlass.AddressSpace.smem, alignment=16) + sSpread = cutlass.Array(cutlass.Int32, 2, space=cutlass.AddressSpace.smem, alignment=8) + n_heads_out = cutlass.Int32(do_.shape[1]) + order_body( + order_gen, + True, + b_t, + ORDER_THREADS, + ORDER_ELEMS, + tidx, + n_heads_out, + n_heads_out * n_batch, + cu_seqlens, + mStaging, + mCount, + mWorkItems, + mSched, + sKey, + sIdx, + sSpread, + ) + build_descs_body( + widx, + base_q, + base_k, + base_v, + base_do, + base_checkpoint, + base_dq, + base_dk, + base_dv, + base_initial_state, + desc_ws, + cu_seqlens, + q, + k, + v, + do_, + state_checkpoints, + dq, + dk, + dv, + state0, + n_batch, + q_rs, + k_rs, + v_rs, + do_rs, + checkpoint_rs, + checkpoint_every_n, + dq_rs, + dk_rs, + dv_rs, + ) + + @cute.jit -def build_descs( +def prologue( io_dtype: cutlass.Constexpr, b_t: cutlass.Constexpr[int], + run_order: cutlass.Constexpr[bool], + order_gen: cutlass.Constexpr[bool], q: cute.Tensor, k: cute.Tensor, v: cute.Tensor, @@ -3234,11 +3388,17 @@ def build_descs( state_checkpoints: cute.Tensor, cu_seqlens: cute.Tensor, state0: Optional[cute.Tensor], + work_item_staging: Optional[cute.Tensor], + work_count: cute.Tensor, + work_items: cute.Tensor, + sched_all: Optional[cute.Tensor], tensormap_workspace: cute.Tensor, stream: cuda.CUstream, ): - """Build the per-(b,h) TMA-descriptor arrays (Q, K, V, dO, checkpoint loads; - dQ, dK, dV stores; the io-dtype initial-state loads when ``state0`` is given) into + """One-launch prologue: LPT-order the work items (with ``run_order``, when + this kernel is the backward pair's first table consumer) and build the + per-(b,h) TMA-descriptor arrays (Q, K, V, dO, checkpoint loads; dQ, dK, dV + stores; the io-dtype initial-state loads when ``state0`` is given) into ``tensormap_workspace``.""" h_q = q.shape[1] h_k = k.shape[1] @@ -3298,8 +3458,10 @@ def build_descs( ) base_desc_state0 = tma.create_tensor_map_tiled_from_view(initial_state_view, box_dims=(64, d_k_state, 1, 1), stride_order=(0, 1, 2, 3), swizzle=swz128) - n_warps = 9 if state0 is not None else 8 - build_all_descs_kernel( + prologue_kernel( + run_order, + order_gen, + b_t, base_desc_q, base_desc_k, base_desc_v, @@ -3320,6 +3482,10 @@ def build_descs( dk, dv, state0, + work_item_staging, + work_count, + work_items, + sched_all, cutlass.Int32(batch_size), cutlass.Int32(q_row_stride), cutlass.Int32(k_row_stride), @@ -3330,7 +3496,7 @@ def build_descs( cutlass.Int32(dq_row_stride), cutlass.Int32(dk_row_stride), cutlass.Int32(dv_row_stride), - ).launch(grid=(1, 1, 1), block=(32 * n_warps, 1, 1), stream=stream) + ).launch(grid=(1, 1, 1), block=(ORDER_THREADS, 1, 1), stream=stream) @cute.jit @@ -3340,6 +3506,8 @@ def host( k: cute.Tensor, v: cute.Tensor, gate: cute.Tensor, + a_log: Optional[cute.Tensor], + dt_bias: Optional[cute.Tensor], beta: cute.Tensor, dgate: cute.Tensor, dbeta: cute.Tensor, @@ -3407,6 +3575,8 @@ def host( kernel( cfg, gate, + a_log, + dt_bias, beta, dgate, dbeta, @@ -3441,6 +3611,8 @@ def host( def kernel( cfg: cutlass.Constexpr, mGate: cute.Tensor, + mA_log: Optional[cute.Tensor], + mDt_bias: Optional[cute.Tensor], mBeta: cute.Tensor, mDgate: cute.Tensor, mDbeta: cute.Tensor, @@ -4112,6 +4284,8 @@ def kernel( mWorkItems, tidx, mGate=mGate, + mA_log=mA_log, + mDt_bias=mDt_bias, mBeta=mBeta, mDgate=mDgate, mDbeta=mDbeta, @@ -4157,6 +4331,8 @@ class GdnBwdCfg: max_active_clusters: int is_GQA: bool log_gate: bool = False + safe_gate: bool = False + beta_sigmoid: bool = False # ---- fixed constants stamped from CFG by build_cfg --------------------------- b_t: int = CFG.B_T @@ -4246,6 +4422,8 @@ def build_cfg( use_dstate_in: bool = False, use_dstate0: bool = False, log_gate: bool = False, + safe_gate: bool = False, + beta_sigmoid: bool = False, dyn_sched: bool = False, ) -> GdnBwdCfg: """Build the per-compile ``GdnBwdCfg`` (io_dtype in {Float16, BFloat16}; @@ -4261,6 +4439,8 @@ def build_cfg( max_active_clusters=max_active_clusters, is_GQA=is_GQA, log_gate=log_gate, + safe_gate=safe_gate, + beta_sigmoid=beta_sigmoid, dyn_sched=dyn_sched, ) n_cg0 = len(cfg.compute_group_0_warp_ids) @@ -4296,7 +4476,11 @@ def get_compiled_cache( use_dstate_in: bool = False, use_dstate0: bool = False, log_gate: bool = False, + safe_gate: bool = False, + beta_sigmoid: bool = False, dyn_sched: bool = False, + run_order: bool = False, + order_gen: bool = False, ): """Return a mutable dict that lazily stores the compiled kernel.""" return {} @@ -4309,6 +4493,8 @@ def compile( use_dstate_in: bool = False, use_dstate0: bool = False, log_gate: bool = False, + safe_gate: bool = False, + beta_sigmoid: bool = False, dyn_sched: bool = False, *, num_sm: int, @@ -4319,6 +4505,8 @@ def compile( k_cute, v_cute, gate_cute, + a_log_cute=None, + dt_bias_cute=None, beta_cute, dgate_cute, dbeta_cute, @@ -4345,6 +4533,8 @@ def compile( use_dstate_in=use_dstate_in, use_dstate0=use_dstate0, log_gate=log_gate, + safe_gate=safe_gate, + beta_sigmoid=beta_sigmoid, dyn_sched=dyn_sched, ) cfg.h_q = h_q @@ -4358,6 +4548,8 @@ def compile( k_cute, v_cute, gate_cute, + a_log_cute, + dt_bias_cute, beta_cute, dgate_cute, dbeta_cute, @@ -4400,7 +4592,14 @@ def chunk_gdn_bwd_sm100( work_items=None, work_count=None, sched_ctr=None, + sched_all=None, + work_item_scratch=None, + order_in_prologue: bool = False, log_gate: bool = False, + safe_gate: bool = False, + a_log=None, + dt_bias=None, + use_beta_sigmoid: bool = False, workspace, stream, ) -> None: @@ -4421,19 +4620,28 @@ def chunk_gdn_bwd_sm100( k: ``(total_tokens, HK, DK)`` float16/bfloat16 v: ``(total_tokens, HV, DV)`` float16/bfloat16 gate: ``(total_tokens, HO)`` float32, forget gate — raw linear - alpha, or the natural-log decay when ``log_gate`` - beta: ``(total_tokens, HO)`` float32, update gate + alpha, or the natural-log decay when ``log_gate``, or raw logits + when ``safe_gate``, which applies the safe-gate transform + ``-exp(a_log) * softplus(gate + dt_bias)`` + beta: ``(total_tokens, HO)`` float32, update gate — post-sigmoid, or + io-dtype logits when ``use_beta_sigmoid`` do: ``(total_tokens, HO, DV)`` float16/bfloat16, output gradient state_checkpoints: ``(total_checkpoints, HO, DK, DV)`` io dtype, per-chunk forward states from the prefill kernel's checkpoint output (``checkpoint_every_n_tokens=B_T``) dq/dk/dv: pre-allocated output gradients, shaped/typed like q/k/v at HO heads - dgate/dbeta: pre-allocated ``(total_tokens, HO)`` float32 gate/beta - gradients + dgate: pre-allocated ``(total_tokens, HO)`` float32 gate gradient + (``safe_gate`` leaves it in the transformed gate space) + dbeta: pre-allocated ``(total_tokens, HO)`` beta gradient; float32, or + io dtype and wrt the raw logits under ``use_beta_sigmoid`` cu_seqlens: ``(num_seqs + 1,)`` int32 initial_state: ``(num_seqs, HO, DK, DV)`` io dtype (matching ``state_checkpoints``), or None scale: attention scale factor (must not be 0) + safe_gate: interpret ``gate`` through the safe-gate transform + a_log: ``(HO,)`` float32, safe-gate per-head log-amplitude (None = 0) + dt_bias: ``(HO,)`` float32, safe-gate per-head bias (None = 0) + use_beta_sigmoid: ``beta`` holds logits; sigmoid in-kernel work_items: ``(max_items, 8)`` int32 work-item table from ``common/split_k.py`` (REQUIRED; an uncut table row is the whole (b, h) sequence). Each item computes chunks ``[wstart, cend)`` @@ -4460,6 +4668,15 @@ def chunk_gdn_bwd_sm100( cu_stream = cuda.CUstream(int(stream)) dyn_sched = sched_ctr is not None + run_order = bool(order_in_prologue) + order_gen = run_order and work_item_scratch is None + if run_order and sched_all is None: + raise ValueError("order_in_prologue requires sched_all (the prologue zeroes both consumers' sched rings)") + if safe_gate and (a_log is None or dt_bias is None): + raise ValueError("safe_gate requires a_log and dt_bias") + if not safe_gate: + a_log = None + dt_bias = None cache = get_compiled_cache( str(q.dtype), str(cu_seqlens.dtype), @@ -4471,7 +4688,11 @@ def chunk_gdn_bwd_sm100( d_final_state is not None, d_initial_state is not None, log_gate, + safe_gate, + use_beta_sigmoid, dyn_sched, + run_order, + order_gen, ) if "compiled" not in cache: @@ -4490,6 +4711,8 @@ def chunk_gdn_bwd_sm100( sched_ctr_cute = None if dyn_sched: sched_ctr_cute = from_dlpack(sched_ctr, assumed_align=4).mark_layout_dynamic() + a_log_cute = from_dlpack(a_log, assumed_align=4) if a_log is not None else None + dt_bias_cute = from_dlpack(dt_bias, assumed_align=4) if dt_bias is not None else None cache["compiled"] = compile( io_dtype, is_GQA, @@ -4497,6 +4720,8 @@ def chunk_gdn_bwd_sm100( use_dstate_in=d_final_state is not None, use_dstate0=d_initial_state is not None, log_gate=log_gate, + safe_gate=safe_gate, + beta_sigmoid=use_beta_sigmoid, dyn_sched=dyn_sched, num_sm=multiprocessor_count(current_device_id()), h_q=HQ, @@ -4506,6 +4731,8 @@ def chunk_gdn_bwd_sm100( k_cute=from_dlpack(k, assumed_align=16).mark_layout_dynamic(leading_dim=2), v_cute=from_dlpack(v, assumed_align=16).mark_layout_dynamic(leading_dim=2), gate_cute=from_dlpack(gate, assumed_align=16).mark_layout_dynamic(leading_dim=1), + a_log_cute=a_log_cute, + dt_bias_cute=dt_bias_cute, beta_cute=from_dlpack(beta, assumed_align=16).mark_layout_dynamic(leading_dim=1), dgate_cute=from_dlpack(dgate, assumed_align=16).mark_layout_dynamic(leading_dim=1), dbeta_cute=from_dlpack(dbeta, assumed_align=16).mark_layout_dynamic(leading_dim=1), @@ -4526,21 +4753,29 @@ def chunk_gdn_bwd_sm100( compiled = cache["compiled"] - # desc build runs every execute by contract (cu contents are data; - # buffer pointers may change) — capture-safe, single tiny launch - if "build_descs" not in cache: - - checkpoints_bc = from_dlpack(state_checkpoints, assumed_align=16).mark_layout_dynamic(leading_dim=3) - cu_bc = from_dlpack(cu_seqlens, assumed_align=4).mark_layout_dynamic() - state0_bc = None + if "prologue" not in cache: + checkpoints_pl = from_dlpack(state_checkpoints, assumed_align=16).mark_layout_dynamic(leading_dim=3) + cu_pl = from_dlpack(cu_seqlens, assumed_align=4).mark_layout_dynamic() + state0_pl = None if initial_state is not None: - state0_bc = from_dlpack(initial_state, assumed_align=16).mark_layout_dynamic(leading_dim=3) - - ws_bc = from_dlpack(workspace, assumed_align=128).mark_layout_dynamic() - cache["build_descs"] = cute.compile( - build_descs, + state0_pl = from_dlpack(initial_state, assumed_align=16).mark_layout_dynamic(leading_dim=3) + staging_pl = None + if run_order and not order_gen: + staging_pl = from_dlpack(work_item_scratch, assumed_align=16) + staging_pl.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1), divisibility=1) + work_count_pl = from_dlpack(work_count, assumed_align=4).mark_layout_dynamic() + work_items_pl = from_dlpack(work_items, assumed_align=16) + work_items_pl.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1), divisibility=1) + sched_all_pl = None + if run_order: + sched_all_pl = from_dlpack(sched_all, assumed_align=4).mark_layout_dynamic() + ws_pl = from_dlpack(workspace, assumed_align=128).mark_layout_dynamic() + cache["prologue"] = cute.compile( + prologue, io_dtype, CFG.B_T, + run_order, + order_gen, from_dlpack(q, assumed_align=16).mark_layout_dynamic(leading_dim=2), from_dlpack(k, assumed_align=16).mark_layout_dynamic(leading_dim=2), from_dlpack(v, assumed_align=16).mark_layout_dynamic(leading_dim=2), @@ -4548,19 +4783,42 @@ def chunk_gdn_bwd_sm100( from_dlpack(dq, assumed_align=16).mark_layout_dynamic(leading_dim=2), from_dlpack(dk, assumed_align=16).mark_layout_dynamic(leading_dim=2), from_dlpack(dv, assumed_align=16).mark_layout_dynamic(leading_dim=2), - checkpoints_bc, - cu_bc, - state0_bc, - ws_bc, + checkpoints_pl, + cu_pl, + state0_pl, + staging_pl, + work_count_pl, + work_items_pl, + sched_all_pl, + ws_pl, cu_stream, options="--enable-tvm-ffi", ) - cache["build_descs"](q, k, v, do, dq, dk, dv, state_checkpoints, cu_seqlens, initial_state, workspace, cu_stream) + cache["prologue"]( + q, + k, + v, + do, + dq, + dk, + dv, + state_checkpoints, + cu_seqlens, + initial_state, + work_item_scratch if (run_order and not order_gen) else None, + work_count, + work_items, + sched_all if run_order else None, + workspace, + cu_stream, + ) compiled( q, k, v, gate, + a_log, + dt_bias, beta, dgate, dbeta, @@ -4578,3 +4836,81 @@ def chunk_gdn_bwd_sm100( workspace, cu_stream, ) + return cache + + +def run_bwd( + cache, + q, + k, + v, + gate, + beta, + do, + state_checkpoints, + dq, + dk, + dv, + dgate, + dbeta, + cu_seqlens, + initial_state, + d_initial_state, + d_final_state, + work_items, + work_count, + sched_ctr, + sched_all, + work_item_scratch, + tensormap_workspace, + scale, + stream, + a_log=None, + dt_bias=None, +) -> None: + """Replay the compiled plan: the prologue launch, then the main launch. + The caller owns the contract, which the plan validated at build, so + nothing here raises.""" + cu_stream = cuda.CUstream(int(stream)) + cache["prologue"]( + q, + k, + v, + do, + dq, + dk, + dv, + state_checkpoints, + cu_seqlens, + initial_state, + work_item_scratch, + work_count, + work_items, + sched_all, + tensormap_workspace, + cu_stream, + ) + cache["compiled"]( + q, + k, + v, + gate, + a_log, + dt_bias, + beta, + dgate, + dbeta, + do, + dq, + dk, + dv, + cu_seqlens, + d_initial_state, + d_final_state, + work_items, + work_count, + sched_ctr, + scale, + tensormap_workspace, + cu_stream, + ) diff --git a/python/cudnn/linear_attention/frost/kernel/gdn_prefill_f16.py b/python/cudnn/linear_attention/frost/kernel/gdn_prefill_f16.py index cec6bd144..82fa23546 100644 --- a/python/cudnn/linear_attention/frost/kernel/gdn_prefill_f16.py +++ b/python/cudnn/linear_attention/frost/kernel/gdn_prefill_f16.py @@ -95,7 +95,8 @@ from cutlass.cutlass_dsl import min from ..common.thd import emit_checkpoint_seq_descs, emit_seq_descs, TENSOR_MAP_QWORDS -from ..common.split_k import decode_work_item +from ..common.split_k import ORDER_CAPACITY, ORDER_ELEMS, ORDER_THREADS, decode_work_item, order_body +from ..common.elementwise import softplus from ..common.host import get_dtype from cudnn.frost.buffers import current_device_id, data_ptr from cudnn.frost.device import multiprocessor_count @@ -626,6 +627,8 @@ def gate_beta_warp( mWorkItems, tidx, mGate, + mA_log, + mDt_bias, mBeta, sCumsumlog, sCumprod, @@ -641,12 +644,19 @@ def gate_beta_warp( beta_index = PipelineState.start(phase=1) lidx = tidx % cfg.threads_per_warp + a_l2 = cutlass.Float32(0.0) + bias = cutlass.Float32(0.0) sched_state = PipelineState.start(phase=0) tile_idx = cutlass.Int32(bidx) while tile_idx < total_tiles: batch_idx, head_idx, batch_start, batch_end, seqlen_b, num_chunks_b, wstart, wend, cstart, cend = decode_work_item(cfg, tile_idx, mWorkItems) n_local = wend - cstart n_padded = ((n_local + 1) // 2) * 2 + if cutlass.const_expr(cfg.safe_gate): + if n_local > 0: + # per-head transform constants, fixed for the whole tile + a_l2 = -cute.math.exp2(mA_log[head_idx].to(cutlass.Float32) * cutlass.Float32(RCP_LN2), fastmath=True) * cutlass.Float32(RCP_LN2) + bias = mDt_bias[head_idx].to(cutlass.Float32) if n_local > 0: for local_idx in cutlass.range(n_padded): # ---- Gate load: GMEM -> SMEM (OOB neutral) ----------------------- @@ -669,7 +679,12 @@ def gate_beta_warp( tok_clamped = min(tok, batch_end - 1) gate_vals[col] = gGateSeq[tok_clamped] if pos_valid[col] else oob_neutral - if cutlass.const_expr(cfg.log_gate): + if cutlass.const_expr(cfg.safe_gate): + # raw logits -> log2-domain decay: a_l2 * softplus(g + bias) (split-K scan arithmetic) + for col in cutlass.range_constexpr(n_cols): + contrib = a_l2 * softplus(gate_vals[col] + bias) + gate_vals[col] = contrib if pos_valid[col] else cutlass.Float32(0.0) + elif cutlass.const_expr(cfg.log_gate): for col in cutlass.range_constexpr(n_cols): gate_vals[col] = gate_vals[col] * cutlass.Float32(RCP_LN2) else: @@ -701,13 +716,25 @@ def gate_beta_warp( beta_idx = beta_index.idx bars.mb_beta_done[beta_idx].wait(beta_index.phase) beta_index = advance(beta_index, cfg.smem_beta_stages) - for col in cutlass.range_constexpr(n_cols): - pos = lidx + col * cfg.threads_per_warp - src = gBeta.iterator + gBeta.layout((pos,)) - dst = sBeta.iterator + sBeta.layout((pos, 0, beta_idx)) - cp_size = cutlass.Int32(4) * cutlass.Int32(pos_valid[col]) - nvvm.cp_async_shared_global(dst, src, 4, nvvm.LoadCacheModifier.CA, cp_size=cp_size) - nvvm.cp_async_mbarrier_arrive(bars.mb_beta_ready[beta_idx].smem_ptr, noinc=True) + if cutlass.const_expr(cfg.beta_sigmoid): + # io-dtype logits -> sigmoid (tanh identity) -> fp32 SMEM + for col in cutlass.range_constexpr(n_cols): + pos = lidx + col * cfg.threads_per_warp + beta_value = cutlass.Float32(0.0) + if pos_valid[col]: + beta_value = gBeta[pos].to(cutlass.Float32) + half = cutlass.Float32(0.5) + beta_value = (cute.math.tanh(beta_value * half, approx=True) * half + half).to(mBeta.element_type).to(cutlass.Float32) + sBeta[pos, 0, beta_idx] = beta_value + bars.mb_beta_ready[beta_idx].arrive() + else: + for col in cutlass.range_constexpr(n_cols): + pos = lidx + col * cfg.threads_per_warp + src = gBeta.iterator + gBeta.layout((pos,)) + dst = sBeta.iterator + sBeta.layout((pos, 0, beta_idx)) + cp_size = cutlass.Int32(4) * cutlass.Int32(pos_valid[col]) + nvvm.cp_async_shared_global(dst, src, 4, nvvm.LoadCacheModifier.CA, cp_size=cp_size) + nvvm.cp_async_mbarrier_arrive(bars.mb_beta_ready[beta_idx].smem_ptr, noinc=True) tile_idx, sched_state = sched_next_tile(cfg, bars, sSched, sched_state, tile_idx, num_ctas) for _ in range(cfg.smem_gate_stages): @@ -1925,13 +1952,14 @@ def compute1_warp_group( checkpoint_cnt = checkpoint_cnt + 1 -@cute.kernel -def build_all_descs_kernel( - base_q: cutlass.GridConstant[tma.TensorMap], - base_k: cutlass.GridConstant[tma.TensorMap], - base_v: cutlass.GridConstant[tma.TensorMap], - base_o: cutlass.GridConstant[tma.TensorMap], - base_checkpoint: cutlass.GridConstant[tma.TensorMap], +@cute.jit +def build_descs_body( + widx, + base_q, + base_k, + base_v, + base_o, + base_checkpoint, desc_ws: cute.Tensor, cu_seqlens: cute.Tensor, q: cute.Tensor, @@ -1947,10 +1975,9 @@ def build_all_descs_kernel( checkpoint_row_stride: cutlass.Int32, checkpoint_every_n: cutlass.Int32, ) -> None: - """Single-launch builder for the per-batch descriptor arrays (one warp - per array).""" - tidx, _, _ = cute.arch.thread_idx() - widx = cutlass.Int32(tidx) // cutlass.Int32(32) + """Per-batch descriptor-array build, one warp per array. Runs inside the + prologue kernel after its order pass; warps past the array count fall + through the widx guards.""" arr_words = n_batch * cutlass.Int32(TENSOR_MAP_QWORDS) sub0 = cute.make_tensor(desc_ws.iterator, cute.make_layout((arr_words,), stride=(1,))) sub1 = cute.make_tensor(desc_ws.iterator + arr_words, cute.make_layout((arr_words,), stride=(1,))) @@ -1982,21 +2009,110 @@ def build_all_descs_kernel( nvvm.fence_proxy_release(nvvm.MemScope.GPU, from_proxy=nvvm.Proxy.GENERIC, to_proxy=nvvm.Proxy.TENSORMAP) +@cute.kernel +def prologue_kernel( + order_gen: cutlass.Constexpr[bool], + has_sched: cutlass.Constexpr[bool], + b_t: cutlass.Constexpr[int], + base_q: cutlass.GridConstant[tma.TensorMap], + base_k: cutlass.GridConstant[tma.TensorMap], + base_v: cutlass.GridConstant[tma.TensorMap], + base_o: cutlass.GridConstant[tma.TensorMap], + base_checkpoint: cutlass.GridConstant[tma.TensorMap], + desc_ws: cute.Tensor, + cu_seqlens: cute.Tensor, + q: cute.Tensor, + k: cute.Tensor, + v: cute.Tensor, + o: Optional[cute.Tensor], + state_checkpoints_out: Optional[cute.Tensor], + mStaging: Optional[cute.Tensor], + mCount: cute.Tensor, + mWorkItems: cute.Tensor, + mSched: Optional[cute.Tensor], + n_batch: cutlass.Int32, + q_row_stride: cutlass.Int32, + k_row_stride: cutlass.Int32, + v_row_stride: cutlass.Int32, + o_row_stride: cutlass.Int32, + checkpoint_row_stride: cutlass.Int32, + checkpoint_every_n: cutlass.Int32, +) -> None: + """Single-CTA prologue: LPT-order the work-item table and zero the sched + rings via :func:`order_body`, then build the per-batch TMA-descriptor + arrays via :func:`build_descs_body`, one warp per array (the extra warps + only take part in the order phase).""" + tidx, _, _ = cute.arch.thread_idx() + tidx = cutlass.Int32(tidx) + widx = tidx // cutlass.Int32(32) + sKey = cutlass.Array(cutlass.Int32, ORDER_CAPACITY, space=cutlass.AddressSpace.smem, alignment=16) + sIdx = cutlass.Array(cutlass.Int32, ORDER_CAPACITY, space=cutlass.AddressSpace.smem, alignment=16) + sSpread = cutlass.Array(cutlass.Int32, 2, space=cutlass.AddressSpace.smem, alignment=8) + n_heads_out = cutlass.Int32(q.shape[1] if q.shape[1] >= v.shape[1] else v.shape[1]) + order_body( + order_gen, + has_sched, + b_t, + ORDER_THREADS, + ORDER_ELEMS, + tidx, + n_heads_out, + n_heads_out * n_batch, + cu_seqlens, + mStaging, + mCount, + mWorkItems, + mSched, + sKey, + sIdx, + sSpread, + ) + build_descs_body( + widx, + base_q, + base_k, + base_v, + base_o, + base_checkpoint, + desc_ws, + cu_seqlens, + q, + k, + v, + o, + state_checkpoints_out, + n_batch, + q_row_stride, + k_row_stride, + v_row_stride, + o_row_stride, + checkpoint_row_stride, + checkpoint_every_n, + ) + + @cute.jit -def build_descs( +def prologue( io_dtype: cutlass.Constexpr, b_t: cutlass.Constexpr[int], + order_gen: cutlass.Constexpr[bool], + has_sched: cutlass.Constexpr[bool], q: cute.Tensor, k: cute.Tensor, v: cute.Tensor, o: Optional[cute.Tensor], cu_seqlens: cute.Tensor, state_checkpoints_out: Optional[cute.Tensor], + work_item_staging: Optional[cute.Tensor], + work_count: cute.Tensor, + work_items: cute.Tensor, + sched_ctr: Optional[cute.Tensor], checkpoint_every_n: cutlass.Int32, tensormap_workspace: cute.Tensor, stream: cuda.CUstream, ): - """Build the 5 per-(b,h) TMA-descriptor arrays (Q, K, V, O, checkpoints) into + """One-launch prologue: LPT-order the work items and build the 5 + per-(b,h) TMA-descriptor arrays (Q, K, V, O, checkpoints) into ``tensormap_workspace``.""" h_q = q.shape[1] h_k = k.shape[1] @@ -2043,8 +2159,10 @@ def build_descs( checkpoint_view, box_dims=(checkpoint_elems_per_128b, d_k_state, 1, 1), stride_order=(0, 1, 2, 3), swizzle=swz128 ) - n_warps = 5 if state_checkpoints_out is not None else (4 if o is not None else 3) - build_all_descs_kernel( + prologue_kernel( + order_gen, + has_sched, + b_t, base_desc_q, base_desc_k, base_desc_v, @@ -2057,6 +2175,10 @@ def build_descs( v, o, state_checkpoints_out, + work_item_staging, + work_count, + work_items, + sched_ctr, cutlass.Int32(batch_size), cutlass.Int32(q_row_stride), cutlass.Int32(k_row_stride), @@ -2064,7 +2186,7 @@ def build_descs( cutlass.Int32(o.stride[0] if o is not None else 0), cutlass.Int32(state_checkpoints_out.stride[0] if state_checkpoints_out is not None else 0), checkpoint_every_n, - ).launch(grid=(1, 1, 1), block=(32 * n_warps, 1, 1), stream=stream) + ).launch(grid=(1, 1, 1), block=(ORDER_THREADS, 1, 1), stream=stream) @cute.jit @@ -2074,6 +2196,8 @@ def host( k: cute.Tensor, v: cute.Tensor, gate: cute.Tensor, + a_log: Optional[cute.Tensor], + dt_bias: Optional[cute.Tensor], beta: cute.Tensor, o: Optional[cute.Tensor], cu_seqlens: cute.Tensor, @@ -2226,6 +2350,8 @@ def host( kernel( cfg, gate, + a_log, + dt_bias, beta, cu_seqlens, state_in, @@ -2256,6 +2382,8 @@ def host( def kernel( cfg: cutlass.Constexpr, mGate: cute.Tensor, + mA_log: Optional[cute.Tensor], + mDt_bias: Optional[cute.Tensor], mBeta: cute.Tensor, cu_seqlens: cute.Tensor, mState_init: Optional[cute.Tensor], @@ -2523,6 +2651,8 @@ def kernel( mWorkItems, tidx=tidx, mGate=mGate, + mA_log=mA_log, + mDt_bias=mDt_bias, mBeta=mBeta, sCumsumlog=sCumsumlog, sCumprod=sCumprod, @@ -2605,6 +2735,8 @@ class GdnCfg: store_final_state: bool enable_checkpoints: bool log_gate: bool = False + safe_gate: bool = False + beta_sigmoid: bool = False dyn_sched: bool = False sched_stages: int = CFG.SMEM_SCHED_STAGES @@ -2681,6 +2813,8 @@ def build_cfg( store_final_state: bool = True, enable_checkpoints: bool = False, log_gate: bool = False, + safe_gate: bool = False, + beta_sigmoid: bool = False, dyn_sched: bool = False, ) -> GdnCfg: """Build the per-compile ``GdnCfg`` (io_dtype ∈ {Float16, BFloat16}; @@ -2697,6 +2831,8 @@ def build_cfg( store_final_state=store_final_state, enable_checkpoints=enable_checkpoints, log_gate=log_gate, + safe_gate=safe_gate, + beta_sigmoid=beta_sigmoid, dyn_sched=dyn_sched, ) cfg.smem_checkpoint_stages = 1 @@ -2740,7 +2876,10 @@ def get_compiled_cache( store_final_state: bool, enable_checkpoints: bool, log_gate: bool, + safe_gate: bool, + beta_sigmoid: bool, dyn_sched: bool, + order_gen: bool, ): """Return a mutable dict that lazily stores the compiled kernel.""" return {} @@ -2754,6 +2893,8 @@ def compile( store_final_state: bool, enable_checkpoints: bool, log_gate: bool = False, + safe_gate: bool = False, + beta_sigmoid: bool = False, dyn_sched: bool = False, *, num_sm: int, @@ -2761,6 +2902,8 @@ def compile( k_cute, v_cute, gate_cute, + a_log_cute=None, + dt_bias_cute=None, beta_cute, o_cute, cu_seqlens_cute, @@ -2784,6 +2927,8 @@ def compile( store_final_state=store_final_state, enable_checkpoints=enable_checkpoints, log_gate=log_gate, + safe_gate=safe_gate, + beta_sigmoid=beta_sigmoid, dyn_sched=dyn_sched, ) @@ -2794,6 +2939,8 @@ def compile( k_cute, v_cute, gate_cute, + a_log_cute, + dt_bias_cute, beta_cute, o_cute, cu_seqlens_cute, @@ -2827,6 +2974,11 @@ def chunk_gdn_sm100( work_count=None, sched_ctr=None, log_gate: bool = False, + safe_gate: bool = False, + a_log=None, + dt_bias=None, + use_beta_sigmoid: bool = False, + work_item_scratch=None, *, workspace, stream, @@ -2842,8 +2994,11 @@ def chunk_gdn_sm100( k: ``(total_tokens, HK, DK)`` float16/bfloat16 v: ``(total_tokens, HV, DK)`` float16/bfloat16 gate: ``(total_tokens, HO)`` float32, forget gate — raw linear - alpha, or the natural-log decay when ``log_gate`` - beta: ``(total_tokens, HO)`` float32, update gate + alpha, or the natural-log decay when ``log_gate``, or raw logits + when ``safe_gate``, which applies the safe-gate transform + ``-exp(a_log) * softplus(gate + dt_bias)`` + beta: ``(total_tokens, HO)`` float32, update gate — post-sigmoid, or + io-dtype logits when ``use_beta_sigmoid`` output: ``(total_tokens, HO, DK)`` float16/bfloat16, pre-allocated cu_seqlens: ``(num_seqs + 1,)`` int32 initial_state: ``(num_seqs, HO, DK, DK)`` float32/bfloat16, or None @@ -2863,6 +3018,10 @@ def chunk_gdn_sm100( log_gate: ``gate`` holds natural-log decay values; the gate warp skips its log2 (rescales by 1/ln2) instead of exponentiating upstream + safe_gate: interpret ``gate`` through the safe-gate transform + a_log: ``(HO,)`` float32, safe-gate per-head log-amplitude (None = 0) + dt_bias: ``(HO,)`` float32, safe-gate per-head bias (None = 0) + use_beta_sigmoid: ``beta`` holds logits; sigmoid in-kernel workspace: ``(>= tensormap_workspace_bytes(module, B) // 8,)`` int64, 128-byte aligned; holds the per-(b,h) TMA descriptors (contents managed here — reuse the same buffer across calls) @@ -2878,7 +3037,13 @@ def chunk_gdn_sm100( enable_checkpoints = checkpoint_every_n_tokens > 0 if work_items is None or work_count is None: raise ValueError("work_items/work_count are required (the split-table stage builds them for every launch)") + if safe_gate and (a_log is None or dt_bias is None): + raise ValueError("safe_gate requires a_log and dt_bias") + if not safe_gate: + a_log = None + dt_bias = None dyn_sched = sched_ctr is not None + order_gen = work_item_scratch is None io_dtype = get_dtype(q.dtype) if initial_state is not None: @@ -2903,7 +3068,10 @@ def chunk_gdn_sm100( store_final_state, enable_checkpoints, log_gate, + safe_gate, + use_beta_sigmoid, dyn_sched, + order_gen, ) if "compiled" not in cache: @@ -2915,6 +3083,8 @@ def chunk_gdn_sm100( v_cute.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1, 2), divisibility=1) gate_cute = from_dlpack(gate, assumed_align=16) gate_cute.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1), divisibility=1) + a_log_cute = from_dlpack(a_log, assumed_align=4) if a_log is not None else None + dt_bias_cute = from_dlpack(dt_bias, assumed_align=4) if dt_bias is not None else None beta_cute = from_dlpack(beta, assumed_align=16) beta_cute.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1), divisibility=1) o_cute = from_dlpack(output, assumed_align=16) @@ -2949,12 +3119,16 @@ def chunk_gdn_sm100( store_final_state, enable_checkpoints, log_gate, + safe_gate, + use_beta_sigmoid, dyn_sched, num_sm=multiprocessor_count(current_device_id()), q_cute=q_cute, k_cute=k_cute, v_cute=v_cute, gate_cute=gate_cute, + a_log_cute=a_log_cute, + dt_bias_cute=dt_bias_cute, beta_cute=beta_cute, o_cute=o_cute, cu_seqlens_cute=cu_seqlens_cute, @@ -2971,47 +3145,64 @@ def chunk_gdn_sm100( compiled = cache["compiled"] - # desc build runs every execute by contract (cu contents are data; - # buffer pointers may change) — capture-safe, single tiny launch - if "build_descs" not in cache: - q_bc = from_dlpack(q, assumed_align=16) - q_bc.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1, 2), divisibility=1) - k_bc = from_dlpack(k, assumed_align=16) - k_bc.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1, 2), divisibility=1) - v_bc = from_dlpack(v, assumed_align=16) - v_bc.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1, 2), divisibility=1) - o_bc = from_dlpack(output, assumed_align=16) - o_bc.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1, 2), divisibility=1) - cu_bc = from_dlpack(cu_seqlens, assumed_align=8 if str(cu_seqlens.dtype).endswith("int64") else 4).mark_layout_dynamic() - checkpoints_bc = None - cu_ckpt_bc = None + if "prologue" not in cache: + q_pl = from_dlpack(q, assumed_align=16) + q_pl.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1, 2), divisibility=1) + k_pl = from_dlpack(k, assumed_align=16) + k_pl.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1, 2), divisibility=1) + v_pl = from_dlpack(v, assumed_align=16) + v_pl.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1, 2), divisibility=1) + o_pl = from_dlpack(output, assumed_align=16) + o_pl.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1, 2), divisibility=1) + cu_pl = from_dlpack(cu_seqlens, assumed_align=8 if str(cu_seqlens.dtype).endswith("int64") else 4).mark_layout_dynamic() + checkpoints_pl = None if enable_checkpoints: - checkpoints_bc = from_dlpack(output_state_checkpoints, assumed_align=16) - checkpoints_bc.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1, 2, 3), divisibility=1) - ws_bc = from_dlpack(workspace, assumed_align=128).mark_layout_dynamic() - cache["build_descs"] = cute.compile( - build_descs, + checkpoints_pl = from_dlpack(output_state_checkpoints, assumed_align=16) + checkpoints_pl.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1, 2, 3), divisibility=1) + staging_pl = None + if not order_gen: + staging_pl = from_dlpack(work_item_scratch, assumed_align=16) + staging_pl.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1), divisibility=1) + work_items_pl = from_dlpack(work_items, assumed_align=16) + work_items_pl.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1), divisibility=1) + work_count_pl = from_dlpack(work_count, assumed_align=4).mark_layout_dynamic() + sched_pl = None + if dyn_sched: + sched_pl = from_dlpack(sched_ctr, assumed_align=4).mark_layout_dynamic() + ws_pl = from_dlpack(workspace, assumed_align=128).mark_layout_dynamic() + cache["prologue"] = cute.compile( + prologue, io_dtype, CFG.B_T, - q_bc, - k_bc, - v_bc, - o_bc, - cu_bc, - checkpoints_bc, - cutlass.Int32(checkpoint_every_n_tokens if enable_checkpoints else 1), - ws_bc, + order_gen, + dyn_sched, + q_pl, + k_pl, + v_pl, + o_pl, + cu_pl, + checkpoints_pl, + staging_pl, + work_count_pl, + work_items_pl, + sched_pl, + cutlass.Int32(checkpoint_every_n_tokens), + ws_pl, cu_stream, options="--enable-tvm-ffi", ) - cache["build_descs"]( + cache["prologue"]( q, k, v, output, cu_seqlens, - output_state_checkpoints, - checkpoint_every_n_tokens if enable_checkpoints else 1, + output_state_checkpoints if enable_checkpoints else None, + work_item_scratch if not order_gen else None, + work_count, + work_items, + sched_ctr, + checkpoint_every_n_tokens, workspace, cu_stream, ) @@ -3020,6 +3211,8 @@ def chunk_gdn_sm100( k, v, gate, + a_log, + dt_bias, beta, output, cu_seqlens, @@ -3033,3 +3226,68 @@ def chunk_gdn_sm100( workspace, cu_stream, ) + return cache + + +def run_prefill( + cache, + q, + k, + v, + gate, + beta, + output, + cu_seqlens, + initial_state, + output_state, + output_state_checkpoints, + work_items, + work_count, + sched_ctr, + work_item_scratch, + tensormap_workspace, + checkpoint_every_n_tokens, + scale, + stream, + a_log=None, + dt_bias=None, +) -> None: + """Replay the compiled plan: the prologue launch, then the main launch. + The caller owns the contract, which the plan validated at build, so + nothing here raises.""" + cu_stream = cuda.CUstream(int(stream)) + cache["prologue"]( + q, + k, + v, + output, + cu_seqlens, + output_state_checkpoints, + work_item_scratch, + work_count, + work_items, + sched_ctr, + checkpoint_every_n_tokens, + tensormap_workspace, + cu_stream, + ) + cache["compiled"]( + q, + k, + v, + gate, + a_log, + dt_bias, + beta, + output, + cu_seqlens, + initial_state, + output_state, + work_items, + work_count, + sched_ctr, + checkpoint_every_n_tokens, + scale, + tensormap_workspace, + cu_stream, + ) diff --git a/python/cudnn/linear_attention/frost/kernel/gdn_recompute_f16.py b/python/cudnn/linear_attention/frost/kernel/gdn_recompute_f16.py index 4ff922c18..eb9067543 100644 --- a/python/cudnn/linear_attention/frost/kernel/gdn_recompute_f16.py +++ b/python/cudnn/linear_attention/frost/kernel/gdn_recompute_f16.py @@ -85,7 +85,8 @@ from cutlass.cutlass_dsl import min from ..common.thd import emit_checkpoint_seq_descs, emit_seq_descs, TENSOR_MAP_QWORDS -from ..common.split_k import decode_work_item +from ..common.split_k import ORDER_CAPACITY, ORDER_ELEMS, ORDER_THREADS, decode_work_item, order_body +from ..common.elementwise import softplus from ..common.host import get_dtype from cudnn.frost.buffers import current_device_id, data_ptr from cudnn.frost.device import multiprocessor_count @@ -529,6 +530,8 @@ def gate_beta_warp( mWorkItems, tidx, mGate, + mA_log, + mDt_bias, mBeta, sCumsumlog, sCumprod, @@ -544,12 +547,19 @@ def gate_beta_warp( beta_index = PipelineState.start(phase=1) lidx = tidx % cfg.threads_per_warp + a_l2 = cutlass.Float32(0.0) + bias = cutlass.Float32(0.0) sched_state = PipelineState.start(phase=0) tile_idx = cutlass.Int32(bidx) while tile_idx < total_tiles: batch_idx, head_idx, batch_start, batch_end, seqlen_b, num_chunks_b, wstart, wend, cstart, cend = decode_work_item(cfg, tile_idx, mWorkItems) n_local = wend - cstart n_padded = ((n_local + 1) // 2) * 2 + if cutlass.const_expr(cfg.safe_gate): + if n_local > 0: + # per-head transform constants, fixed for the whole tile + a_l2 = -cute.math.exp2(mA_log[head_idx].to(cutlass.Float32) * cutlass.Float32(RCP_LN2), fastmath=True) * cutlass.Float32(RCP_LN2) + bias = mDt_bias[head_idx].to(cutlass.Float32) if n_local > 0: for local_idx in cutlass.range(n_padded): # ---- Gate load: GMEM -> SMEM (OOB neutral) ----------------------- @@ -572,7 +582,12 @@ def gate_beta_warp( tok_clamped = min(tok, batch_end - 1) gate_vals[col] = gGateSeq[tok_clamped] if pos_valid[col] else oob_neutral - if cutlass.const_expr(cfg.log_gate): + if cutlass.const_expr(cfg.safe_gate): + # raw logits -> log2-domain decay: a_l2 * softplus(g + bias) (split-K scan arithmetic) + for col in cutlass.range_constexpr(n_cols): + contrib = a_l2 * softplus(gate_vals[col] + bias) + gate_vals[col] = contrib if pos_valid[col] else cutlass.Float32(0.0) + elif cutlass.const_expr(cfg.log_gate): for col in cutlass.range_constexpr(n_cols): gate_vals[col] = gate_vals[col] * cutlass.Float32(RCP_LN2) else: @@ -605,13 +620,25 @@ def gate_beta_warp( beta_idx = beta_index.idx bars.mb_beta_done[beta_idx].wait(beta_index.phase) beta_index = advance(beta_index, cfg.smem_beta_stages) - for col in cutlass.range_constexpr(n_cols): - pos = lidx + col * cfg.threads_per_warp - src = gBeta.iterator + gBeta.layout((pos,)) - dst = sBeta.iterator + sBeta.layout((pos, 0, beta_idx)) - cp_size = cutlass.Int32(4) * cutlass.Int32(pos_valid[col]) - nvvm.cp_async_shared_global(dst, src, 4, nvvm.LoadCacheModifier.CA, cp_size=cp_size) - nvvm.cp_async_mbarrier_arrive(bars.mb_beta_ready[beta_idx].smem_ptr, noinc=True) + if cutlass.const_expr(cfg.beta_sigmoid): + # io-dtype logits -> sigmoid (tanh identity) -> fp32 SMEM + for col in cutlass.range_constexpr(n_cols): + pos = lidx + col * cfg.threads_per_warp + beta_value = cutlass.Float32(0.0) + if pos_valid[col]: + beta_value = gBeta[pos].to(cutlass.Float32) + half = cutlass.Float32(0.5) + beta_value = (cute.math.tanh(beta_value * half, approx=True) * half + half).to(mBeta.element_type).to(cutlass.Float32) + sBeta[pos, 0, beta_idx] = beta_value + bars.mb_beta_ready[beta_idx].arrive() + else: + for col in cutlass.range_constexpr(n_cols): + pos = lidx + col * cfg.threads_per_warp + src = gBeta.iterator + gBeta.layout((pos,)) + dst = sBeta.iterator + sBeta.layout((pos, 0, beta_idx)) + cp_size = cutlass.Int32(4) * cutlass.Int32(pos_valid[col]) + nvvm.cp_async_shared_global(dst, src, 4, nvvm.LoadCacheModifier.CA, cp_size=cp_size) + nvvm.cp_async_mbarrier_arrive(bars.mb_beta_ready[beta_idx].smem_ptr, noinc=True) tile_idx, sched_state = sched_next_tile(cfg, bars, sSched, sched_state, tile_idx, num_ctas) for _ in range(cfg.smem_gate_stages): @@ -1648,11 +1675,12 @@ def compute1_warp_group( checkpoint_cnt = checkpoint_cnt + 1 -@cute.kernel -def build_all_descs_kernel( - base_k: cutlass.GridConstant[tma.TensorMap], - base_v: cutlass.GridConstant[tma.TensorMap], - base_checkpoint: cutlass.GridConstant[tma.TensorMap], +@cute.jit +def build_descs_body( + widx, + base_k, + base_v, + base_checkpoint, desc_ws: cute.Tensor, cu_seqlens: cute.Tensor, k: cute.Tensor, @@ -1664,10 +1692,9 @@ def build_all_descs_kernel( checkpoint_row_stride: cutlass.Int32, checkpoint_every_n: cutlass.Int32, ) -> None: - """Single-launch builder for the per-BATCH descriptor arrays (one warp - per array).""" - tidx, _, _ = cute.arch.thread_idx() - widx = cutlass.Int32(tidx) // cutlass.Int32(32) + """Per-batch descriptor-array build, one warp per array. Runs inside the + prologue kernel after its order pass; warps past the array count fall + through the widx guards.""" arr_words = n_batch * cutlass.Int32(TENSOR_MAP_QWORDS) desc_k_arr = cute.make_tensor(desc_ws.iterator, cute.make_layout((arr_words,), stride=(1,))) desc_v_arr = cute.make_tensor(desc_ws.iterator + arr_words, cute.make_layout((arr_words,), stride=(1,))) @@ -1690,19 +1717,101 @@ def build_all_descs_kernel( nvvm.fence_proxy_release(nvvm.MemScope.GPU, from_proxy=nvvm.Proxy.GENERIC, to_proxy=nvvm.Proxy.TENSORMAP) +@cute.kernel +def prologue_kernel( + run_order: cutlass.Constexpr[bool], + order_gen: cutlass.Constexpr[bool], + b_t: cutlass.Constexpr[int], + base_k: cutlass.GridConstant[tma.TensorMap], + base_v: cutlass.GridConstant[tma.TensorMap], + base_checkpoint: cutlass.GridConstant[tma.TensorMap], + desc_ws: cute.Tensor, + cu_seqlens: cute.Tensor, + k: cute.Tensor, + v: cute.Tensor, + gate: cute.Tensor, + state_checkpoints_out: Optional[cute.Tensor], + mStaging: Optional[cute.Tensor], + mCount: cute.Tensor, + mWorkItems: cute.Tensor, + mSched: Optional[cute.Tensor], + n_batch: cutlass.Int32, + k_row_stride: cutlass.Int32, + v_row_stride: cutlass.Int32, + checkpoint_row_stride: cutlass.Int32, + checkpoint_every_n: cutlass.Int32, +) -> None: + """Single-CTA prologue. Under ``run_order`` this kernel is the first + work-item-table consumer, so it LPT-orders the table and zeroes both + consumers' sched rings via :func:`order_body`; it then builds the + per-batch TMA-descriptor arrays via :func:`build_descs_body`, one warp + per array (the extra warps only take part in the order phase).""" + tidx, _, _ = cute.arch.thread_idx() + tidx = cutlass.Int32(tidx) + widx = tidx // cutlass.Int32(32) + if cutlass.const_expr(run_order): + sKey = cutlass.Array(cutlass.Int32, ORDER_CAPACITY, space=cutlass.AddressSpace.smem, alignment=16) + sIdx = cutlass.Array(cutlass.Int32, ORDER_CAPACITY, space=cutlass.AddressSpace.smem, alignment=16) + sSpread = cutlass.Array(cutlass.Int32, 2, space=cutlass.AddressSpace.smem, alignment=8) + n_heads_out = cutlass.Int32(gate.shape[1]) + order_body( + order_gen, + True, + b_t, + ORDER_THREADS, + ORDER_ELEMS, + tidx, + n_heads_out, + n_heads_out * n_batch, + cu_seqlens, + mStaging, + mCount, + mWorkItems, + mSched, + sKey, + sIdx, + sSpread, + ) + build_descs_body( + widx, + base_k, + base_v, + base_checkpoint, + desc_ws, + cu_seqlens, + k, + v, + state_checkpoints_out, + n_batch, + k_row_stride, + v_row_stride, + checkpoint_row_stride, + checkpoint_every_n, + ) + + @cute.jit -def build_descs( +def prologue( io_dtype: cutlass.Constexpr, b_t: cutlass.Constexpr[int], + run_order: cutlass.Constexpr[bool], + order_gen: cutlass.Constexpr[bool], k: cute.Tensor, v: cute.Tensor, + gate: cute.Tensor, cu_seqlens: cute.Tensor, state_checkpoints_out: Optional[cute.Tensor], + work_item_staging: Optional[cute.Tensor], + work_count: cute.Tensor, + work_items: cute.Tensor, + sched_all: Optional[cute.Tensor], checkpoint_every_n: cutlass.Int32, tensormap_workspace: cute.Tensor, stream: cuda.CUstream, ): - """Build the 3 per-batch TMA-descriptor arrays (K, V, checkpoints) into + """One-launch prologue: LPT-order the work items (with ``run_order``, when + this kernel is the backward pair's first table consumer) and build the 3 + per-batch TMA-descriptor arrays (K, V, checkpoints) into ``tensormap_workspace``.""" h_k = k.shape[1] h_v = v.shape[1] @@ -1739,8 +1848,10 @@ def build_descs( checkpoint_view, box_dims=(checkpoint_granu, d_k_state, 1, 1), stride_order=(0, 1, 2, 3), swizzle=swz128 ) - n_warps = 3 if state_checkpoints_out is not None else 2 - build_all_descs_kernel( + prologue_kernel( + run_order, + order_gen, + b_t, base_desc_k, base_desc_v, base_desc_checkpoint, @@ -1748,13 +1859,18 @@ def build_descs( cu_seqlens, k, v, + gate, state_checkpoints_out, + work_item_staging, + work_count, + work_items, + sched_all, cutlass.Int32(batch_size), cutlass.Int32(k_row_stride), cutlass.Int32(v_row_stride), cutlass.Int32(state_checkpoints_out.stride[0] if state_checkpoints_out is not None else 0), checkpoint_every_n, - ).launch(grid=(1, 1, 1), block=(32 * n_warps, 1, 1), stream=stream) + ).launch(grid=(1, 1, 1), block=(ORDER_THREADS, 1, 1), stream=stream) @cute.jit @@ -1763,6 +1879,8 @@ def host( k: cute.Tensor, v: cute.Tensor, gate: cute.Tensor, + a_log: Optional[cute.Tensor], + dt_bias: Optional[cute.Tensor], beta: cute.Tensor, cu_seqlens: cute.Tensor, state_in: Optional[cute.Tensor], @@ -1883,6 +2001,8 @@ def host( kernel( cfg, gate, + a_log, + dt_bias, beta, cu_seqlens, state_in, @@ -1910,6 +2030,8 @@ def host( def kernel( cfg: cutlass.Constexpr, mGate: cute.Tensor, + mA_log: Optional[cute.Tensor], + mDt_bias: Optional[cute.Tensor], mBeta: cute.Tensor, cu_seqlens: cute.Tensor, mState_init: Optional[cute.Tensor], @@ -2131,6 +2253,8 @@ def kernel( mWorkItems, tidx=tidx, mGate=mGate, + mA_log=mA_log, + mDt_bias=mDt_bias, mBeta=mBeta, sCumsumlog=sCumsumlog, sCumprod=sCumprod, @@ -2209,6 +2333,8 @@ class GdnRecomputeCfg: store_final_state: bool enable_checkpoints: bool log_gate: bool = False + safe_gate: bool = False + beta_sigmoid: bool = False dyn_sched: bool = False sched_stages: int = CFG.SMEM_SCHED_STAGES @@ -2277,6 +2403,8 @@ def build_cfg( store_final_state: bool = True, enable_checkpoints: bool = False, log_gate: bool = False, + safe_gate: bool = False, + beta_sigmoid: bool = False, dyn_sched: bool = False, ) -> GdnRecomputeCfg: """Build the per-compile ``GdnRecomputeCfg`` (io_dtype ∈ {Float16, BFloat16}; @@ -2293,6 +2421,8 @@ def build_cfg( store_final_state=store_final_state, enable_checkpoints=enable_checkpoints, log_gate=log_gate, + safe_gate=safe_gate, + beta_sigmoid=beta_sigmoid, dyn_sched=dyn_sched, ) cfg.smem_checkpoint_stages = 1 @@ -2335,7 +2465,11 @@ def get_compiled_cache( store_final_state: bool, enable_checkpoints: bool, log_gate: bool, + safe_gate: bool, + beta_sigmoid: bool, dyn_sched: bool, + run_order: bool, + order_gen: bool, ): """Return a mutable dict that lazily stores the compiled kernel.""" return {} @@ -2349,12 +2483,16 @@ def compile( store_final_state: bool, enable_checkpoints: bool, log_gate: bool = False, + safe_gate: bool = False, + beta_sigmoid: bool = False, dyn_sched: bool = False, *, num_sm: int, k_cute, v_cute, gate_cute, + a_log_cute=None, + dt_bias_cute=None, beta_cute, cu_seqlens_cute, state_in_cute, @@ -2376,6 +2514,8 @@ def compile( store_final_state=store_final_state, enable_checkpoints=enable_checkpoints, log_gate=log_gate, + safe_gate=safe_gate, + beta_sigmoid=beta_sigmoid, dyn_sched=dyn_sched, ) @@ -2385,6 +2525,8 @@ def compile( k_cute, v_cute, gate_cute, + a_log_cute, + dt_bias_cute, beta_cute, cu_seqlens_cute, state_in_cute, @@ -2412,7 +2554,14 @@ def chunk_gdn_recompute_sm100( work_items=None, work_count=None, sched_ctr=None, + sched_all=None, + work_item_scratch=None, + order_in_prologue: bool = False, log_gate: bool = False, + safe_gate: bool = False, + a_log=None, + dt_bias=None, + use_beta_sigmoid: bool = False, *, workspace, stream, @@ -2428,8 +2577,11 @@ def chunk_gdn_recompute_sm100( k: ``(total_tokens, HK, DK)`` float16/bfloat16 v: ``(total_tokens, HV, DK)`` float16/bfloat16 gate: ``(total_tokens, HO)`` float32, forget gate — raw linear - alpha, or the natural-log decay when ``log_gate`` - beta: ``(total_tokens, HO)`` float32, update gate + alpha, or the natural-log decay when ``log_gate``, or raw logits + when ``safe_gate``, which applies the safe-gate transform + ``-exp(a_log) * softplus(gate + dt_bias)`` + beta: ``(total_tokens, HO)`` float32, update gate — post-sigmoid, or + io-dtype logits when ``use_beta_sigmoid`` cu_seqlens: ``(num_seqs + 1,)`` int32 initial_state: ``(num_seqs, HO, DK, DK)`` float32/bfloat16, or None output_state: ``(num_seqs, HO, DK, DK)`` float32/bfloat16, or None @@ -2447,6 +2599,10 @@ def chunk_gdn_recompute_sm100( log_gate: ``gate`` holds natural-log decay values; the gate warp skips its log2 (rescales by 1/ln2) instead of exponentiating upstream + safe_gate: interpret ``gate`` through the safe-gate transform + a_log: ``(HO,)`` float32, safe-gate per-head log-amplitude (None = 0) + dt_bias: ``(HO,)`` float32, safe-gate per-head bias (None = 0) + use_beta_sigmoid: ``beta`` holds logits; sigmoid in-kernel workspace: ``(>= tensormap_workspace_bytes(module, B) // 8,)`` int64, 128-byte aligned; holds the per-(b,h) TMA descriptors (contents managed here — reuse the same buffer across calls) @@ -2464,8 +2620,17 @@ def chunk_gdn_recompute_sm100( if work_items is None or work_count is None: raise ValueError("work_items/work_count are required (the split-table stage builds them for every launch)") dyn_sched = sched_ctr is not None + run_order = bool(order_in_prologue) + order_gen = work_item_scratch is None + if run_order and sched_all is None: + raise ValueError("order_in_prologue requires sched_all (the prologue zeroes both consumers' sched rings)") if not (enable_checkpoints or store_final_state): raise ValueError("output_state_checkpoints or output_state is required") + if safe_gate and (a_log is None or dt_bias is None): + raise ValueError("safe_gate requires a_log and dt_bias") + if not safe_gate: + a_log = None + dt_bias = None io_dtype = get_dtype(k.dtype) if initial_state is not None: @@ -2490,7 +2655,11 @@ def chunk_gdn_recompute_sm100( store_final_state, enable_checkpoints, log_gate, + safe_gate, + use_beta_sigmoid, dyn_sched, + run_order, + order_gen, ) if "compiled" not in cache: @@ -2500,6 +2669,8 @@ def chunk_gdn_recompute_sm100( v_cute.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1, 2), divisibility=1) gate_cute = from_dlpack(gate, assumed_align=16) gate_cute.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1), divisibility=1) + a_log_cute = from_dlpack(a_log, assumed_align=4) if a_log is not None else None + dt_bias_cute = from_dlpack(dt_bias, assumed_align=4) if dt_bias is not None else None beta_cute = from_dlpack(beta, assumed_align=16) beta_cute.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1), divisibility=1) cu_seqlens_cute = from_dlpack(cu_seqlens, assumed_align=8 if str(cu_seqlens.dtype).endswith("int64") else 4).mark_layout_dynamic() @@ -2532,11 +2703,15 @@ def chunk_gdn_recompute_sm100( store_final_state, enable_checkpoints, log_gate, + safe_gate, + use_beta_sigmoid, dyn_sched, num_sm=multiprocessor_count(current_device_id()), k_cute=k_cute, v_cute=v_cute, gate_cute=gate_cute, + a_log_cute=a_log_cute, + dt_bias_cute=dt_bias_cute, beta_cute=beta_cute, cu_seqlens_cute=cu_seqlens_cute, state_in_cute=state_in_cute, @@ -2551,39 +2726,60 @@ def chunk_gdn_recompute_sm100( compiled = cache["compiled"] - # desc build runs every execute by contract (cu contents are data; - # buffer pointers may change) — capture-safe, single tiny launch - if "build_descs" not in cache: - k_bc = from_dlpack(k, assumed_align=16) - k_bc.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1, 2), divisibility=1) - v_bc = from_dlpack(v, assumed_align=16) - v_bc.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1, 2), divisibility=1) - cu_bc = from_dlpack(cu_seqlens, assumed_align=8 if str(cu_seqlens.dtype).endswith("int64") else 4).mark_layout_dynamic() - checkpoints_bc = None - cu_ckpt_bc = None + if "prologue" not in cache: + k_pl = from_dlpack(k, assumed_align=16) + k_pl.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1, 2), divisibility=1) + v_pl = from_dlpack(v, assumed_align=16) + v_pl.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1, 2), divisibility=1) + gate_pl = from_dlpack(gate, assumed_align=16) + gate_pl.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1), divisibility=1) + cu_pl = from_dlpack(cu_seqlens, assumed_align=8 if str(cu_seqlens.dtype).endswith("int64") else 4).mark_layout_dynamic() + checkpoints_pl = None if enable_checkpoints: - checkpoints_bc = from_dlpack(output_state_checkpoints, assumed_align=16) - checkpoints_bc.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1, 2, 3), divisibility=1) - ws_bc = from_dlpack(workspace, assumed_align=128).mark_layout_dynamic() - cache["build_descs"] = cute.compile( - build_descs, + checkpoints_pl = from_dlpack(output_state_checkpoints, assumed_align=16) + checkpoints_pl.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1, 2, 3), divisibility=1) + staging_pl = None + if not order_gen: + staging_pl = from_dlpack(work_item_scratch, assumed_align=16) + staging_pl.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1), divisibility=1) + work_count_pl = from_dlpack(work_count, assumed_align=4).mark_layout_dynamic() + work_items_pl = from_dlpack(work_items, assumed_align=16) + work_items_pl.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1), divisibility=1) + sched_all_pl = None + if run_order: + sched_all_pl = from_dlpack(sched_all, assumed_align=4).mark_layout_dynamic() + ws_pl = from_dlpack(workspace, assumed_align=128).mark_layout_dynamic() + cache["prologue"] = cute.compile( + prologue, io_dtype, CFG.B_T, - k_bc, - v_bc, - cu_bc, - checkpoints_bc, - cutlass.Int32(checkpoint_every_n_tokens if enable_checkpoints else 1), - ws_bc, + run_order, + order_gen, + k_pl, + v_pl, + gate_pl, + cu_pl, + checkpoints_pl, + staging_pl, + work_count_pl, + work_items_pl, + sched_all_pl, + cutlass.Int32(checkpoint_every_n_tokens), + ws_pl, cu_stream, options="--enable-tvm-ffi", ) - cache["build_descs"]( + cache["prologue"]( k, v, + gate, cu_seqlens, - output_state_checkpoints, - checkpoint_every_n_tokens if enable_checkpoints else 1, + output_state_checkpoints if enable_checkpoints else None, + work_item_scratch if not order_gen else None, + work_count, + work_items, + sched_all if run_order else None, + checkpoint_every_n_tokens, workspace, cu_stream, ) @@ -2591,6 +2787,8 @@ def chunk_gdn_recompute_sm100( k, v, gate, + a_log, + dt_bias, beta, cu_seqlens, initial_state, @@ -2602,3 +2800,62 @@ def chunk_gdn_recompute_sm100( workspace, cu_stream, ) + return cache + + +def run_recompute( + cache, + k, + v, + gate, + beta, + cu_seqlens, + initial_state, + output_state, + output_state_checkpoints, + work_items, + work_count, + sched_ctr, + sched_all, + work_item_scratch, + tensormap_workspace, + checkpoint_every_n_tokens, + stream, + a_log=None, + dt_bias=None, +) -> None: + """Replay the compiled plan: the prologue launch, then the main launch. + The caller owns the contract, which the plan validated at build, so + nothing here raises.""" + cu_stream = cuda.CUstream(int(stream)) + cache["prologue"]( + k, + v, + gate, + cu_seqlens, + output_state_checkpoints, + work_item_scratch, + work_count, + work_items, + sched_all, + checkpoint_every_n_tokens, + tensormap_workspace, + cu_stream, + ) + cache["compiled"]( + k, + v, + gate, + a_log, + dt_bias, + beta, + cu_seqlens, + initial_state, + output_state, + work_items, + work_count, + sched_ctr, + checkpoint_every_n_tokens, + tensormap_workspace, + cu_stream, + ) diff --git a/python/cudnn/linear_attention/frost/kernel/kda_bprop_f16.py b/python/cudnn/linear_attention/frost/kernel/kda_bprop_f16.py index e9576c212..76bb60230 100644 --- a/python/cudnn/linear_attention/frost/kernel/kda_bprop_f16.py +++ b/python/cudnn/linear_attention/frost/kernel/kda_bprop_f16.py @@ -35,8 +35,13 @@ ABI: state_checkpoints `[total_checkpoints, HO, DK, DV]` (KV, v contiguous) io dtype, the plain per-chunk checkpoint series with NO initial-state slot (entry `c - 1` = state entering chunk c >= 1; chunk 0 seeds from `initial_state`); dq/dk/dv io at HO heads; dgate `[T, HO, DK]` fp32 -(natural-log gate domain); dbeta `[T, HO]` fp32; d_initial_state / -d_final_state fp32 `[N, HO, DK, DV]` (K-major). +(natural-log gate domain; with SAFE_GATE the gradient stays wrt the +transformed log-decay); dbeta `[T, HO]` fp32, io dtype with BETA_SIGMOID +(post-sigmoid space, or wrt the raw logits under BETA_SIGMOID). Gate +arrives natural-log fp32 unless SAFE_GATE (safe-gate transform from raw gate ++ a_log/dt_bias); beta arrives post-sigmoid fp32, or io-dtype logits with +BETA_SIGMOID; d_initial_state / d_final_state fp32 `[N, HO, DK, DV]` +(K-major). Warp assignments (16 warps = 512 threads): warps 0-3 : WG0 - Gate prefix scan + decay/restore operands (all chunks) + Beta scalar gather @@ -59,12 +64,13 @@ import cutlass.cute as cute from cutlass.cute.runtime import from_dlpack -from ..common.split_k import decode_work_item +from ..common.split_k import ORDER_CAPACITY, ORDER_ELEMS, ORDER_THREADS, decode_work_item, order_body from ..common.host import get_dtype from cudnn.frost.buffers import current_device_id, data_ptr from cudnn.frost.device import multiprocessor_count from ..common.thd import TENSOR_MAP_QWORDS, emit_copy_desc, emit_checkpoint_seq_descs, emit_seq_descs from .kda_bprop_config import CFG + from cudnn.frost.tile_dsl.barrier import ( advance, MBarrier, @@ -76,18 +82,23 @@ from cudnn.frost.tile_dsl.swizzle import swizzle_lin_S, swizzle_xor_128b from cudnn.frost.tile_dsl.tma import tma_load_tile, tma_store_commit, tma_store_tile, tma_store_wait, tma_tensormap_acquire from cudnn.frost.tile_dsl.pointwise import ( - opaque_f32_zero, f16x2_to_f32, - fmul2, + fadd2, ffma2, + fmul2, + fp32_to_fp16, movmatrix_16b, mul_f16x2, - fp32_to_fp16, + opaque_f32_zero, sub_f16x2, ) LOG2_E: float = 1.4426950408889634 + +DEFAULT_GATE_LOWER_BOUND: float = -5.0 + + L2_NORM_EPS: float = 1.0e-12 @@ -768,14 +779,10 @@ def super_mma_warp( tinv_lo1, tinv_hi1 = f16x2_to_f32(tinv_p1, dtype=cfg.io_dtype) tinv_lo2, tinv_hi2 = f16x2_to_f32(tinv_p2, dtype=cfg.io_dtype) tinv_lo3, tinv_hi3 = f16x2_to_f32(tinv_p3, dtype=cfg.io_dtype) - tinv_acc[0] = tinv_lo0 + upd_acc[0] - tinv_acc[1] = tinv_hi0 + upd_acc[1] - tinv_acc[2] = tinv_lo1 + upd_acc[2] - tinv_acc[3] = tinv_hi1 + upd_acc[3] - tinv_acc[4] = tinv_lo2 + upd_acc[4] - tinv_acc[5] = tinv_hi2 + upd_acc[5] - tinv_acc[6] = tinv_lo3 + upd_acc[6] - tinv_acc[7] = tinv_hi3 + upd_acc[7] + tinv_acc[0], tinv_acc[1] = fadd2(tinv_lo0, tinv_hi0, upd_acc[0], upd_acc[1]) + tinv_acc[2], tinv_acc[3] = fadd2(tinv_lo1, tinv_hi1, upd_acc[2], upd_acc[3]) + tinv_acc[4], tinv_acc[5] = fadd2(tinv_lo2, tinv_hi2, upd_acc[4], upd_acc[5]) + tinv_acc[6], tinv_acc[7] = fadd2(tinv_lo3, tinv_hi3, upd_acc[6], upd_acc[7]) nvvm.stmatrix( sIntermediate_ptr + 1 * (cfg.b_t * cfg.b_t) + stsm_idx, @@ -1680,6 +1687,18 @@ def tmaldg_warp( tile_idx = next_tile +@cute.jit +def gate_scale(cfg, raw_gate: cutlass.Float32) -> cutlass.Float32: + """Map raw gate to the log2-domain decay increment used by KDA.""" + + if cutlass.const_expr(cfg.safe_gate): + half = cutlass.Float32(0.5) + sigmoid = cute.math.tanh(raw_gate * half, approx=True) * half + half + return cfg.gate_scale_log2 * sigmoid + # Default ABI: Gate arrives in natural-log space + return raw_gate * cutlass.Float32(LOG2_E) + + @cute.jit def compute0_warp_group( cfg, @@ -1693,6 +1712,8 @@ def compute0_warp_group( tmem_base_holder, warp_idx, scale, + mA_log, + mDt_bias, mBeta, sBeta_raw, sK_inv_raw, @@ -1714,6 +1735,9 @@ def compute0_warp_group( H -> TMEM f16 at the chunk tail.""" nvvm.setmaxregister(cfg.num_regs_compute_group_0, nvvm.SetMaxRegisterAction.INCREASE) cg0_warp = warp_idx - cfg.compute_group_0_warp_ids[0] + prefix_dim = cg0_warp * cfg.threads_per_warp + lane + cg0_a_log_exp = cutlass.Float32(1.0) + cg0_dt_bias_value = cutlass.Float32(0.0) nvvm.barrier_cta_sync(cfg.tmem_lifecycle_barrier_id, thread_count=cfg.tmem_user_threads) tmem_base = tmem_base_holder.load() tmem_col = tmem_base & 0xFFFF @@ -1730,6 +1754,10 @@ def compute0_warp_group( while tile_idx < total_tiles: batch_idx, head_idx, batch_start, batch_end, seqlen_b, num_chunks_b, wstart, wend, cstart, cend = decode_work_item(cfg, tile_idx, mWorkItems) num_compute_chunks = cend - wstart + if cutlass.const_expr(cfg.safe_gate): + if num_compute_chunks > 0: + cg0_a_log_exp = cute.math.exp2(mA_log[head_idx].to(cutlass.Float32) * LOG2_E, fastmath=True) + cg0_dt_bias_value = mDt_bias[head_idx, prefix_dim].to(cutlass.Float32) for rev_idx in cutlass.range(num_compute_chunks, unroll=1): chunk_idx = cend - cutlass.Int32(1) - rev_idx chunk_serial = chunk_serial_base + rev_idx @@ -1754,6 +1782,9 @@ def compute0_warp_group( beta_value = cutlass.Float32(0.0) if token_idx < seqlen_b: beta_value = mBeta[batch_start + token_idx, head_idx].to(cutlass.Float32) + if cutlass.const_expr(cfg.beta_sigmoid): + half = cutlass.Float32(0.5) + beta_value = (cute.math.tanh(beta_value * half, approx=True) * half + half).to(mBeta.element_type).to(cutlass.Float32) sBeta_raw[beta_stage * cfg.b_t + lane] = beta_value bars.mb_beta_ready[beta_stage].arrive() @@ -1776,14 +1807,38 @@ def compute0_warp_group( prefix_idx = f32_segment * (cfg.b_t * 32) + row * 32 + swizzle_xor_128b(row, f32_segment_dim, elem_bytes=4) gate_raw[row] = (sGate_ptr + prefix_idx).load() g_prefix_regs = cutlass.Array(cutlass.Float32, cfg.b_t, alignment=16) - for row in cutlass.range_constexpr(cfg.b_t): - gate = gate_raw[row] - token_idx = chunk_idx * cutlass.Int32(cfg.b_t) + cutlass.Int32(row) - if token_idx < seqlen_b: - gate = gate * cutlass.Float32(LOG2_E) - else: - gate = cutlass.Float32(0.0) - g_prefix_regs[row] = gate + if cutlass.const_expr(cfg.safe_gate): + valid_rows = seqlen_b - chunk_idx * cutlass.Int32(cfg.b_t) + valid_mask = cutlass.vector.create_mask([cfg.b_t], [valid_rows]) + for row_pair in cutlass.range_constexpr(cfg.b_t // 2): + row0 = row_pair * 2 + row1 = row0 + 1 + gate0 = cg0_a_log_exp * (gate_raw[row0] + cg0_dt_bias_value) + gate1 = cg0_a_log_exp * (gate_raw[row1] + cg0_dt_bias_value) + gate0 = gate_scale( + cfg, + gate0, + ) + gate1 = gate_scale( + cfg, + gate1, + ) + gate_pair = cutlass.Vector.from_elements((gate0, gate1), cutlass.Float32) + gate_pair = cutlass.vector.where(valid_mask[row0 : row1 + 1], gate_pair, 0.0) + g_prefix_regs[row0] = gate_pair[0] + g_prefix_regs[row1] = gate_pair[1] + else: + for row in cutlass.range_constexpr(cfg.b_t): + gate = gate_raw[row] + token_idx = chunk_idx * cutlass.Int32(cfg.b_t) + cutlass.Int32(row) + if token_idx < seqlen_b: + gate = gate_scale( + cfg, + gate, + ) + else: + gate = cutlass.Float32(0.0) + g_prefix_regs[row] = gate prefix_acc = cutlass.Float32(0.0) for row_pair in cutlass.range_constexpr(cfg.b_t // 2): @@ -2306,6 +2361,11 @@ def compute1_warp_group( beta_c1 = (sBeta_ptr + (lane % 4) * 2 + 1).load().to(cutlass.Float32) beta_c8 = (sBeta_ptr + (lane % 4) * 2 + 8).load().to(cutlass.Float32) beta_c9 = (sBeta_ptr + (lane % 4) * 2 + 9).load().to(cutlass.Float32) + # latched before the stage is released; the dbeta store runs after the refill + beta_self = cutlass.Float32(0.0) + if cutlass.const_expr(cfg.beta_sigmoid): + if cg1_tidx < cfg.b_t: + beta_self = (sBeta_ptr + cg1_tidx).load().to(cutlass.Float32) bars.mb_beta_done[chunk_serial % cfg.smem_beta_stages].arrive() beta_dy_regs0 = cutlass.Array(cutlass.Float32, 8, alignment=16) beta_dy_regs1 = cutlass.Array(cutlass.Float32, 8, alignment=16) @@ -2389,9 +2449,11 @@ def compute1_warp_group( for w in cutlass.range_constexpr(4): acc = acc + sRed_raw[w * cfg.b_t + cg1_tidx] db_val = acc + sBetaM_raw[cg1_tidx] + if cutlass.const_expr(cfg.beta_sigmoid): + db_val = db_val * (beta_self - beta_self * beta_self) token_idx = chunk_idx * cutlass.Int32(cfg.b_t) + cg1_tidx if token_idx < seqlen_b and chunk_idx < wend: - mDbeta[batch_start + token_idx, head_idx] = db_val + mDbeta[batch_start + token_idx, head_idx] = db_val.to(mDbeta.element_type) nvvm.barrier_cta_sync(cfg.cg1_sync_barrier_id, thread_count=cfg.cg1_threads) # ---- dH capture for the next ----------------------------------------- @@ -2576,8 +2638,9 @@ def compute2_warp_group( hval1 = (sDstate_raw.data_ptr() + dstate_addr1).load().to(cutlass.Float32) hacc[(2 * j) % 8] = hacc[(2 * j) % 8] + hval0 * state_pair[0].to(cutlass.Float32) hacc[(2 * j + 1) % 8] = hacc[(2 * j + 1) % 8] + hval1 * state_pair[1].to(cutlass.Float32) - part_a = (hacc[0] + hacc[4]) + (hacc[1] + hacc[5]) - part_b = (hacc[2] + hacc[6]) + (hacc[3] + hacc[7]) + pa0, pb0 = fadd2(hacc[0], hacc[2], hacc[4], hacc[6]) + pa1, pb1 = fadd2(hacc[1], hacc[3], hacc[5], hacc[7]) + part_a, part_b = fadd2(pa0, pb0, pa1, pb1) dgate_last_val = dgate_last_val + (part_a + part_b) bars.mb_dstate_smem_cg2_done.arrive() bars.mb_state_inp_cg2_done[chunk_serial % 2].arrive() @@ -2658,26 +2721,36 @@ def compute2_warp_group( dots = cutlass.Array(cutlass.Float32, cfg.b_t, alignment=16) for half in cutlass.range_constexpr(2): p_words = nvvm.tcgen05_ld("32x32b", nvvm.make_tmem_ptr(row_addr + qk_col + half * (cfg.b_t // 4), cutlass.Float32), num=cfg.b_t // 4) - for tt in cutlass.range_constexpr(cfg.b_t // 2): - t = half * (cfg.b_t // 2) + tt - p_pair = cutlass.Vector.from_elements((p_words[tt // 2],), cutlass.Float32).bitcast(cfg.io_dtype) - dots[t] = grad[t] * p_pair[tt % 2].to(cutlass.Float32) * sNorm_raw[norm_base + inv_off + t] + for tt2 in cutlass.range_constexpr(cfg.b_t // 4): + t = half * (cfg.b_t // 2) + 2 * tt2 + p_pair = cutlass.Vector.from_elements((p_words[tt2],), cutlass.Float32).bitcast(cfg.io_dtype) + gp_lo, gp_hi = fmul2(grad[t], grad[t + 1], p_pair[0].to(cutlass.Float32), p_pair[1].to(cutlass.Float32)) + dots[t], dots[t + 1] = fmul2(gp_lo, gp_hi, sNorm_raw[norm_base + inv_off + t], sNorm_raw[norm_base + inv_off + t + 1]) for off in cutlass.range_constexpr(5): step = cutlass.const_expr(1 << off) - for t in cutlass.range_constexpr(cfg.b_t): - dots[t] = dots[t] + cutlass.Float32(nvvm.shfl_sync(0xFFFFFFFF, dots[t], step, 31, kind=nvvm.Shfl.BFLY)) + for t2 in cutlass.range_constexpr(cfg.b_t // 2): + t = 2 * t2 + sh_lo = cutlass.Float32(nvvm.shfl_sync(0xFFFFFFFF, dots[t], step, 31, kind=nvvm.Shfl.BFLY)) + sh_hi = cutlass.Float32(nvvm.shfl_sync(0xFFFFFFFF, dots[t + 1], step, 31, kind=nvvm.Shfl.BFLY)) + dots[t], dots[t + 1] = fadd2(dots[t], dots[t + 1], sh_lo, sh_hi) if lane == 0: for t in cutlass.range_constexpr(cfg.b_t): sRed1_raw[tmem_subpartition * cfg.b_t + t] = dots[t] nvvm.barrier_cta_sync(cfg.cg2_sync_barrier_id, thread_count=cfg.cg2_threads) for half in cutlass.range_constexpr(2): a_words = nvvm.tcgen05_ld("32x32b", nvvm.make_tmem_ptr(row_addr + qk_col + half * (cfg.b_t // 4), cutlass.Float32), num=cfg.b_t // 4) - for tt in cutlass.range_constexpr(cfg.b_t // 2): - t = half * (cfg.b_t // 2) + tt - a_pair = cutlass.Vector.from_elements((a_words[tt // 2],), cutlass.Float32).bitcast(cfg.io_dtype) - total_dot = sRed1_raw[t] + sRed1_raw[cfg.b_t + t] + sRed1_raw[2 * cfg.b_t + t] + sRed1_raw[3 * cfg.b_t + t] - norm_t = sNorm_raw[norm_base + inv_off + t] - grad[t] = (grad[t] - a_pair[tt % 2].to(cutlass.Float32) * norm_t * total_dot) * norm_t + for tt2 in cutlass.range_constexpr(cfg.b_t // 4): + t = half * (cfg.b_t // 2) + 2 * tt2 + a_pair = cutlass.Vector.from_elements((a_words[tt2],), cutlass.Float32).bitcast(cfg.io_dtype) + dot_lo, dot_hi = fadd2(sRed1_raw[t], sRed1_raw[t + 1], sRed1_raw[cfg.b_t + t], sRed1_raw[cfg.b_t + t + 1]) + dot_lo, dot_hi = fadd2(dot_lo, dot_hi, sRed1_raw[2 * cfg.b_t + t], sRed1_raw[2 * cfg.b_t + t + 1]) + dot_lo, dot_hi = fadd2(dot_lo, dot_hi, sRed1_raw[3 * cfg.b_t + t], sRed1_raw[3 * cfg.b_t + t + 1]) + norm_lo = sNorm_raw[norm_base + inv_off + t] + norm_hi = sNorm_raw[norm_base + inv_off + t + 1] + an_lo, an_hi = fmul2(a_pair[0].to(cutlass.Float32), a_pair[1].to(cutlass.Float32), norm_lo, norm_hi) + sub_lo = grad[t] - an_lo * dot_lo + sub_hi = grad[t + 1] - an_hi * dot_hi + grad[t], grad[t + 1] = fmul2(sub_lo, sub_hi, norm_lo, norm_hi) nvvm.barrier_cta_sync(cfg.cg2_sync_barrier_id, thread_count=cfg.cg2_threads) bars.mb_qk_raw_done[qk_raw_stage].arrive() @@ -2729,19 +2802,20 @@ def compute2_warp_group( # --------------------------------------------------------------------------- -@cute.kernel -def build_all_descs_kernel( - base_q: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_k: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_v: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_gate: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_do: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_dq: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_dk: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_dv: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_dgate: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_checkpoint: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_initial_state: cutlass.GridConstant[cuda.tensor_map.TensorMap], +@cute.jit +def build_descs_body( + widx, + base_q, + base_k, + base_v, + base_gate, + base_do, + base_dq, + base_dk, + base_dv, + base_dgate, + base_checkpoint, + base_initial_state, desc_ws: cute.Tensor, cu_seqlens: cute.Tensor, q: cute.Tensor, @@ -2768,10 +2842,9 @@ def build_all_descs_kernel( checkpoint_row_stride: cutlass.Int32, checkpoint_every_n: cutlass.Int32, ) -> None: - """Single-launch builder for the per-batch descriptor arrays (one warp - per array).""" - tidx, _, _ = cute.arch.thread_idx() - widx = cutlass.Int32(tidx) // cutlass.Int32(32) + """Per-batch descriptor-array build, one warp per array. Runs inside the + prologue kernel after its order pass; warps past the array count fall + through the widx guards.""" arr_words = n_batch * cutlass.Int32(TENSOR_MAP_QWORDS) sub0 = cute.make_tensor(desc_ws.iterator, cute.make_layout((arr_words,), stride=(1,))) sub1 = cute.make_tensor(desc_ws.iterator + arr_words, cute.make_layout((arr_words,), stride=(1,))) @@ -2832,10 +2905,132 @@ def build_all_descs_kernel( nvvm.fence_proxy_release(nvvm.MemScope.GPU, from_proxy=nvvm.Proxy.GENERIC, to_proxy=nvvm.Proxy.TENSORMAP) +@cute.kernel +def prologue_kernel( + run_order: cutlass.Constexpr[bool], + order_gen: cutlass.Constexpr[bool], + has_sched: cutlass.Constexpr[bool], + b_t: cutlass.Constexpr[int], + base_q: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_k: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_v: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_gate: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_do: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_dq: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_dk: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_dv: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_dgate: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_checkpoint: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_initial_state: cutlass.GridConstant[cuda.tensor_map.TensorMap], + desc_ws: cute.Tensor, + cu_seqlens: cute.Tensor, + q: cute.Tensor, + k: cute.Tensor, + v: cute.Tensor, + gate: cute.Tensor, + do: cute.Tensor, + dq: cute.Tensor, + dk: cute.Tensor, + dv: cute.Tensor, + dgate: cute.Tensor, + state_checkpoints: cute.Tensor, + state0: cute.Tensor | None, + mStaging: cute.Tensor | None, + mCount: cute.Tensor, + mWorkItems: cute.Tensor | None, + mSched: cute.Tensor | None, + n_batch: cutlass.Int32, + q_row_stride: cutlass.Int32, + k_row_stride: cutlass.Int32, + v_row_stride: cutlass.Int32, + gate_row_stride: cutlass.Int32, + do_row_stride: cutlass.Int32, + dq_row_stride: cutlass.Int32, + dk_row_stride: cutlass.Int32, + dv_row_stride: cutlass.Int32, + dgate_row_stride: cutlass.Int32, + checkpoint_row_stride: cutlass.Int32, + checkpoint_every_n: cutlass.Int32, +) -> None: + """Single-CTA prologue. Under ``run_order`` this kernel is the first + work-item-table consumer, so it LPT-orders the table and zeroes both + consumers' sched rings via :func:`order_body`; it then builds the + per-batch TMA-descriptor arrays via :func:`build_descs_body`, one warp + per array (the extra warps only take part in the order phase).""" + tidx, _, _ = cute.arch.thread_idx() + tidx = cutlass.Int32(tidx) + widx = tidx // cutlass.Int32(32) + if cutlass.const_expr(run_order): + sKey = cutlass.Array(cutlass.Int32, ORDER_CAPACITY, space=cutlass.AddressSpace.smem, alignment=16) + sIdx = cutlass.Array(cutlass.Int32, ORDER_CAPACITY, space=cutlass.AddressSpace.smem, alignment=16) + sSpread = cutlass.Array(cutlass.Int32, 2, space=cutlass.AddressSpace.smem, alignment=8) + n_heads_out = cutlass.Int32(gate.shape[1]) + order_body( + order_gen, + has_sched, + b_t, + ORDER_THREADS, + ORDER_ELEMS, + tidx, + n_heads_out, + n_heads_out * n_batch, + cu_seqlens, + mStaging, + mCount, + mWorkItems, + mSched, + sKey, + sIdx, + sSpread, + ) + build_descs_body( + widx, + base_q, + base_k, + base_v, + base_gate, + base_do, + base_dq, + base_dk, + base_dv, + base_dgate, + base_checkpoint, + base_initial_state, + desc_ws, + cu_seqlens, + q, + k, + v, + gate, + do, + dq, + dk, + dv, + dgate, + state_checkpoints, + state0, + n_batch, + q_row_stride, + k_row_stride, + v_row_stride, + gate_row_stride, + do_row_stride, + dq_row_stride, + dk_row_stride, + dv_row_stride, + dgate_row_stride, + checkpoint_row_stride, + checkpoint_every_n, + ) + + @cute.jit -def build_descs( +def prologue( io_dtype: cutlass.Constexpr, b_t: cutlass.Constexpr[int], + run_order: cutlass.Constexpr[bool], + order_gen: cutlass.Constexpr[bool], + has_sched: cutlass.Constexpr[bool], q: cute.Tensor, k: cute.Tensor, v: cute.Tensor, @@ -2848,10 +3043,15 @@ def build_descs( state_checkpoints: cute.Tensor, state0: cute.Tensor | None, cu_seqlens: cute.Tensor, + work_item_staging: cute.Tensor | None, + work_count: cute.Tensor, + work_items: cute.Tensor | None, + sched_all: cute.Tensor | None, tensormap_workspace: cute.Tensor, stream: cuda_driver.CUstream, ): - """Build the 11 per-(batch, head) capped TMA-descriptor arrays into + """One-launch prologue: LPT-order the work items (when ``run_order``) and + build the 11 per-(batch, head) capped TMA-descriptor arrays into ``tensormap_workspace`` (sequence-relative coordinates; tail loads zero-fill and tail stores clip in hardware).""" h_q = q.shape[1] @@ -2905,8 +3105,11 @@ def build_descs( ) base_initial_state = cuda.create_tensor_map_tiled_from_view(initial_state_view, box_dims=(64, d_k, 1, 1), stride_order=(0, 1, 2, 3), swizzle=swz) - n_warps = 11 if state0 is not None else 10 - build_all_descs_kernel( + prologue_kernel( + run_order, + order_gen, + has_sched, + b_t, base_q, base_k, base_v, @@ -2931,6 +3134,10 @@ def build_descs( dgate, state_checkpoints, state0, + work_item_staging, + work_count, + work_items, + sched_all, cutlass.Int32(batch_size), cutlass.Int32(q.stride[0]), cutlass.Int32(k.stride[0]), @@ -2943,12 +3150,14 @@ def build_descs( cutlass.Int32(dgate.stride[0]), cutlass.Int32(state_checkpoints.stride[0]), cutlass.Int32(b_t), - ).launch(grid=(1, 1, 1), block=(32 * n_warps, 1, 1), stream=stream) + ).launch(grid=(1, 1, 1), block=(ORDER_THREADS, 1, 1), stream=stream) @cute.jit def host( cfg: cutlass.Constexpr, + a_log: cute.Tensor | None, + dt_bias: cute.Tensor | None, beta: cute.Tensor, state_checkpoints: cute.Tensor, mState_init: cute.Tensor | None, @@ -2973,6 +3182,8 @@ def host( cfg, tensormap_workspace, n_desc, + a_log, + dt_bias, beta, cu_seqlens, dgate, @@ -2996,6 +3207,8 @@ def kernel( cfg: cutlass.Constexpr, tensormap_workspace: cute.Tensor, n_desc: cutlass.Int32, + mA_log: cute.Tensor | None, + mDt_bias: cute.Tensor | None, mBeta: cute.Tensor, cu_seqlens: cute.Tensor, mDgate: cute.Tensor, @@ -3015,9 +3228,10 @@ def kernel( lane = tidx % cfg.threads_per_warp total_tiles = mCount[0] - assert mBeta.element_type == cutlass.Float32 + beta_expected = cfg.io_dtype if cutlass.const_expr(cfg.beta_sigmoid) else cutlass.Float32 + assert mBeta.element_type == beta_expected assert cu_seqlens.element_type in (cutlass.Int32, cutlass.Int64) - assert mDgate.element_type == cutlass.Float32 and mDbeta.element_type == cutlass.Float32 + assert mDgate.element_type == cutlass.Float32 and mDbeta.element_type == beta_expected desc_base_words = tensormap_workspace.iterator.raw_ptr() arr_words = n_desc * cutlass.Int32(TENSOR_MAP_QWORDS) @@ -3381,6 +3595,8 @@ def kernel( tmem_base_holder, warp_idx, scale, + mA_log, + mDt_bias, mBeta, sBeta_raw, sK_inv_raw, @@ -3454,6 +3670,9 @@ class KdaBwdCfg: use_dstate_in: bool use_dstate0: bool l2norm: bool + safe_gate: bool + gate_scale_log2: float + beta_sigmoid: bool use_initial_state: bool q_ratio: int k_ratio: int @@ -3549,6 +3768,9 @@ def build_cfg( use_dstate_in: bool, use_dstate0: bool, l2norm: bool, + safe_gate: bool, + gate_scale_log2: float, + beta_sigmoid: bool, use_initial_state: bool, q_ratio: int, k_ratio: int, @@ -3564,6 +3786,9 @@ def build_cfg( use_dstate_in=use_dstate_in, use_dstate0=use_dstate0, l2norm=l2norm, + safe_gate=safe_gate, + gate_scale_log2=gate_scale_log2, + beta_sigmoid=beta_sigmoid, use_initial_state=use_initial_state, q_ratio=q_ratio, k_ratio=k_ratio, @@ -3635,8 +3860,13 @@ def get_compiled_cache( use_dstate_in: bool, use_dstate0: bool, l2norm: bool, + safe_gate: bool, + gate_lower_bound: float, + beta_sigmoid: bool, use_initial_state: bool, dyn_sched: bool, + run_order: bool, + order_gen: bool, ): return {} @@ -3661,9 +3891,17 @@ def chunk_kda_bwd_sm100( d_initial_state=None, d_final_state=None, use_qk_l2norm_in_kernel: bool = False, + safe_gate: bool = False, + gate_lower_bound: float = DEFAULT_GATE_LOWER_BOUND, + a_log=None, + dt_bias=None, + use_beta_sigmoid: bool = False, work_items=None, work_count=None, sched_ctr=None, + sched_all=None, + work_item_scratch=None, + order_in_prologue: bool = False, tensormap_workspace, stream, ) -> None: @@ -3675,15 +3913,23 @@ def chunk_kda_bwd_sm100( q: ``(total_tokens, HQ, DK)`` float16/bfloat16 k: ``(total_tokens, HK, DK)`` float16/bfloat16 v: ``(total_tokens, HV, DV)`` float16/bfloat16 - gate: ``(total_tokens, HO, DK)`` fp32 natural-log per-channel decay - beta: ``(total_tokens, HO)`` fp32 post-sigmoid + gate: ``(total_tokens, HO, DK)`` fp32. Natural-log per-channel decay + unless ``safe_gate``, which applies the safe-gate transform + ``lower_bound * sigmoid(exp(a_log) * (gate + dt_bias))``. + beta: ``(total_tokens, HO)``. Post-sigmoid float32, or io-dtype + logits when ``use_beta_sigmoid`` do: ``(total_tokens, HO, DV)`` io dtype state_checkpoints: ``(total_checkpoints, HO, DK, DV)`` io dtype (KV, v contiguous), the PLAIN per-chunk checkpoint series with no initial-state slot: sequence-local entry ``c - 1`` is the state ENTERING chunk c >= 1 of sequence b; chunk 0 seeds from ``initial_state`` dq/dk/dv: io dtype at ``HO = max(HQ, HV)`` heads, pre-allocated - dgate: ``(total_tokens, HO, DK)`` fp32 (dL/d ln alpha), pre-allocated - dbeta: ``(total_tokens, HO)`` fp32, pre-allocated + dgate: ``(total_tokens, HO, DK)`` fp32 (dL/d ln alpha), pre-allocated. + With ``safe_gate`` this stays the gradient wrt the TRANSFORMED + log-decay; a host-side helper converts to d(raw gate) afterward + dbeta: ``(total_tokens, HO)`` fp32, or io dtype with + ``use_beta_sigmoid``, pre-allocated. Gradient wrt the post-sigmoid + beta, or wrt the raw logits under ``use_beta_sigmoid`` (the kernel + folds the sigmoid derivative into its own dbeta write) cu_seqlens: ``(num_seqs + 1,)`` int32 scale: attention scale factor initial_state: ``(num_seqs, HO, DK, DV)`` io dtype (KV) -- the state @@ -3692,6 +3938,10 @@ def chunk_kda_bwd_sm100( d_final_state: fp32 ``(num_seqs, HO, DK, DV)`` IN (dL/d final state) use_qk_l2norm_in_kernel: q/k arrive raw; the kernel normalizes for the recompute math and chains the L2-norm backward into dq/dk + safe_gate: interpret ``gate`` through the safe-gate transform + a_log: ``(HO,)`` float32, safe-gate per-head log-amplitude (None = 0) + dt_bias: ``(HO, DK)`` float32, safe-gate channel bias (None = 0) + use_beta_sigmoid: ``beta`` holds logits; sigmoid in-kernel work_items/work_count: split-K table (``common/split_k.py``, REQUIRED; an uncut table row is the whole (b, h) sequence); each item computes chunks ``[wstart, cend)`` backward and writes @@ -3712,6 +3962,10 @@ def chunk_kda_bwd_sm100( if work_items is None or work_count is None: raise ValueError("work_items/work_count are required (the split-table stage builds them for every launch)") dyn_sched = sched_ctr is not None + run_order = order_in_prologue + order_gen = order_in_prologue and work_item_scratch is None + if run_order and sched_all is None: + raise ValueError("order in the prologue requires sched_all (the prologue zeroes both consumers' sched rings)") for name, t in (("state_checkpoints", state_checkpoints),) + ((("initial_state", initial_state),) if use_initial_state else ()): if str(t.dtype).split(".")[-1] != str(q.dtype).split(".")[-1]: raise ValueError(f"{name} dtype must match the io dtype: got {t.dtype} with io {q.dtype}") @@ -3719,6 +3973,13 @@ def chunk_kda_bwd_sm100( if HO % hh != 0: raise ValueError(f"{name}={hh} must divide {HO}") B = cu_seqlens.shape[0] - 1 + gate_scale_log2 = gate_lower_bound * LOG2_E + + if safe_gate and (a_log is None or dt_bias is None): + raise ValueError("safe_gate requires a_log and dt_bias") + if not safe_gate: + a_log = None + dt_bias = None cu_stream = cuda_driver.CUstream(int(stream)) cache = get_compiled_cache( @@ -3730,8 +3991,13 @@ def chunk_kda_bwd_sm100( use_dstate_in, use_dstate0, use_qk_l2norm_in_kernel, + safe_gate, + gate_lower_bound, + use_beta_sigmoid, use_initial_state, dyn_sched, + run_order, + order_gen, ) if "compiled" not in cache: @@ -3741,6 +4007,9 @@ def chunk_kda_bwd_sm100( use_dstate_in=use_dstate_in, use_dstate0=use_dstate0, l2norm=use_qk_l2norm_in_kernel, + safe_gate=safe_gate, + gate_scale_log2=gate_scale_log2, + beta_sigmoid=use_beta_sigmoid, use_initial_state=use_initial_state, q_ratio=HO // HQ, k_ratio=HO // HK, @@ -3765,6 +4034,8 @@ def chunk_kda_bwd_sm100( tensormap_ws_cute = from_dlpack(tensormap_workspace, assumed_align=128).mark_layout_dynamic() + a_log_cute = from_dlpack(a_log, assumed_align=4) if a_log is not None else None + dt_bias_cute = from_dlpack(dt_bias, assumed_align=16) if dt_bias is not None else None beta_cute = from_dlpack(beta, assumed_align=4).mark_layout_dynamic(leading_dim=len(beta.shape) - 1) state_checkpoints_cute = from_dlpack(state_checkpoints, assumed_align=16).mark_layout_dynamic(leading_dim=len(state_checkpoints.shape) - 1) initial_state_cute = ( @@ -3775,6 +4046,8 @@ def chunk_kda_bwd_sm100( cache["compiled"] = cute.compile( host, cfg, + a_log_cute, + dt_bias_cute, beta_cute, state_checkpoints_cute, initial_state_cute, @@ -3792,48 +4065,81 @@ def chunk_kda_bwd_sm100( options="--enable-tvm-ffi --opt-level 2", ) - # ---- per-(batch, head) descriptor arrays: rebuild on input change ------------ - # desc build runs every execute by contract (cu contents are data; - # buffer pointers may change) -- capture-safe, single tiny launch - if "build_descs" not in cache: + if "prologue" not in cache: io_dtype = get_dtype(q.dtype) - - q_bd = from_dlpack(q, assumed_align=16).mark_layout_dynamic(leading_dim=2) - k_bd = from_dlpack(k, assumed_align=16).mark_layout_dynamic(leading_dim=2) - v_bd = from_dlpack(v, assumed_align=16).mark_layout_dynamic(leading_dim=2) - gate_bd = from_dlpack(gate, assumed_align=16).mark_layout_dynamic(leading_dim=2) - do_bd = from_dlpack(do, assumed_align=16).mark_layout_dynamic(leading_dim=2) - dq_bd = from_dlpack(dq, assumed_align=16).mark_layout_dynamic(leading_dim=2) - dk_bd = from_dlpack(dk, assumed_align=16).mark_layout_dynamic(leading_dim=2) - dv_bd = from_dlpack(dv, assumed_align=16).mark_layout_dynamic(leading_dim=2) - dgate_bd = from_dlpack(dgate, assumed_align=16).mark_layout_dynamic(leading_dim=2) - state_checkpoints_bd = from_dlpack(state_checkpoints, assumed_align=16).mark_layout_dynamic(leading_dim=3) - initial_state_bd = from_dlpack(initial_state, assumed_align=16).mark_layout_dynamic(leading_dim=3) if use_initial_state else None - - cu_bd = from_dlpack(cu_seqlens, assumed_align=8).mark_layout_dynamic() - ws_bd = from_dlpack(tensormap_workspace, assumed_align=128).mark_layout_dynamic() - cache["build_descs"] = cute.compile( - build_descs, + q_pl = from_dlpack(q, assumed_align=16).mark_layout_dynamic(leading_dim=2) + k_pl = from_dlpack(k, assumed_align=16).mark_layout_dynamic(leading_dim=2) + v_pl = from_dlpack(v, assumed_align=16).mark_layout_dynamic(leading_dim=2) + gate_pl = from_dlpack(gate, assumed_align=16).mark_layout_dynamic(leading_dim=2) + do_pl = from_dlpack(do, assumed_align=16).mark_layout_dynamic(leading_dim=2) + dq_pl = from_dlpack(dq, assumed_align=16).mark_layout_dynamic(leading_dim=2) + dk_pl = from_dlpack(dk, assumed_align=16).mark_layout_dynamic(leading_dim=2) + dv_pl = from_dlpack(dv, assumed_align=16).mark_layout_dynamic(leading_dim=2) + dgate_pl = from_dlpack(dgate, assumed_align=16).mark_layout_dynamic(leading_dim=2) + state_checkpoints_pl = from_dlpack(state_checkpoints, assumed_align=16).mark_layout_dynamic(leading_dim=3) + initial_state_pl = from_dlpack(initial_state, assumed_align=16).mark_layout_dynamic(leading_dim=3) if use_initial_state else None + cu_pl = from_dlpack(cu_seqlens, assumed_align=8).mark_layout_dynamic() + ws_pl = from_dlpack(tensormap_workspace, assumed_align=128).mark_layout_dynamic() + staging_pl = None + if run_order and not order_gen: + staging_pl = from_dlpack(work_item_scratch, assumed_align=16) + staging_pl.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1), divisibility=1) + work_items_pl = from_dlpack(work_items, assumed_align=16) + work_items_pl.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1), divisibility=1) + work_count_pl = from_dlpack(work_count, assumed_align=4).mark_layout_dynamic() + sched_pl = None + if run_order: + sched_pl = from_dlpack(sched_all, assumed_align=4).mark_layout_dynamic() + cache["prologue"] = cute.compile( + prologue, io_dtype, CFG.B_T, - q_bd, - k_bd, - v_bd, - gate_bd, - do_bd, - dq_bd, - dk_bd, - dv_bd, - dgate_bd, - state_checkpoints_bd, - initial_state_bd, - cu_bd, - ws_bd, + run_order, + order_gen, + run_order, + q_pl, + k_pl, + v_pl, + gate_pl, + do_pl, + dq_pl, + dk_pl, + dv_pl, + dgate_pl, + state_checkpoints_pl, + initial_state_pl, + cu_pl, + staging_pl, + work_count_pl, + work_items_pl, + sched_pl, + ws_pl, cu_stream, options="--enable-tvm-ffi", ) - cache["build_descs"](q, k, v, gate, do, dq, dk, dv, dgate, state_checkpoints, initial_state, cu_seqlens, tensormap_workspace, cu_stream) + cache["prologue"]( + q, + k, + v, + gate, + do, + dq, + dk, + dv, + dgate, + state_checkpoints, + initial_state, + cu_seqlens, + work_item_scratch if run_order else None, + work_count, + work_items, + sched_all if run_order else None, + tensormap_workspace, + cu_stream, + ) cache["compiled"]( + a_log, + dt_bias, beta, state_checkpoints, initial_state, @@ -3849,6 +4155,77 @@ def chunk_kda_bwd_sm100( scale, cu_stream, ) + return cache -# ---- Engine-side helpers: checkpoint-series bounds + entry-0 state seeding ---------------- +def run_bwd( + cache, + q, + k, + v, + gate, + beta, + do, + state_checkpoints, + dq, + dk, + dv, + dgate, + dbeta, + cu_seqlens, + initial_state, + d_initial_state, + d_final_state, + work_items, + work_count, + sched_ctr, + sched_all, + work_item_scratch, + tensormap_workspace, + scale, + stream, + a_log=None, + dt_bias=None, +) -> None: + """Replay the compiled plan: the prologue launch, then the main launch. + The caller owns the contract, which the plan validated at build, so + nothing here raises.""" + cu_stream = cuda_driver.CUstream(int(stream)) + cache["prologue"]( + q, + k, + v, + gate, + do, + dq, + dk, + dv, + dgate, + state_checkpoints, + initial_state, + cu_seqlens, + work_item_scratch, + work_count, + work_items, + sched_all, + tensormap_workspace, + cu_stream, + ) + cache["compiled"]( + a_log, + dt_bias, + beta, + state_checkpoints, + initial_state, + dgate, + dbeta, + cu_seqlens, + d_initial_state, + d_final_state, + work_items, + work_count, + sched_ctr, + tensormap_workspace, + scale, + cu_stream, + ) diff --git a/python/cudnn/linear_attention/frost/kernel/kda_prefill_f16.py b/python/cudnn/linear_attention/frost/kernel/kda_prefill_f16.py index b3819a9fd..4be7b04cc 100644 --- a/python/cudnn/linear_attention/frost/kernel/kda_prefill_f16.py +++ b/python/cudnn/linear_attention/frost/kernel/kda_prefill_f16.py @@ -98,12 +98,13 @@ import cutlass.cute as cute from cutlass.cute.runtime import from_dlpack -from ..common.split_k import decode_work_item +from ..common.split_k import ORDER_CAPACITY, ORDER_ELEMS, ORDER_THREADS, decode_work_item, order_body from ..common.host import get_dtype from cudnn.frost.buffers import current_device_id, data_ptr from cudnn.frost.device import multiprocessor_count from ..common.thd import TENSOR_MAP_QWORDS, emit_checkpoint_seq_descs, emit_seq_descs from .kda_prefill_config import CFG + from cudnn.frost.tile_dsl.barrier import ( advance, MBarrier, @@ -577,14 +578,10 @@ def super_mma_warp( tinv_lo1, tinv_hi1 = f16x2_to_f32(tinv_p1, dtype=cfg.io_dtype) tinv_lo2, tinv_hi2 = f16x2_to_f32(tinv_p2, dtype=cfg.io_dtype) tinv_lo3, tinv_hi3 = f16x2_to_f32(tinv_p3, dtype=cfg.io_dtype) - tinv_acc[0] = tinv_lo0 + upd_acc[0] - tinv_acc[1] = tinv_hi0 + upd_acc[1] - tinv_acc[2] = tinv_lo1 + upd_acc[2] - tinv_acc[3] = tinv_hi1 + upd_acc[3] - tinv_acc[4] = tinv_lo2 + upd_acc[4] - tinv_acc[5] = tinv_hi2 + upd_acc[5] - tinv_acc[6] = tinv_lo3 + upd_acc[6] - tinv_acc[7] = tinv_hi3 + upd_acc[7] + tinv_acc[0], tinv_acc[1] = fadd2(tinv_lo0, tinv_hi0, upd_acc[0], upd_acc[1]) + tinv_acc[2], tinv_acc[3] = fadd2(tinv_lo1, tinv_hi1, upd_acc[2], upd_acc[3]) + tinv_acc[4], tinv_acc[5] = fadd2(tinv_lo2, tinv_hi2, upd_acc[4], upd_acc[5]) + tinv_acc[6], tinv_acc[7] = fadd2(tinv_lo3, tinv_hi3, upd_acc[6], upd_acc[7]) bars.mb_t_inv_done[intermediate_stage].wait(t_inv_free.phase) t_inv_free = advance(t_inv_free, cfg.smem_intermediate_stages) @@ -1990,6 +1987,7 @@ def host( stream, ) -> None: num_sequences = cu_seqlens.shape[0] - 1 + grid_shape = (cfg.max_active_clusters, 1, 1) kernel( cfg, @@ -2500,14 +2498,15 @@ def build_cfg( TENSORMAP_STATIC_SLOTS = 0 -@cute.kernel -def build_all_descs_kernel( - base_q: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_k: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_v: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_gate: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_o: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_checkpoint: cutlass.GridConstant[cuda.tensor_map.TensorMap], +@cute.jit +def build_descs_body( + widx, + base_q, + base_k, + base_v, + base_gate, + base_o, + base_checkpoint, desc_ws: cute.Tensor, cu_seqlens: cute.Tensor, q: cute.Tensor, @@ -2525,10 +2524,9 @@ def build_all_descs_kernel( checkpoint_entry_stride: cutlass.Int32, checkpoint_every_n: cutlass.Int32, ) -> None: - """Single-launch builder for the per-BATCH descriptor arrays (one warp - per array).""" - tidx, _, _ = cute.arch.thread_idx() - widx = cutlass.Int32(tidx) // cutlass.Int32(32) + """Per-batch descriptor-array build, one warp per array. Runs inside the + prologue kernel after its order pass; warps past the array count fall + through the widx guards.""" arr_words = n_batch * cutlass.Int32(TENSOR_MAP_QWORDS) desc_q_arr = cute.make_tensor(desc_ws.iterator, cute.make_layout((arr_words,), stride=(1,))) desc_k_arr = cute.make_tensor(desc_ws.iterator + arr_words, cute.make_layout((arr_words,), stride=(1,))) @@ -2566,10 +2564,100 @@ def build_all_descs_kernel( nvvm.fence_proxy_release(nvvm.MemScope.GPU, from_proxy=nvvm.Proxy.GENERIC, to_proxy=nvvm.Proxy.TENSORMAP) +@cute.kernel +def prologue_kernel( + order_gen: cutlass.Constexpr[bool], + has_sched: cutlass.Constexpr[bool], + b_t: cutlass.Constexpr[int], + base_q: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_k: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_v: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_gate: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_o: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_checkpoint: cutlass.GridConstant[cuda.tensor_map.TensorMap], + desc_ws: cute.Tensor, + cu_seqlens: cute.Tensor, + q: cute.Tensor, + k: cute.Tensor, + v: cute.Tensor, + gate: cute.Tensor, + o: cute.Tensor, + state_checkpoints: cute.Tensor | None, + mStaging: cute.Tensor | None, + mCount: cute.Tensor, + mWorkItems: cute.Tensor, + mSched: cute.Tensor | None, + n_batch: cutlass.Int32, + q_token_stride: cutlass.Int32, + k_token_stride: cutlass.Int32, + v_token_stride: cutlass.Int32, + gate_token_stride: cutlass.Int32, + o_token_stride: cutlass.Int32, + checkpoint_entry_stride: cutlass.Int32, + checkpoint_every_n: cutlass.Int32, +) -> None: + """Single-CTA prologue: LPT-order the work-item table and zero the sched + rings via :func:`order_body`, then build the per-batch TMA-descriptor + arrays via :func:`build_descs_body`, one warp per array (the extra warps + only take part in the order phase).""" + tidx, _, _ = cute.arch.thread_idx() + tidx = cutlass.Int32(tidx) + widx = tidx // cutlass.Int32(32) + sKey = cutlass.Array(cutlass.Int32, ORDER_CAPACITY, space=cutlass.AddressSpace.smem, alignment=16) + sIdx = cutlass.Array(cutlass.Int32, ORDER_CAPACITY, space=cutlass.AddressSpace.smem, alignment=16) + sSpread = cutlass.Array(cutlass.Int32, 2, space=cutlass.AddressSpace.smem, alignment=8) + n_heads_out = cutlass.Int32(gate.shape[1]) + order_body( + order_gen, + has_sched, + b_t, + ORDER_THREADS, + ORDER_ELEMS, + tidx, + n_heads_out, + n_heads_out * n_batch, + cu_seqlens, + mStaging, + mCount, + mWorkItems, + mSched, + sKey, + sIdx, + sSpread, + ) + build_descs_body( + widx, + base_q, + base_k, + base_v, + base_gate, + base_o, + base_checkpoint, + desc_ws, + cu_seqlens, + q, + k, + v, + gate, + o, + state_checkpoints, + n_batch, + q_token_stride, + k_token_stride, + v_token_stride, + gate_token_stride, + o_token_stride, + checkpoint_entry_stride, + checkpoint_every_n, + ) + + @cute.jit -def build_descs( +def prologue( io_dtype: cutlass.Constexpr, b_t: cutlass.Constexpr[int], + order_gen: cutlass.Constexpr[bool], + has_sched: cutlass.Constexpr[bool], q: cute.Tensor, k: cute.Tensor, v: cute.Tensor, @@ -2577,12 +2665,17 @@ def build_descs( o: cute.Tensor, state_checkpoints: cute.Tensor | None, cu_seqlens: cute.Tensor, + work_item_staging: cute.Tensor | None, + work_count: cute.Tensor, + work_items: cute.Tensor, + sched_ctr: cute.Tensor | None, tensormap_workspace: cute.Tensor, checkpoint_every_n: cutlass.Int32, stream: cuda_driver.CUstream, ): - """Build the 6 per-(batch, head) TMA-descriptor arrays (q, k, v, gate, - o, state_checkpoints) into ``tensormap_workspace``.""" + """One-launch prologue: LPT-order the work items and build the 6 + per-(batch, head) TMA-descriptor arrays (q, k, v, gate, o, + state_checkpoints) into ``tensormap_workspace``.""" h_q = q.shape[1] h_k = k.shape[1] h_v = v.shape[1] @@ -2617,8 +2710,10 @@ def build_descs( ), ) base_checkpoint = cuda.create_tensor_map_tiled_from_view(checkpoint_view, box_dims=(tma_granu_elems, d_k, 1, 1), stride_order=(0, 1, 2, 3), swizzle=swz) - n_warps = 6 if state_checkpoints is not None else 5 - build_all_descs_kernel( + prologue_kernel( + order_gen, + has_sched, + b_t, base_q, base_k, base_v, @@ -2633,6 +2728,10 @@ def build_descs( gate, o, state_checkpoints, + work_item_staging, + work_count, + work_items, + sched_ctr, cutlass.Int32(batch_size), cutlass.Int32(q.stride[0]), cutlass.Int32(k.stride[0]), @@ -2641,7 +2740,7 @@ def build_descs( cutlass.Int32(o.stride[0]), cutlass.Int32(state_checkpoints.stride[0] if state_checkpoints is not None else 0), checkpoint_every_n, - ).launch(grid=(1, 1, 1), block=(32 * n_warps, 1, 1), stream=stream) + ).launch(grid=(1, 1, 1), block=(ORDER_THREADS, 1, 1), stream=stream) # ---- Torch adapter / host-side compilation --------------------------------------- @@ -2663,6 +2762,7 @@ def get_compiled_cache( gate_lower_bound: float, beta_sigmoid: bool, dyn_sched: bool, + order_gen: bool, ): """Return a mutable dict that lazily stores the compiled kernel.""" return {} @@ -2770,6 +2870,7 @@ def chunk_kda_sm100( work_items=None, work_count=None, sched_ctr=None, + work_item_scratch=None, *, tensormap_workspace, stream, @@ -2835,6 +2936,7 @@ def chunk_kda_sm100( if work_items is None or work_count is None: raise ValueError("work_items/work_count are required (the split-table stage builds them for every launch)") dyn_sched = sched_ctr is not None + order_gen = work_item_scratch is None if initial_state is not None: state_dtype_src = initial_state.dtype @@ -2873,6 +2975,7 @@ def chunk_kda_sm100( gate_lower_bound, use_beta_sigmoid_in_kernel, dyn_sched, + order_gen, ) if "compiled" not in cache: @@ -2944,40 +3047,51 @@ def chunk_kda_sm100( compiled = cache["compiled"] state_checkpoints_for_descs = output_state_checkpoints if enable_checkpoints else None - # desc build runs every execute by contract (cu contents are data; - # buffer pointers may change) — capture-safe, single tiny launch - if cache.get("build_descs_has_state_checkpoints") != (state_checkpoints_for_descs is not None): - cache.pop("build_descs", None) - cache["build_descs_has_state_checkpoints"] = state_checkpoints_for_descs is not None - if "build_descs" not in cache: + if "prologue" not in cache: io_dtype = get_dtype(q.dtype) - q_bd = from_dlpack(q, assumed_align=16).mark_layout_dynamic(leading_dim=2) - k_bd = from_dlpack(k, assumed_align=16).mark_layout_dynamic(leading_dim=2) - v_bd = from_dlpack(v, assumed_align=16).mark_layout_dynamic(leading_dim=2) - gate_bd = from_dlpack(gate, assumed_align=16).mark_layout_dynamic(leading_dim=2) - o_bd = from_dlpack(output, assumed_align=16).mark_layout_dynamic(leading_dim=2) - cu_bd = from_dlpack(cu_seqlens, assumed_align=8).mark_layout_dynamic() - ws_bd = from_dlpack(tensormap_workspace, assumed_align=128).mark_layout_dynamic() - state_checkpoints_bd = None + q_pl = from_dlpack(q, assumed_align=16).mark_layout_dynamic(leading_dim=2) + k_pl = from_dlpack(k, assumed_align=16).mark_layout_dynamic(leading_dim=2) + v_pl = from_dlpack(v, assumed_align=16).mark_layout_dynamic(leading_dim=2) + gate_pl = from_dlpack(gate, assumed_align=16).mark_layout_dynamic(leading_dim=2) + o_pl = from_dlpack(output, assumed_align=16).mark_layout_dynamic(leading_dim=2) + cu_pl = from_dlpack(cu_seqlens, assumed_align=8).mark_layout_dynamic() + ws_pl = from_dlpack(tensormap_workspace, assumed_align=128).mark_layout_dynamic() + state_checkpoints_pl = None if state_checkpoints_for_descs is not None: - state_checkpoints_bd = from_dlpack(state_checkpoints_for_descs, assumed_align=16).mark_layout_dynamic(leading_dim=3) - cache["build_descs"] = cute.compile( - build_descs, + state_checkpoints_pl = from_dlpack(state_checkpoints_for_descs, assumed_align=16).mark_layout_dynamic(leading_dim=3) + staging_pl = None + if not order_gen: + staging_pl = from_dlpack(work_item_scratch, assumed_align=16) + staging_pl.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1), divisibility=1) + work_items_pl = from_dlpack(work_items, assumed_align=16) + work_items_pl.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1), divisibility=1) + work_count_pl = from_dlpack(work_count, assumed_align=4).mark_layout_dynamic() + sched_pl = None + if dyn_sched: + sched_pl = from_dlpack(sched_ctr, assumed_align=4).mark_layout_dynamic() + cache["prologue"] = cute.compile( + prologue, io_dtype, CFG.B_T, - q_bd, - k_bd, - v_bd, - gate_bd, - o_bd, - state_checkpoints_bd, - cu_bd, - ws_bd, + order_gen, + dyn_sched, + q_pl, + k_pl, + v_pl, + gate_pl, + o_pl, + state_checkpoints_pl, + cu_pl, + staging_pl, + work_count_pl, + work_items_pl, + sched_pl, + ws_pl, cutlass.Int32(checkpoint_every_n_tokens), cu_stream, options="--enable-tvm-ffi", ) - cache["build_descs"]( + cache["prologue"]( q, k, v, @@ -2985,6 +3099,10 @@ def chunk_kda_sm100( output, state_checkpoints_for_descs, cu_seqlens, + work_item_scratch if not order_gen else None, + work_count, + work_items, + sched_ctr, tensormap_workspace, checkpoint_every_n_tokens, cu_stream, @@ -3009,3 +3127,69 @@ def chunk_kda_sm100( scale, cu_stream, ) + return cache + + +def run_prefill( + cache, + q, + k, + v, + gate, + a_log, + dt_bias, + beta, + cu_seqlens, + initial_state, + output, + output_state, + output_state_checkpoints, + work_items, + work_count, + sched_ctr, + work_item_scratch, + tensormap_workspace, + checkpoint_every_n_tokens, + scale, + stream, +) -> None: + """Replay the compiled plan: the prologue launch, then the main launch. + The caller owns the contract, which the plan validated at build, so + nothing here raises.""" + cu_stream = cuda_driver.CUstream(int(stream)) + cache["prologue"]( + q, + k, + v, + gate, + output, + output_state_checkpoints, + cu_seqlens, + work_item_scratch, + work_count, + work_items, + sched_ctr, + tensormap_workspace, + checkpoint_every_n_tokens, + cu_stream, + ) + cache["compiled"]( + q, + k, + v, + gate, + a_log, + dt_bias, + beta, + cu_seqlens, + initial_state, + output, + output_state, + work_items, + work_count, + sched_ctr, + tensormap_workspace, + checkpoint_every_n_tokens, + scale, + cu_stream, + ) diff --git a/python/cudnn/linear_attention/frost/kernel/kda_recompute_f16.py b/python/cudnn/linear_attention/frost/kernel/kda_recompute_f16.py index c78fbf232..50774e7d1 100644 --- a/python/cudnn/linear_attention/frost/kernel/kda_recompute_f16.py +++ b/python/cudnn/linear_attention/frost/kernel/kda_recompute_f16.py @@ -94,12 +94,13 @@ import cutlass.cute as cute from cutlass.cute.runtime import from_dlpack -from ..common.split_k import decode_work_item +from ..common.split_k import ORDER_CAPACITY, ORDER_ELEMS, ORDER_THREADS, decode_work_item, order_body from ..common.host import get_dtype from cudnn.frost.buffers import current_device_id, data_ptr from cudnn.frost.device import multiprocessor_count from ..common.thd import TENSOR_MAP_QWORDS, emit_checkpoint_seq_descs, emit_seq_descs from .kda_recompute_config import CFG + from cudnn.frost.tile_dsl.barrier import ( advance, MBarrier, @@ -530,14 +531,10 @@ def super_mma_warp( tinv_lo1, tinv_hi1 = f16x2_to_f32(tinv_p1, dtype=cfg.io_dtype) tinv_lo2, tinv_hi2 = f16x2_to_f32(tinv_p2, dtype=cfg.io_dtype) tinv_lo3, tinv_hi3 = f16x2_to_f32(tinv_p3, dtype=cfg.io_dtype) - tinv_acc[0] = tinv_lo0 + upd_acc[0] - tinv_acc[1] = tinv_hi0 + upd_acc[1] - tinv_acc[2] = tinv_lo1 + upd_acc[2] - tinv_acc[3] = tinv_hi1 + upd_acc[3] - tinv_acc[4] = tinv_lo2 + upd_acc[4] - tinv_acc[5] = tinv_hi2 + upd_acc[5] - tinv_acc[6] = tinv_lo3 + upd_acc[6] - tinv_acc[7] = tinv_hi3 + upd_acc[7] + tinv_acc[0], tinv_acc[1] = fadd2(tinv_lo0, tinv_hi0, upd_acc[0], upd_acc[1]) + tinv_acc[2], tinv_acc[3] = fadd2(tinv_lo1, tinv_hi1, upd_acc[2], upd_acc[3]) + tinv_acc[4], tinv_acc[5] = fadd2(tinv_lo2, tinv_hi2, upd_acc[4], upd_acc[5]) + tinv_acc[6], tinv_acc[7] = fadd2(tinv_lo3, tinv_hi3, upd_acc[6], upd_acc[7]) bars.mb_t_inv_done[intermediate_stage].wait(t_inv_free.phase) t_inv_free = advance(t_inv_free, cfg.smem_intermediate_stages) @@ -1019,10 +1016,8 @@ def compute0_warp_group( # ---- optional K L2-norm + K_inv staging ------------------------------ if cutlass.const_expr(cfg.l2norm): - kk0_lo = opaque_f32_zero() - kk0_hi = opaque_f32_zero() - kk1_lo = opaque_f32_zero() - kk1_hi = opaque_f32_zero() + kk_lo = opaque_f32_zero() + kk_hi = opaque_f32_zero() for dim_half in cutlass.range_constexpr(2): dim_base = dim_half * (cfg.d_k // 2) + lane_in_row_group * 8 reg_base = dim_half * 8 @@ -1034,15 +1029,15 @@ def compute0_warp_group( for dim_offset in cutlass.range_constexpr(8): k_val = raw_k_vec_f32[dim_offset] raw_k_regs[reg_base + dim_offset] = k_val - if cutlass.const_expr(cfg.l2norm): - if cutlass.const_expr(dim_offset % 2 == 0): - kk0_lo, kk0_hi = ffma2(k_val, k_val, k_val, k_val, kk0_lo, kk0_hi) - else: - kk1_lo, kk1_hi = ffma2(k_val, k_val, k_val, k_val, kk1_lo, kk1_hi) + if cutlass.const_expr(cfg.l2norm): + for dim_pair in cutlass.range_constexpr(4): + k_even = raw_k_vec_f32[2 * dim_pair] + k_odd = raw_k_vec_f32[2 * dim_pair + 1] + kk_lo, kk_hi = ffma2(k_even, k_odd, k_even, k_odd, kk_lo, kk_hi) k_inv_norm = opaque_one if cutlass.const_expr(cfg.l2norm): - k_sum_sq = kk0_hi + kk1_hi + k_sum_sq = kk_lo + kk_hi k_sum_sq = k_sum_sq + cutlass.Float32(nvvm.shfl_sync(0xFFFFFFFF, k_sum_sq, 4, 31, kind=nvvm.Shfl.BFLY)) k_sum_sq = k_sum_sq + cutlass.Float32(nvvm.shfl_sync(0xFFFFFFFF, k_sum_sq, 2, 31, kind=nvvm.Shfl.BFLY)) k_sum_sq = k_sum_sq + cutlass.Float32(nvvm.shfl_sync(0xFFFFFFFF, k_sum_sq, 1, 31, kind=nvvm.Shfl.BFLY)) @@ -2062,12 +2057,13 @@ def build_cfg( TENSORMAP_STATIC_SLOTS = 0 -@cute.kernel -def build_all_descs_kernel( - base_k: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_v: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_gate: cutlass.GridConstant[cuda.tensor_map.TensorMap], - base_checkpoint: cutlass.GridConstant[cuda.tensor_map.TensorMap], +@cute.jit +def build_descs_body( + widx, + base_k, + base_v, + base_gate, + base_checkpoint, desc_ws: cute.Tensor, cu_seqlens: cute.Tensor, k: cute.Tensor, @@ -2081,10 +2077,9 @@ def build_all_descs_kernel( checkpoint_row_stride: cutlass.Int32, checkpoint_every_n: cutlass.Int32, ) -> None: - """Single-launch builder kernel: one warp emits each per-batch TMA - descriptor array.""" - tidx, _, _ = cute.arch.thread_idx() - widx = cutlass.Int32(tidx) // cutlass.Int32(32) + """Per-batch descriptor-array build, one warp per array. Runs inside the + prologue kernel after its order pass; warps past the array count fall + through the widx guards.""" arr_words = n_batch * cutlass.Int32(TENSOR_MAP_QWORDS) desc_words_k = cute.make_tensor(desc_ws.iterator, cute.make_layout((arr_words,), stride=(1,))) desc_words_v = cute.make_tensor(desc_ws.iterator + arr_words, cute.make_layout((arr_words,), stride=(1,))) @@ -2112,20 +2107,107 @@ def build_all_descs_kernel( nvvm.fence_proxy_release(nvvm.MemScope.GPU, from_proxy=nvvm.Proxy.GENERIC, to_proxy=nvvm.Proxy.TENSORMAP) +@cute.kernel +def prologue_kernel( + run_order: cutlass.Constexpr[bool], + order_gen: cutlass.Constexpr[bool], + has_sched: cutlass.Constexpr[bool], + b_t: cutlass.Constexpr[int], + base_k: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_v: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_gate: cutlass.GridConstant[cuda.tensor_map.TensorMap], + base_checkpoint: cutlass.GridConstant[cuda.tensor_map.TensorMap], + desc_ws: cute.Tensor, + cu_seqlens: cute.Tensor, + k: cute.Tensor, + v: cute.Tensor, + gate: cute.Tensor, + state_checkpoints: cute.Tensor | None, + mStaging: cute.Tensor | None, + mCount: cute.Tensor, + mWorkItems: cute.Tensor | None, + mSched: cute.Tensor | None, + n_batch: cutlass.Int32, + k_row_stride: cutlass.Int32, + v_row_stride: cutlass.Int32, + gate_row_stride: cutlass.Int32, + checkpoint_row_stride: cutlass.Int32, + checkpoint_every_n: cutlass.Int32, +) -> None: + """Single-CTA prologue. Under ``run_order`` this kernel is the first + work-item-table consumer, so it LPT-orders the table and zeroes both + consumers' sched rings via :func:`order_body`; it then builds the + per-batch TMA-descriptor arrays via :func:`build_descs_body`, one warp + per array (the extra warps only take part in the order phase).""" + tidx, _, _ = cute.arch.thread_idx() + tidx = cutlass.Int32(tidx) + widx = tidx // cutlass.Int32(32) + if cutlass.const_expr(run_order): + sKey = cutlass.Array(cutlass.Int32, ORDER_CAPACITY, space=cutlass.AddressSpace.smem, alignment=16) + sIdx = cutlass.Array(cutlass.Int32, ORDER_CAPACITY, space=cutlass.AddressSpace.smem, alignment=16) + sSpread = cutlass.Array(cutlass.Int32, 2, space=cutlass.AddressSpace.smem, alignment=8) + n_heads_out = cutlass.Int32(gate.shape[1]) + order_body( + order_gen, + has_sched, + b_t, + ORDER_THREADS, + ORDER_ELEMS, + tidx, + n_heads_out, + n_heads_out * n_batch, + cu_seqlens, + mStaging, + mCount, + mWorkItems, + mSched, + sKey, + sIdx, + sSpread, + ) + build_descs_body( + widx, + base_k, + base_v, + base_gate, + base_checkpoint, + desc_ws, + cu_seqlens, + k, + v, + gate, + state_checkpoints, + n_batch, + k_row_stride, + v_row_stride, + gate_row_stride, + checkpoint_row_stride, + checkpoint_every_n, + ) + + @cute.jit -def build_descs( +def prologue( io_dtype: cutlass.Constexpr, b_t: cutlass.Constexpr[int], + run_order: cutlass.Constexpr[bool], + order_gen: cutlass.Constexpr[bool], + has_sched: cutlass.Constexpr[bool], k: cute.Tensor, v: cute.Tensor, gate: cute.Tensor, state_checkpoints: cute.Tensor | None, cu_seqlens: cute.Tensor, + work_item_staging: cute.Tensor | None, + work_count: cute.Tensor, + work_items: cute.Tensor | None, + sched_all: cute.Tensor | None, tensormap_workspace: cute.Tensor, checkpoint_every_n: cutlass.Int32, stream: cuda_driver.CUstream, ): - """Build the per-batch K/V/Gate/checkpoint TMA-descriptor arrays into + """One-launch prologue: LPT-order the work items (when ``run_order``) and + build the per-batch K/V/Gate/checkpoint TMA-descriptor arrays into ``tensormap_workspace``.""" h_k = k.shape[1] h_v = v.shape[1] @@ -2156,8 +2238,11 @@ def build_descs( ), ) base_checkpoint = cuda.create_tensor_map_tiled_from_view(checkpoint_view, box_dims=(tma_granu_elems, d_k, 1, 1), stride_order=(0, 1, 2, 3), swizzle=swz) - n_warps = 4 if state_checkpoints is not None else 3 - build_all_descs_kernel( + prologue_kernel( + run_order, + order_gen, + has_sched, + b_t, base_k, base_v, base_gate, @@ -2168,13 +2253,17 @@ def build_descs( v, gate, state_checkpoints, + work_item_staging, + work_count, + work_items, + sched_all, cutlass.Int32(batch_size), cutlass.Int32(k.stride[0]), cutlass.Int32(v.stride[0]), cutlass.Int32(gate.stride[0]), cutlass.Int32(state_checkpoints.stride[0] if state_checkpoints is not None else 0), checkpoint_every_n, - ).launch(grid=(1, 1, 1), block=(32 * n_warps, 1, 1), stream=stream) + ).launch(grid=(1, 1, 1), block=(ORDER_THREADS, 1, 1), stream=stream) # ---- Torch adapter / host-side compilation --------------------------------------- @@ -2196,6 +2285,8 @@ def get_compiled_cache( gate_lower_bound: float, beta_sigmoid: bool, dyn_sched: bool, + run_order: bool, + order_gen: bool, ): """Return a mutable dict that lazily stores the compiled kernel.""" return {} @@ -2288,10 +2379,13 @@ def chunk_kda_recompute_sm100( gate_lower_bound: float = DEFAULT_GATE_LOWER_BOUND, a_log=None, dt_bias=None, - use_beta_sigmoid_in_kernel: bool = False, + use_beta_sigmoid: bool = False, work_items=None, work_count=None, sched_ctr=None, + sched_all=None, + work_item_scratch=None, + order_in_prologue: bool = False, *, tensormap_workspace, stream, @@ -2310,7 +2404,7 @@ def chunk_kda_recompute_sm100( ``safe_gate``, which applies the safe-gate transform ``lower_bound * sigmoid(exp(a_log) * (gate + dt_bias))``. beta: ``(total_tokens, HO)``. Post-sigmoid float32, or io-dtype - logits when ``use_beta_sigmoid_in_kernel`` + logits when ``use_beta_sigmoid`` cu_seqlens: ``(num_seqs + 1,)`` int32 initial_state: ``(num_seqs, HO, DK, DV)`` float32/bfloat16, or None output_state: ``(num_seqs, HO, DK, DV)`` float32/bfloat16, or None @@ -2327,7 +2421,7 @@ def chunk_kda_recompute_sm100( safe_gate: interpret ``gate`` through the safe-gate transform a_log: ``(HO,)`` float32, safe-gate per-head log-amplitude (None = 0) dt_bias: ``(HO, DK)`` float32, safe-gate channel bias (None = 0) - use_beta_sigmoid_in_kernel: ``beta`` holds logits; sigmoid in-kernel + use_beta_sigmoid: ``beta`` holds logits; sigmoid in-kernel work_items: ``(max_items, 8)`` int32 work-item table from ``common/split_k.py`` (REQUIRED; an uncut table row is the whole (b, h) sequence). Each item computes chunks ``[cstart, wend)`` @@ -2354,6 +2448,10 @@ def chunk_kda_recompute_sm100( if work_items is None or work_count is None: raise ValueError("work_items/work_count are required (the split-table stage builds them for every launch)") dyn_sched = sched_ctr is not None + run_order = order_in_prologue + order_gen = order_in_prologue and work_item_scratch is None + if run_order and sched_all is None: + raise ValueError("order in the prologue requires sched_all (the prologue zeroes both consumers' sched rings)") if initial_state is not None: state_dtype_src = initial_state.dtype @@ -2389,8 +2487,10 @@ def chunk_kda_recompute_sm100( use_qk_l2norm_in_kernel, safe_gate, gate_lower_bound, - use_beta_sigmoid_in_kernel, + use_beta_sigmoid, dyn_sched, + run_order, + order_gen, ) if "compiled" not in cache: @@ -2431,7 +2531,7 @@ def chunk_kda_recompute_sm100( use_qk_l2norm_in_kernel, safe_gate, gate_scale_log2, - use_beta_sigmoid_in_kernel, + use_beta_sigmoid, k_ratio, v_ratio, HO, @@ -2456,41 +2556,57 @@ def chunk_kda_recompute_sm100( compiled = cache["compiled"] state_checkpoints_for_descs = output_state_checkpoints if enable_checkpoints else None - # desc build runs every execute by contract (cu contents are data; - # buffer pointers may change) - capture-safe, single tiny launch - if cache.get("build_descs_has_state_checkpoints") != (state_checkpoints_for_descs is not None): - cache.pop("build_descs", None) - cache["build_descs_has_state_checkpoints"] = state_checkpoints_for_descs is not None - if "build_descs" not in cache: + if "prologue" not in cache: io_dtype = get_dtype(k.dtype) - k_bd = from_dlpack(k, assumed_align=16).mark_layout_dynamic(leading_dim=2) - v_bd = from_dlpack(v, assumed_align=16).mark_layout_dynamic(leading_dim=2) - gate_bd = from_dlpack(gate, assumed_align=16).mark_layout_dynamic(leading_dim=2) - cu_bd = from_dlpack(cu_seqlens, assumed_align=8).mark_layout_dynamic() - ws_bd = from_dlpack(tensormap_workspace, assumed_align=128).mark_layout_dynamic() - state_checkpoints_bd = None + k_pl = from_dlpack(k, assumed_align=16).mark_layout_dynamic(leading_dim=2) + v_pl = from_dlpack(v, assumed_align=16).mark_layout_dynamic(leading_dim=2) + gate_pl = from_dlpack(gate, assumed_align=16).mark_layout_dynamic(leading_dim=2) + cu_pl = from_dlpack(cu_seqlens, assumed_align=8).mark_layout_dynamic() + ws_pl = from_dlpack(tensormap_workspace, assumed_align=128).mark_layout_dynamic() + state_checkpoints_pl = None if state_checkpoints_for_descs is not None: - state_checkpoints_bd = from_dlpack(state_checkpoints_for_descs, assumed_align=16).mark_layout_dynamic(leading_dim=3) - cache["build_descs"] = cute.compile( - build_descs, + state_checkpoints_pl = from_dlpack(state_checkpoints_for_descs, assumed_align=16).mark_layout_dynamic(leading_dim=3) + staging_pl = None + if run_order and not order_gen: + staging_pl = from_dlpack(work_item_scratch, assumed_align=16) + staging_pl.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1), divisibility=1) + work_items_pl = from_dlpack(work_items, assumed_align=16) + work_items_pl.mark_compact_shape_dynamic(mode=0, stride_order=(0, 1), divisibility=1) + work_count_pl = from_dlpack(work_count, assumed_align=4).mark_layout_dynamic() + sched_pl = None + if run_order: + sched_pl = from_dlpack(sched_all, assumed_align=4).mark_layout_dynamic() + cache["prologue"] = cute.compile( + prologue, io_dtype, CFG.B_T, - k_bd, - v_bd, - gate_bd, - state_checkpoints_bd, - cu_bd, - ws_bd, + run_order, + order_gen, + run_order, + k_pl, + v_pl, + gate_pl, + state_checkpoints_pl, + cu_pl, + staging_pl, + work_count_pl, + work_items_pl, + sched_pl, + ws_pl, cutlass.Int32(checkpoint_every_n_tokens), cu_stream, options="--enable-tvm-ffi", ) - cache["build_descs"]( + cache["prologue"]( k, v, gate, state_checkpoints_for_descs, cu_seqlens, + work_item_scratch if run_order else None, + work_count, + work_items, + sched_all if run_order else None, tensormap_workspace, checkpoint_every_n_tokens, cu_stream, @@ -2512,3 +2628,62 @@ def chunk_kda_recompute_sm100( checkpoint_every_n_tokens, cu_stream, ) + return cache + + +def run_recompute( + cache, + k, + v, + gate, + a_log, + dt_bias, + beta, + cu_seqlens, + initial_state, + output_state, + output_state_checkpoints, + work_items, + work_count, + sched_ctr, + sched_all, + work_item_scratch, + tensormap_workspace, + checkpoint_every_n_tokens, + stream, +) -> None: + """Replay the compiled plan: the prologue launch, then the main launch. + The caller owns the contract, which the plan validated at build, so + nothing here raises.""" + cu_stream = cuda_driver.CUstream(int(stream)) + cache["prologue"]( + k, + v, + gate, + output_state_checkpoints, + cu_seqlens, + work_item_scratch, + work_count, + work_items, + sched_all, + tensormap_workspace, + checkpoint_every_n_tokens, + cu_stream, + ) + cache["compiled"]( + k, + v, + gate, + a_log, + dt_bias, + beta, + cu_seqlens, + initial_state, + output_state, + work_items, + work_count, + sched_ctr, + tensormap_workspace, + checkpoint_every_n_tokens, + cu_stream, + ) diff --git a/python/cudnn/linear_attention/graph_analyzer.py b/python/cudnn/linear_attention/graph_analyzer.py index e3b7fa6a3..2d6ee18f6 100644 --- a/python/cudnn/linear_attention/graph_analyzer.py +++ b/python/cudnn/linear_attention/graph_analyzer.py @@ -8,10 +8,6 @@ family names in ``engines/manifest.py``; PLANNING runs it once per frozen graph and attaches the record, so the family's engines read that same record back instead of each parsing the node. - -Also hosts the engine-side helpers shared by the LA engines: the -check_support gates over the facts record, the execute-time buffer-layout -gate, and the compiled-plan wrapper. """ from __future__ import annotations @@ -20,9 +16,6 @@ from typing import Any, Optional import cudnn -from cudnn.engines.base import CompiledPlan, NodeBuffers, bind_ports -from cudnn.frost import buffers -from cudnn.frost.workspace import Workspace BUFFER_NAME_FROM_CUDNN = { cudnn.data_type.HALF: "float16", @@ -98,6 +91,8 @@ class LaGraphFacts: dbeta_dtype: Any = None dw_dtype: Any = None d_initial_state_dtype: Any = None + d_a_log_dtype: Any = None + d_dt_bias_dtype: Any = None # ports present / requested has_initial_state: bool = False @@ -109,6 +104,7 @@ class LaGraphFacts: use_qk_l2norm: bool = False safe_gate: bool = False use_beta_sigmoid: bool = False + gate_lower_bound: Optional[float] = None checkpoint_every_n_tokens: int = 0 batch_invariant: bool = False @@ -151,6 +147,12 @@ def analyze(graph: "cudnn.pygraph") -> Optional[LaGraphFacts]: invalid = "safe_gate requires a_log and dt_bias inputs" elif not safe_gate and ("a_log" in ins or "dt_bias" in ins): invalid = "a_log/dt_bias require safe_gate=True" + elif ("d_a_log" in outs or "d_dt_bias" in outs) and not (is_bwd and safe_gate): + invalid = "d_a_log/d_dt_bias require safe_gate=True on a bwd node" + elif is_bwd and safe_gate and ("d_a_log" not in outs or "d_dt_bias" not in outs): + invalid = "safe_gate on a bwd node requires the d_a_log and d_dt_bias outputs" + elif is_bwd and safe_gate and any(list(outs[d].dim or []) != list(ins[p].dim or []) for d, p in (("d_a_log", "a_log"), ("d_dt_bias", "dt_bias"))): + invalid = "d_a_log/d_dt_bias dims must match a_log/dt_bias" elif params.get("gate_lower_bound") is not None and not safe_gate: invalid = "gate_lower_bound requires safe_gate=True" elif ckpt < 0: @@ -212,6 +214,8 @@ def analyze(graph: "cudnn.pygraph") -> Optional[LaGraphFacts]: dbeta_dtype=out_dt.get("dBeta"), dw_dtype=out_dt.get("dW"), d_initial_state_dtype=out_dt.get("d_initial_state"), + d_a_log_dtype=out_dt.get("d_a_log"), + d_dt_bias_dtype=out_dt.get("d_dt_bias"), has_initial_state="initial_state" in ins, wants_d_initial_state="d_initial_state" in outs, wants_state_checkpoints="state_checkpoints" in outs, @@ -219,138 +223,7 @@ def analyze(graph: "cudnn.pygraph") -> Optional[LaGraphFacts]: use_qk_l2norm=bool(params.get("use_qk_l2norm", False)), safe_gate=safe_gate, use_beta_sigmoid=bool(params.get("use_beta_sigmoid", False)), + gate_lower_bound=float(params["gate_lower_bound"]) if params.get("gate_lower_bound") is not None else None, checkpoint_every_n_tokens=ckpt, batch_invariant=bool(params.get("batch_invariant", False)), ) - - -# --------------------------------------------------------------------------- -# Engine-side helpers shared by the LA engines -# --------------------------------------------------------------------------- - - -def require(engine: str, port: str, got, want) -> None: - """check_support dtype gate over a facts field: unset passes (the kernel - validates the buffer), anything else must be the kernel-native dtype.""" - if got is None: - return - wanted = want if isinstance(want, tuple) else (want,) - if got not in wanted: - names = "/".join(w.name for w in wanted) - raise NotImplementedError(f"{engine}: '{port}' must be {names} (the kernel-native dtype; no staging), got {got}") - - -def frost_la_gate(engine: str, facts, op: str) -> None: - """The FROST LA engines' shared check_support core: the analyzer record, - the device/DSL environment, and the gates common to all three kernels.""" - if facts is None or facts.op != op: - raise NotImplementedError(f"{engine} supports exactly one {op}/{op}_BWD node") - if facts.invalid: - raise NotImplementedError(f"{engine}: {facts.invalid}") - sm = buffers.current_sm() - if sm is None or not (100 <= sm <= 103): - raise NotImplementedError(f"{engine} requires SM100-SM103 (found {sm})") - installed, version = buffers.cutedsl_state() - if not installed: - raise NotImplementedError(f"{engine} requires the cutedsl extra (nvidia-cutlass-dsl), which is not installed") - if buffers.cutedsl_too_old(version): - want = ".".join(str(v) for v in buffers.CUTEDSL_MIN_VERSION) - raise NotImplementedError(f"{engine} requires nvidia-cutlass-dsl >= {want}; found {version[1]}") - if not facts.uniform_io: - raise NotImplementedError(f"{engine}: q/k/v dtypes must match") - require(engine, "q/k/v", facts.io_dtype, (cudnn.data_type.BFLOAT16, cudnn.data_type.HALF)) - if not facts.thd_layout: - raise NotImplementedError(f"{engine}: q/k/v must be THD [total_T, heads, dim]") - if facts.d_qk != 128 or facts.d_v != 128: - raise NotImplementedError(f"{engine}: head dims must be 128 (the recurrent state is 128x128), got K={facts.d_qk} V={facts.d_v}") - if facts.h_k not in (facts.h_q, facts.h_v): - raise NotImplementedError(f"{engine}: k heads ({facts.h_k}) must match q's ({facts.h_q}) or v's ({facts.h_v}; canonical GQA shares grouped k/v heads)") - if facts.h_v != facts.h_q and max(facts.h_q, facts.h_v) % min(facts.h_q, facts.h_v) != 0: - raise NotImplementedError(f"{engine}: q heads ({facts.h_q}) and v heads ({facts.h_v}) must be equal or one a multiple of the other") - require(engine, "g", facts.g_dtype, cudnn.data_type.FLOAT) - require(engine, "cu_seqlens", facts.cu_dtype, (cudnn.data_type.INT32, cudnn.data_type.INT64)) - - -class FrostLaPlan(CompiledPlan): - """A compiled LA executor, driven from the normalized variant pack: the - port-to-slot join is a property of the graph, so it happens once and is - kept; between executes only the buffer addresses move.""" - - takes_variant_pack = True - - def __init__(self, compiled): - self.compiled = compiled - self.ports = None - - def get_workspace_size(self) -> int: - return self.compiled.workspace_bytes() - - def execute(self, graph, variant_pack, ctx) -> None: - ports = self.ports - if ports is None: - ports = self.ports = bind_ports(graph, variant_pack) - ok, offender = variant_pack.all_dense_layout() - if not ok: - raise ValueError(dense_layout_message(self.compiled.plan_name, ports, offender)) - node_buffers = {} - for node, slots in ports.items(): - names = list(slots.inputs) + list(slots.outputs) - views = variant_pack.operands(list(slots.inputs.values()) + list(slots.outputs.values())) - split = len(slots.inputs) - node_buffers[node] = NodeBuffers(dict(zip(names[:split], views[:split])), dict(zip(names[split:], views[split:]))) - workspace = Workspace.over(variant_pack, self.compiled.workspace_bytes(), type(self.compiled).__name__) - self.compiled(node_buffers, workspace=workspace, stream=ctx.stream) - - -def expect_table(node, align) -> dict: - """Build-time ``{port: (dims, dtype_name, align_bytes)}`` for - :func:`check_layouts`: bound buffers must match the node's frozen - geometry exactly (one graph per shape), and base pointers must satisfy - the kernel entry's ``assumed_align`` claim. ``align`` maps port name -> - bytes (family table read from the entry's from_dlpack calls; absent - ports default to 16).""" - table = {} - for ports in (node.inputs, node.outputs): - for name, t in ports.items(): - if t is None: - continue - dims = tuple(int(d) for d in t.dim) if t.dim else None - table[name] = (dims, BUFFER_NAME_FROM_CUDNN.get(t.get_data_type()), align.get(name, 16)) - return table - - -def check_layouts_compact(plan_name: str, expect, nb) -> None: - """Execute-time gate for the cuTile backend: every bound buffer must be - CONTIGUOUS (the kernels stage rank-merged views and whole-buffer zero - fills), and must match the node's build-time dims/dtype and base - alignment per ``expect`` (see :func:`expect_table`).""" - for ports in (nb.inputs, nb.outputs): - for name, b in ports.items(): - if b is None: - continue - ptr, shape, strides, dtype, _dev = buffers.probe(b) - exp = expect.get(name) if expect else None - if exp is not None: - dims, dtype_name, align = exp - if dims is not None and tuple(shape) != dims: - raise ValueError(f"{plan_name}: buffer for {name!r} must match the graph's build-time dims {dims}; got {tuple(shape)}") - if dtype_name is not None and dtype != dtype_name: - raise ValueError(f"{plan_name}: buffer for {name!r} must be {dtype_name} (the node's declared dtype); got {dtype}") - if align and ptr % align != 0: - raise ValueError(f"{plan_name}: buffer for {name!r} base pointer must be {align}-byte aligned; got 0x{ptr:x}") - if not buffers.is_contiguous(shape, strides): - raise ValueError( - f"{plan_name}: buffer for {name!r} must be contiguous (the cuTile backend stages rank-merged views); got shape {shape} strides {strides}" - ) - - -def dense_layout_message(plan_name, ports, offender) -> str: - """Name the port behind ``all_dense_layout``'s failing slot. Buffers pass - straight to the stride-plumbed kernels, so the one execute-time rule is a - stride-1 innermost dim; this walk only runs on the way to raising.""" - for slots in ports.values(): - for direction in (slots.inputs, slots.outputs): - for port, slot in direction.items(): - if slot == offender: - return f"{plan_name}: buffer for {port!r} must have a stride-1 innermost dim (buffers pass straight to the kernel)" - return f"{plan_name}: the buffer at variant-pack slot {offender} must have a stride-1 innermost dim" diff --git a/python/cudnn/linear_attention/ops/gdn.py b/python/cudnn/linear_attention/ops/gdn.py index 33a3f7666..b5b283f75 100644 --- a/python/cudnn/linear_attention/ops/gdn.py +++ b/python/cudnn/linear_attention/ops/gdn.py @@ -133,6 +133,8 @@ def _make_fprop_cache_key( output_final_state, use_qk_l2norm, batch_invariant, + use_beta_sigmoid, + safe_gate, has_initial_state, ckpt, device, @@ -157,6 +159,8 @@ def _make_fprop_cache_key( bool(output_final_state), bool(use_qk_l2norm), bool(batch_invariant), + bool(use_beta_sigmoid), + bool(safe_gate), bool(has_initial_state), ckpt, device, @@ -185,6 +189,8 @@ def _make_bprop_cache_key( scale, use_qk_l2norm, batch_invariant, + use_beta_sigmoid, + safe_gate, device, plan_name, ): @@ -210,6 +216,8 @@ def _make_bprop_cache_key( float(scale), bool(use_qk_l2norm), bool(batch_invariant), + bool(use_beta_sigmoid), + bool(safe_gate), device, plan_name, ) @@ -221,7 +229,25 @@ def _make_bprop_cache_key( def _build_fprop_graph( - total, N, H, HK, HV, K, V, io_dtype, g_dtype, beta_dtype, state_dtype, cu_dtype, scale, output_final_state, use_qk_l2norm, batch_invariant, ckpt + total, + N, + H, + HK, + HV, + K, + V, + io_dtype, + g_dtype, + beta_dtype, + state_dtype, + cu_dtype, + scale, + output_final_state, + use_qk_l2norm, + batch_invariant, + ckpt, + use_beta_sigmoid=False, + safe_gate=False, ): graph = cudnn.pygraph() HO = max(H, HV) @@ -234,6 +260,11 @@ def _build_fprop_graph( state0_t = None if state_dtype is not None: state0_t = graph.tensor([N, HO, K, V], data_type=state_dtype, name="initial_state") + a_log_t = None + dt_bias_t = None + if safe_gate: + a_log_t = graph.tensor([HO], data_type=cudnn.data_type.FLOAT, name="a_log") + dt_bias_t = graph.tensor([HO], data_type=cudnn.data_type.FLOAT, name="dt_bias") O_t, fs_t, state_checkpoints_t = graph.gdn( q=q_t, k=k_t, @@ -242,14 +273,31 @@ def _build_fprop_graph( beta=beta_t, cu_seqlens=cu_t, initial_state=state0_t, + a_log=a_log_t, + dt_bias=dt_bias_t, scale=scale, output_final_state=output_final_state, use_qk_l2norm=use_qk_l2norm, batch_invariant=batch_invariant, + use_beta_sigmoid=use_beta_sigmoid or None, + safe_gate=safe_gate or None, checkpoint_every_n_tokens=ckpt, name="gdn", ) - return graph, dict(q=q_t, k=k_t, v=v_t, g=g_t, beta=beta_t, cu=cu_t, state0=state0_t, O=O_t, fs=fs_t, state_checkpoints=state_checkpoints_t) + return graph, dict( + q=q_t, + k=k_t, + v=v_t, + g=g_t, + beta=beta_t, + cu=cu_t, + state0=state0_t, + a_log=a_log_t, + dt_bias=dt_bias_t, + O=O_t, + fs=fs_t, + state_checkpoints=state_checkpoints_t, + ) # --------------------------------------------------------------------------- @@ -270,6 +318,10 @@ def _gdn_fwd( output_final_state: bool = False, use_qk_l2norm_in_kernel: bool = False, batch_invariant: bool = False, + use_beta_sigmoid_in_kernel: bool = False, + safe_gate: bool = False, + a_log: Optional[torch.Tensor] = None, + dt_bias: Optional[torch.Tensor] = None, checkpoint_every_n_tokens: int = 0, plan_name: Optional[str] = None, ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: @@ -290,12 +342,31 @@ def _gdn_fwd( raise ValueError(f"gated_delta_net: cu_seqlens must be int32 or int64; got {cu_seqlens.dtype}") cu = cu_seqlens _check_dtype("g", g, torch.float32) - _check_dtype("beta", beta, torch.float32) + if use_beta_sigmoid_in_kernel: + _check_dtype("beta", beta, q.dtype) + else: + _check_dtype("beta", beta, torch.float32) + if safe_gate: + if a_log is None or dt_bias is None: + raise ValueError("gated_delta_net: safe_gate requires a_log and dt_bias") + _check_dtype("a_log", a_log, torch.float32) + _check_dtype("dt_bias", dt_bias, torch.float32) + elif a_log is not None or dt_bias is not None: + raise ValueError("gated_delta_net: a_log/dt_bias require safe_gate=True") if initial_state is not None: _check_dtype("initial_state", initial_state, torch.float32) if initial_state.shape[0] != N: raise ValueError(f"initial_state must carry one state per sequence: got {initial_state.shape[0]} for {N} sequences") - for _name, _t in (("k", k), ("v", v), ("g", g), ("beta", beta), ("cu_seqlens", cu_seqlens), ("initial_state", initial_state)): + for _name, _t in ( + ("k", k), + ("v", v), + ("g", g), + ("beta", beta), + ("cu_seqlens", cu_seqlens), + ("initial_state", initial_state), + ("a_log", a_log), + ("dt_bias", dt_bias), + ): if _t is not None and _t.device != device: raise ValueError(f"gated_delta_net: {_name} must be on q's device ({device}); got {_t.device}") g32 = g @@ -321,6 +392,8 @@ def _gdn_fwd( output_final_state, use_qk_l2norm_in_kernel, batch_invariant, + use_beta_sigmoid_in_kernel, + safe_gate, state0 is not None, ckpt, device, @@ -337,7 +410,7 @@ def _gdn_fwd( V, _torch_dtype_to_cudnn(q.dtype), cudnn.data_type.FLOAT, - cudnn.data_type.FLOAT, + _torch_dtype_to_cudnn(beta.dtype), cudnn.data_type.FLOAT if state0 is not None else None, _torch_dtype_to_cudnn(cu_seqlens.dtype), float(scale), @@ -345,6 +418,8 @@ def _gdn_fwd( bool(use_qk_l2norm_in_kernel), bool(batch_invariant), ckpt, + use_beta_sigmoid=bool(use_beta_sigmoid_in_kernel), + safe_gate=bool(safe_gate), ) select_plan(_fprop_cache[cache_key][0], plan_name) @@ -363,6 +438,9 @@ def _gdn_fwd( } if state0 is not None: variant_pack[t["state0"]] = state0 + if safe_gate: + variant_pack[t["a_log"]] = a_log + variant_pack[t["dt_bias"]] = dt_bias final_state = torch.empty(0, dtype=torch.float32, device=device) if output_final_state: final_state = torch.empty(N, HO, K, V, dtype=torch.float32, device=device) @@ -389,6 +467,10 @@ def _gdn_fwd_fake( output_final_state=False, use_qk_l2norm_in_kernel=False, batch_invariant=False, + use_beta_sigmoid_in_kernel=False, + safe_gate=False, + a_log=None, + dt_bias=None, checkpoint_every_n_tokens=0, plan_name: Optional[str] = None, ): @@ -419,7 +501,25 @@ def _gdn_fwd_fake( def _build_bprop_graph( - total, N, H, HK, HV, K, V, io_dtype, g_dtype, beta_dtype, state_dtype, dstate_in_dtype, cu_dtype, ckpt_rows, scale, use_qk_l2norm, batch_invariant + total, + N, + H, + HK, + HV, + K, + V, + io_dtype, + g_dtype, + beta_dtype, + state_dtype, + dstate_in_dtype, + cu_dtype, + ckpt_rows, + scale, + use_qk_l2norm, + batch_invariant, + use_beta_sigmoid=False, + safe_gate=False, ): graph = cudnn.pygraph() HO = max(H, HV) @@ -439,7 +539,12 @@ def _build_bprop_graph( ckpts_t = None if ckpt_rows is not None: ckpts_t = graph.tensor([ckpt_rows, HO, K, V], data_type=io_dtype, name="state_checkpoints") - dQ_t, dK_t, dV_t, dG_t, dBeta_t, dstate0_t = graph.gdn_bwd( + a_log_t = None + dt_bias_t = None + if safe_gate: + a_log_t = graph.tensor([HO], data_type=cudnn.data_type.FLOAT, name="a_log") + dt_bias_t = graph.tensor([HO], data_type=cudnn.data_type.FLOAT, name="dt_bias") + dQ_t, dK_t, dV_t, dG_t, dBeta_t, dstate0_t, dA_t, dDt_t = graph.gdn_bwd( q=q_t, k=k_t, v=v_t, @@ -450,9 +555,13 @@ def _build_bprop_graph( state_checkpoints=ckpts_t, initial_state=state0_t, d_final_state=dfs_t, + a_log=a_log_t, + dt_bias=dt_bias_t, scale=scale, use_qk_l2norm=use_qk_l2norm, batch_invariant=batch_invariant, + use_beta_sigmoid=use_beta_sigmoid or None, + safe_gate=safe_gate or None, name="gdn_bwd", ) return graph, dict( @@ -465,12 +574,16 @@ def _build_bprop_graph( dO=dO_t, state0=state0_t, dfs=dfs_t, + a_log=a_log_t, + dt_bias=dt_bias_t, dQ=dQ_t, dK=dK_t, dV=dV_t, dG=dG_t, dBeta=dBeta_t, dstate0=dstate0_t, + d_a_log=dA_t, + d_dt_bias=dDt_t, ckpts=ckpts_t, ) @@ -495,15 +608,23 @@ def _gdn_bwd( state_checkpoints: Optional[torch.Tensor] = None, use_qk_l2norm_in_kernel: bool = False, batch_invariant: bool = False, + use_beta_sigmoid_in_kernel: bool = False, + safe_gate: bool = False, + a_log: Optional[torch.Tensor] = None, + dt_bias: Optional[torch.Tensor] = None, plan_name: Optional[str] = None, -) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: """GDN backward (internal): a cached single-node GDN_BWD pygraph, THD layout. ``state_checkpoints`` is the forward's per-chunk state series (io dtype, chunk cadence); when given, the engine consumes it instead of running the checkpoint recompute pass. Returns ``(dq, dk, dv, dg, dbeta, - d_initial_state)``; ``d_initial_state`` is a zero-size tensor when - ``initial_state`` is ``None``. + d_initial_state, d_a_log, d_dt_bias)``; ``d_initial_state`` is a + zero-size tensor when ``initial_state`` is ``None``, and ``d_a_log`` / + ``d_dt_bias`` are zero-size tensors unless ``safe_gate``. With + ``safe_gate``, ``g`` is the raw logits and ``dg`` is the raw-logit + gradient; with ``use_beta_sigmoid_in_kernel``, ``beta`` is io-dtype + logits and ``dbeta`` is the raw-logit gradient. """ total, H, K = q.shape # autograd materializes reduction grads as broadcast (stride-0) @@ -522,7 +643,17 @@ def _gdn_bwd( raise ValueError(f"gated_delta_net: cu_seqlens must be int32 or int64; got {cu_seqlens.dtype}") cu = cu_seqlens _check_dtype("g", g, torch.float32) - _check_dtype("beta", beta, torch.float32) + if use_beta_sigmoid_in_kernel: + _check_dtype("beta", beta, q.dtype) + else: + _check_dtype("beta", beta, torch.float32) + if safe_gate: + if a_log is None or dt_bias is None: + raise ValueError("gated_delta_net: safe_gate requires a_log and dt_bias") + _check_dtype("a_log", a_log, torch.float32) + _check_dtype("dt_bias", dt_bias, torch.float32) + elif a_log is not None or dt_bias is not None: + raise ValueError("gated_delta_net: a_log/dt_bias require safe_gate=True") if initial_state is not None: _check_dtype("initial_state", initial_state, torch.float32) if initial_state.shape[0] != N: @@ -538,6 +669,8 @@ def _gdn_bwd( ("dO", dO), ("d_final_state", d_final_state), ("state_checkpoints", state_checkpoints), + ("a_log", a_log), + ("dt_bias", dt_bias), ): if _t is not None and _t.device != device: raise ValueError(f"gated_delta_net: {_name} must be on q's device ({device}); got {_t.device}") @@ -569,6 +702,8 @@ def _gdn_bwd( scale, use_qk_l2norm_in_kernel, batch_invariant, + use_beta_sigmoid_in_kernel, + safe_gate, device, plan_name, ) @@ -583,7 +718,7 @@ def _gdn_bwd( V, _torch_dtype_to_cudnn(q.dtype), cudnn.data_type.FLOAT, - cudnn.data_type.FLOAT, + _torch_dtype_to_cudnn(beta.dtype), cudnn.data_type.FLOAT if state0 is not None else None, cudnn.data_type.FLOAT if dstate_in is not None else None, _torch_dtype_to_cudnn(cu_seqlens.dtype), @@ -591,6 +726,8 @@ def _gdn_bwd( float(scale), bool(use_qk_l2norm_in_kernel), bool(batch_invariant), + use_beta_sigmoid=bool(use_beta_sigmoid_in_kernel), + safe_gate=bool(safe_gate), ) select_plan(_bprop_cache[cache_key][0], plan_name) @@ -601,7 +738,7 @@ def _gdn_bwd( dk = torch.empty(total, HK, K, dtype=q.dtype, device=device) dv = torch.empty(total, HV, V, dtype=q.dtype, device=device) dg32 = torch.empty(total, HO, dtype=torch.float32, device=device) - dbeta32 = torch.empty(total, HO, dtype=torch.float32, device=device) + dbeta = torch.empty(total, HO, dtype=beta.dtype, device=device) variant_pack = { t["q"]: q, t["k"]: k, @@ -614,7 +751,7 @@ def _gdn_bwd( t["dK"]: dk, t["dV"]: dv, t["dG"]: dg32, - t["dBeta"]: dbeta32, + t["dBeta"]: dbeta, } dstate0 = None if state0 is not None: @@ -625,10 +762,19 @@ def _gdn_bwd( variant_pack[t["dfs"]] = dstate_in if state_checkpoints is not None: variant_pack[t["ckpts"]] = state_checkpoints + d_a_log = torch.empty(0, dtype=torch.float32, device=device) + d_dt_bias = torch.empty(0, dtype=torch.float32, device=device) + if safe_gate: + variant_pack[t["a_log"]] = a_log + variant_pack[t["dt_bias"]] = dt_bias + d_a_log = torch.empty(HO, dtype=torch.float32, device=device) + d_dt_bias = torch.empty(HO, dtype=torch.float32, device=device) + variant_pack[t["d_a_log"]] = d_a_log + variant_pack[t["d_dt_bias"]] = d_dt_bias graph.execute(variant_pack, workspace=_graph_workspace(graph, device), handle=_get_handle(device)) if dstate0 is None: dstate0 = torch.empty(0, dtype=torch.float32, device=device) - return dq, dk, dv, dg32, dbeta32, dstate0 + return dq, dk, dv, dg32, dbeta, dstate0, d_a_log, d_dt_bias @_gdn_bwd.register_fake @@ -646,9 +792,17 @@ def _gdn_bwd_fake( state_checkpoints=None, use_qk_l2norm_in_kernel=False, batch_invariant=False, + use_beta_sigmoid_in_kernel=False, + safe_gate=False, + a_log=None, + dt_bias=None, plan_name=None, ): + if safe_gate and (a_log is None or dt_bias is None): + raise ValueError("gated_delta_net: safe_gate requires a_log and dt_bias") dstate0 = torch.empty_like(initial_state) if initial_state is not None else q.new_empty(0, dtype=torch.float32) + d_a_log = torch.empty_like(a_log) if safe_gate else q.new_empty(0, dtype=torch.float32) + d_dt_bias = torch.empty_like(dt_bias) if safe_gate else q.new_empty(0, dtype=torch.float32) return ( torch.empty_like(q), torch.empty_like(k), @@ -656,6 +810,8 @@ def _gdn_bwd_fake( torch.empty_like(g), torch.empty_like(beta), dstate0, + d_a_log, + d_dt_bias, ) @@ -665,36 +821,60 @@ def _gdn_bwd_fake( def _gdn_setup_context(ctx, inputs, output): - q, k, v, g, beta, cu_seqlens, scale, initial_state, output_final_state, use_qk_l2norm_in_kernel, batch_invariant, checkpoint_every_n_tokens, plan_name = ( - inputs - ) + ( + q, + k, + v, + g, + beta, + cu_seqlens, + scale, + initial_state, + output_final_state, + use_qk_l2norm_in_kernel, + batch_invariant, + use_beta_sigmoid_in_kernel, + safe_gate, + a_log, + dt_bias, + checkpoint_every_n_tokens, + plan_name, + ) = inputs # save_for_backward cannot hold None; keep initial_state as an attribute. + # g/beta are saved as passed: raw logits under safe_gate / use_beta_sigmoid. saved = [q, k, v, g, beta, cu_seqlens] ctx.ckpt_reuse = checkpoint_every_n_tokens == 64 and output[2].numel() > 0 if ctx.ckpt_reuse: saved.append(output[2]) + if safe_gate: + saved.extend([a_log, dt_bias]) ctx.save_for_backward(*saved) ctx.initial_state = initial_state ctx.scale = scale ctx.use_qk_l2norm_in_kernel = use_qk_l2norm_in_kernel ctx.batch_invariant = batch_invariant ctx.plan_name = plan_name + ctx.use_beta_sigmoid_in_kernel = bool(use_beta_sigmoid_in_kernel) + ctx.safe_gate = bool(safe_gate) ctx.set_materialize_grads(False) ctx.mark_non_differentiable(output[2]) def _gdn_backward(ctx, dO, dFinal, _dstate_checkpoints): + a_log = dt_bias = None + if ctx.safe_gate: + a_log, dt_bias = ctx.saved_tensors[-2:] if ctx.ckpt_reuse: - q, k, v, g, beta, cu_seqlens, state_checkpoints = ctx.saved_tensors + q, k, v, g, beta, cu_seqlens, state_checkpoints = ctx.saved_tensors[:7] else: - q, k, v, g, beta, cu_seqlens = ctx.saved_tensors + q, k, v, g, beta, cu_seqlens = ctx.saved_tensors[:6] state_checkpoints = None initial_state = ctx.initial_state if dO is None: dO = torch.zeros(q.shape[0], max(q.shape[1], v.shape[1]), v.shape[2], dtype=q.dtype, device=q.device) dstate_in = dFinal if (dFinal is not None and dFinal.numel() > 0) else None - dq, dk, dv, dg, dbeta, dstate0 = torch.ops.cudnn.gated_delta_net_bwd( + dq, dk, dv, dg, dbeta, dstate0, d_a_log, d_dt_bias = torch.ops.cudnn.gated_delta_net_bwd( dO, q, k, @@ -708,10 +888,15 @@ def _gdn_backward(ctx, dO, dFinal, _dstate_checkpoints): state_checkpoints=state_checkpoints, use_qk_l2norm_in_kernel=ctx.use_qk_l2norm_in_kernel, batch_invariant=ctx.batch_invariant, + use_beta_sigmoid_in_kernel=ctx.use_beta_sigmoid_in_kernel, + safe_gate=ctx.safe_gate, + a_log=a_log, + dt_bias=dt_bias, plan_name=ctx.plan_name, ) # q, k, v, g, beta, cu_seqlens, scale, initial_state, output_final_state, - # use_qk_l2norm_in_kernel, batch_invariant, checkpoint_every_n_tokens, plan_name + # use_qk_l2norm_in_kernel, batch_invariant, use_beta_sigmoid_in_kernel, + # safe_gate, a_log, dt_bias, checkpoint_every_n_tokens, plan_name return ( dq, dk, @@ -726,6 +911,10 @@ def _gdn_backward(ctx, dO, dFinal, _dstate_checkpoints): None, None, None, + d_a_log if ctx.safe_gate else None, + d_dt_bias if ctx.safe_gate else None, + None, + None, ) @@ -753,6 +942,10 @@ def gated_delta_net( output_final_state: bool = False, use_qk_l2norm_in_kernel: bool = False, batch_invariant: bool = False, + use_beta_sigmoid_in_kernel: bool = False, + safe_gate: bool = False, + a_log: Optional[torch.Tensor] = None, + dt_bias: Optional[torch.Tensor] = None, checkpoint_every_n_tokens: int = 0, plan_name: Optional[str] = None, ): @@ -769,13 +962,16 @@ def gated_delta_net( A dense batch of N equal-length sequences is expressed as ``cu_seqlens = [0, T, 2T, ...]`` over the flattened tokens. - Dtypes are kernel-native and strict (callers convert): ``g``, ``beta`` - and the states are float32; ``final_state``, ``dG``, ``dBeta`` and - ``d_initial_state`` are returned in float32. + Dtypes are kernel-native and strict (callers convert): ``g`` and the + states are float32; ``final_state``, ``dG`` and ``d_initial_state`` are + returned in float32. ``beta`` and ``dBeta`` are float32, or io dtype + under ``use_beta_sigmoid_in_kernel``. Args: - g: log-space scalar decay per token (``alpha = exp(g) in (0, 1]``). - beta: per-token write strength. + g: log-space scalar decay per token (``alpha = exp(g) in (0, 1]``), + or raw pre-activation logits when ``safe_gate=True``. + beta: per-token write strength (float32), or io-dtype logits when + ``use_beta_sigmoid_in_kernel=True``. cu_seqlens: ``[N+1]`` int32 sequence boundaries over the packed tokens. scale: attention scale applied to ``q``. Defaults to ``1 / sqrt(K)``. initial_state: optional recurrent state (otherwise zero). @@ -786,6 +982,12 @@ def gated_delta_net( batch_invariant: if ``True``, each sequence's results are bitwise independent of the batch composition (whole-sequence scheduling; disables split-K load balancing). + use_beta_sigmoid_in_kernel: apply ``sigmoid(beta)`` inside the kernel. + safe_gate: interpret ``g`` through the safe-gate transform + ``-exp(a_log) * softplus(g + dt_bias)``. Requires ``a_log`` and + ``dt_bias``. + a_log: ``[HO]`` float32 safe-gate per-head log-amplitude. + dt_bias: ``[HO]`` float32 safe-gate per-head bias. checkpoint_every_n_tokens: if ``> 0``, also return the per-chunk recurrent state series ``state_checkpoints`` (``[total_checkpoints, HO, K, V]`` io dtype, one entry per N tokens strictly before each sequence end; the @@ -816,6 +1018,10 @@ def gated_delta_net( output_final_state=bool(output_final_state), use_qk_l2norm_in_kernel=bool(use_qk_l2norm_in_kernel), batch_invariant=bool(batch_invariant), + use_beta_sigmoid_in_kernel=bool(use_beta_sigmoid_in_kernel), + safe_gate=bool(safe_gate), + a_log=a_log, + dt_bias=dt_bias, checkpoint_every_n_tokens=int(checkpoint_every_n_tokens), plan_name=plan_name, ) diff --git a/python/cudnn/linear_attention/ops/gdn2.py b/python/cudnn/linear_attention/ops/gdn2.py index 33da40d16..bee0a35a5 100644 --- a/python/cudnn/linear_attention/ops/gdn2.py +++ b/python/cudnn/linear_attention/ops/gdn2.py @@ -132,6 +132,7 @@ def _make_fprop_cache_key( output_final_state, use_qk_l2norm, batch_invariant, + use_beta_sigmoid, safe_gate, gate_lower_bound, has_initial_state, @@ -158,6 +159,7 @@ def _make_fprop_cache_key( bool(output_final_state), bool(use_qk_l2norm), bool(batch_invariant), + bool(use_beta_sigmoid), bool(safe_gate), float(gate_lower_bound) if gate_lower_bound is not None else None, bool(has_initial_state), @@ -188,6 +190,9 @@ def _make_bprop_cache_key( scale, use_qk_l2norm, batch_invariant, + use_beta_sigmoid, + safe_gate, + gate_lower_bound, device, plan_name, ): @@ -213,6 +218,9 @@ def _make_bprop_cache_key( float(scale), bool(use_qk_l2norm), bool(batch_invariant), + bool(use_beta_sigmoid), + bool(safe_gate), + float(gate_lower_bound) if gate_lower_bound is not None else None, device, plan_name, ) @@ -243,6 +251,7 @@ def _build_fprop_graph( safe_gate, gate_lower_bound, ckpt, + use_beta_sigmoid=False, ): graph = cudnn.pygraph() HO = max(H, HV) @@ -276,6 +285,7 @@ def _build_fprop_graph( output_final_state=output_final_state, use_qk_l2norm=use_qk_l2norm, batch_invariant=batch_invariant, + use_beta_sigmoid=use_beta_sigmoid or None, safe_gate=safe_gate, gate_lower_bound=gate_lower_bound, checkpoint_every_n_tokens=ckpt, @@ -317,6 +327,7 @@ def _gdn2_fwd( output_final_state: bool = False, use_qk_l2norm_in_kernel: bool = False, batch_invariant: bool = False, + use_beta_sigmoid_in_kernel: bool = False, safe_gate: bool = False, gate_lower_bound: Optional[float] = None, a_log: Optional[torch.Tensor] = None, @@ -392,6 +403,7 @@ def _gdn2_fwd( output_final_state, use_qk_l2norm_in_kernel, batch_invariant, + use_beta_sigmoid_in_kernel, safe_gate, gate_lower_bound, state0 is not None, @@ -420,6 +432,7 @@ def _gdn2_fwd( bool(safe_gate), float(gate_lower_bound) if gate_lower_bound is not None else None, ckpt, + use_beta_sigmoid=bool(use_beta_sigmoid_in_kernel), ) select_plan(_fprop_cache[cache_key][0], plan_name) @@ -468,6 +481,7 @@ def _gdn2_fwd_fake( output_final_state=False, use_qk_l2norm_in_kernel=False, batch_invariant=False, + use_beta_sigmoid_in_kernel=False, safe_gate=False, gate_lower_bound=None, a_log=None, @@ -502,7 +516,26 @@ def _gdn2_fwd_fake( def _build_bprop_graph( - total, N, H, HK, HV, K, V, io_dtype, g_dtype, gate_dtype, state_dtype, dstate_in_dtype, cu_dtype, ckpt_rows, scale, use_qk_l2norm, batch_invariant + total, + N, + H, + HK, + HV, + K, + V, + io_dtype, + g_dtype, + gate_dtype, + state_dtype, + dstate_in_dtype, + cu_dtype, + ckpt_rows, + scale, + use_qk_l2norm, + batch_invariant, + use_beta_sigmoid=False, + safe_gate=False, + gate_lower_bound=None, ): graph = cudnn.pygraph() HO = max(H, HV) @@ -523,7 +556,12 @@ def _build_bprop_graph( ckpts_t = None if ckpt_rows is not None: ckpts_t = graph.tensor([ckpt_rows, HO, K, V], data_type=io_dtype, name="state_checkpoints") - dQ_t, dK_t, dV_t, dG_t, dBeta_t, dW_t, dstate0_t = graph.gdn2_bwd( + a_log_t = None + dt_bias_t = None + if safe_gate: + a_log_t = graph.tensor([HO], data_type=cudnn.data_type.FLOAT, name="a_log") + dt_bias_t = graph.tensor([HO, K], data_type=cudnn.data_type.FLOAT, name="dt_bias") + dQ_t, dK_t, dV_t, dG_t, dBeta_t, dW_t, dstate0_t, dA_t, dDt_t = graph.gdn2_bwd( q=q_t, k=k_t, v=v_t, @@ -535,9 +573,14 @@ def _build_bprop_graph( state_checkpoints=ckpts_t, initial_state=state0_t, d_final_state=dfs_t, + a_log=a_log_t, + dt_bias=dt_bias_t, scale=scale, use_qk_l2norm=use_qk_l2norm, batch_invariant=batch_invariant, + use_beta_sigmoid=use_beta_sigmoid or None, + safe_gate=safe_gate or None, + gate_lower_bound=gate_lower_bound, name="gdn2_bwd", ) return graph, dict( @@ -551,6 +594,8 @@ def _build_bprop_graph( dO=dO_t, state0=state0_t, dfs=dfs_t, + a_log=a_log_t, + dt_bias=dt_bias_t, dQ=dQ_t, dK=dK_t, dV=dV_t, @@ -558,6 +603,8 @@ def _build_bprop_graph( dBeta=dBeta_t, dW=dW_t, dstate0=dstate0_t, + d_a_log=dA_t, + d_dt_bias=dDt_t, ckpts=ckpts_t, ) @@ -583,15 +630,24 @@ def _gdn2_bwd( state_checkpoints: Optional[torch.Tensor] = None, use_qk_l2norm_in_kernel: bool = False, batch_invariant: bool = False, + use_beta_sigmoid_in_kernel: bool = False, + safe_gate: bool = False, + gate_lower_bound: Optional[float] = None, + a_log: Optional[torch.Tensor] = None, + dt_bias: Optional[torch.Tensor] = None, plan_name: Optional[str] = None, -) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: """GDN-2 backward (internal): a cached single-node GDN2_BWD pygraph, THD layout. ``state_checkpoints`` is the forward's per-chunk state series (io dtype, chunk cadence); when given, the engine consumes it instead of running the checkpoint recompute pass. Returns ``(dq, dk, dv, dg, dbeta, dw, - d_initial_state)``; ``d_initial_state`` is a zero-size tensor when - ``initial_state`` is ``None``. + d_initial_state, d_a_log, d_dt_bias)``; ``d_initial_state`` is a + zero-size tensor when ``initial_state`` is ``None``, and ``d_a_log`` / + ``d_dt_bias`` are zero-size tensors unless ``safe_gate``. With + ``safe_gate``, ``g`` is the raw logits and ``dg`` is the raw-logit + gradient; with ``use_beta_sigmoid_in_kernel``, ``beta`` is io-dtype + logits and ``dbeta`` is the raw-logit gradient. """ total, H, K = q.shape # autograd materializes reduction grads as broadcast (stride-0) @@ -613,6 +669,13 @@ def _gdn2_bwd( _check_dtype("g", g, torch.float32) _check_dtype("beta", beta, q.dtype) _check_dtype("w", w, q.dtype) + if safe_gate: + if a_log is None or dt_bias is None: + raise ValueError("gated_delta_net_v2: safe_gate requires a_log and dt_bias") + _check_dtype("a_log", a_log, torch.float32) + _check_dtype("dt_bias", dt_bias, torch.float32) + elif a_log is not None or dt_bias is not None: + raise ValueError("gated_delta_net_v2: a_log/dt_bias require safe_gate=True") if initial_state is not None: _check_dtype("initial_state", initial_state, torch.float32) if initial_state.shape[0] != N: @@ -631,6 +694,8 @@ def _gdn2_bwd( ("dO", dO), ("d_final_state", d_final_state), ("state_checkpoints", state_checkpoints), + ("a_log", a_log), + ("dt_bias", dt_bias), ): if _t is not None and _t.device != device: raise ValueError(f"gated_delta_net_v2: {_name} must be on q's device ({device}); got {_t.device}") @@ -658,6 +723,9 @@ def _gdn2_bwd( scale, use_qk_l2norm_in_kernel, batch_invariant, + use_beta_sigmoid_in_kernel, + safe_gate, + gate_lower_bound, device, plan_name, ) @@ -680,6 +748,9 @@ def _gdn2_bwd( float(scale), bool(use_qk_l2norm_in_kernel), bool(batch_invariant), + use_beta_sigmoid=bool(use_beta_sigmoid_in_kernel), + safe_gate=bool(safe_gate), + gate_lower_bound=float(gate_lower_bound) if gate_lower_bound is not None else None, ) select_plan(_bprop_cache[cache_key][0], plan_name) @@ -716,10 +787,19 @@ def _gdn2_bwd( variant_pack[t["dfs"]] = dstate_in if state_checkpoints is not None: variant_pack[t["ckpts"]] = state_checkpoints + d_a_log = torch.empty(0, dtype=torch.float32, device=device) + d_dt_bias = torch.empty(0, dtype=torch.float32, device=device) + if safe_gate: + variant_pack[t["a_log"]] = a_log + variant_pack[t["dt_bias"]] = dt_bias + d_a_log = torch.empty(HO, dtype=torch.float32, device=device) + d_dt_bias = torch.empty(HO, K, dtype=torch.float32, device=device) + variant_pack[t["d_a_log"]] = d_a_log + variant_pack[t["d_dt_bias"]] = d_dt_bias graph.execute(variant_pack, workspace=_graph_workspace(graph, device), handle=_get_handle(device)) if dstate0 is None: dstate0 = torch.empty(0, dtype=torch.float32, device=device) - return dq, dk, dv, dg, dbeta, dw, dstate0 + return dq, dk, dv, dg, dbeta, dw, dstate0, d_a_log, d_dt_bias @_gdn2_bwd.register_fake @@ -738,9 +818,18 @@ def _gdn2_bwd_fake( state_checkpoints=None, use_qk_l2norm_in_kernel=False, batch_invariant=False, + use_beta_sigmoid_in_kernel=False, + safe_gate=False, + gate_lower_bound=None, + a_log=None, + dt_bias=None, plan_name=None, ): + if safe_gate and (a_log is None or dt_bias is None): + raise ValueError("gated_delta_net_v2: safe_gate requires a_log and dt_bias") dstate0 = torch.empty_like(initial_state) if initial_state is not None else q.new_empty(0, dtype=torch.float32) + d_a_log = torch.empty_like(a_log) if safe_gate else q.new_empty(0, dtype=torch.float32) + d_dt_bias = torch.empty_like(dt_bias) if safe_gate else q.new_empty(0, dtype=torch.float32) return ( torch.empty_like(q), torch.empty_like(k), @@ -749,6 +838,8 @@ def _gdn2_bwd_fake( torch.empty_like(beta), torch.empty_like(w), dstate0, + d_a_log, + d_dt_bias, ) @@ -771,6 +862,7 @@ def _gdn2_setup_context(ctx, inputs, output): output_final_state, use_qk_l2norm_in_kernel, batch_invariant, + use_beta_sigmoid_in_kernel, safe_gate, gate_lower_bound, a_log, @@ -779,35 +871,41 @@ def _gdn2_setup_context(ctx, inputs, output): plan_name, ) = inputs # save_for_backward cannot hold None; keep initial_state as an attribute. + # g/beta are saved as passed: raw logits under safe_gate / use_beta_sigmoid. saved = [q, k, v, g, beta, w, cu_seqlens] ctx.ckpt_reuse = checkpoint_every_n_tokens == 16 and output[2].numel() > 0 if ctx.ckpt_reuse: saved.append(output[2]) + if safe_gate: + saved.extend([a_log, dt_bias]) ctx.save_for_backward(*saved) ctx.initial_state = initial_state ctx.scale = scale ctx.use_qk_l2norm_in_kernel = use_qk_l2norm_in_kernel ctx.batch_invariant = batch_invariant ctx.plan_name = plan_name + ctx.use_beta_sigmoid_in_kernel = bool(use_beta_sigmoid_in_kernel) ctx.safe_gate = bool(safe_gate) + ctx.gate_lower_bound = gate_lower_bound ctx.set_materialize_grads(False) ctx.mark_non_differentiable(output[2]) def _gdn2_backward(ctx, dO, dFinal, _dstate_checkpoints): + a_log = dt_bias = None if ctx.safe_gate: - raise NotImplementedError("gated_delta_net_v2: safe_gate is forward-only (GDN2_BWD takes post-activation gates)") + a_log, dt_bias = ctx.saved_tensors[-2:] if ctx.ckpt_reuse: - q, k, v, g, beta, w, cu_seqlens, state_checkpoints = ctx.saved_tensors + q, k, v, g, beta, w, cu_seqlens, state_checkpoints = ctx.saved_tensors[:8] else: - q, k, v, g, beta, w, cu_seqlens = ctx.saved_tensors + q, k, v, g, beta, w, cu_seqlens = ctx.saved_tensors[:7] state_checkpoints = None initial_state = ctx.initial_state if dO is None: dO = torch.zeros(q.shape[0], max(q.shape[1], v.shape[1]), v.shape[2], dtype=q.dtype, device=q.device) dstate_in = dFinal if (dFinal is not None and dFinal.numel() > 0) else None - dq, dk, dv, dg, dbeta, dw, dstate0 = torch.ops.cudnn.gated_delta_net_v2_bwd( + dq, dk, dv, dg, dbeta, dw, dstate0, d_a_log, d_dt_bias = torch.ops.cudnn.gated_delta_net_v2_bwd( dO, q, k, @@ -822,12 +920,17 @@ def _gdn2_backward(ctx, dO, dFinal, _dstate_checkpoints): state_checkpoints=state_checkpoints, use_qk_l2norm_in_kernel=ctx.use_qk_l2norm_in_kernel, batch_invariant=ctx.batch_invariant, + use_beta_sigmoid_in_kernel=ctx.use_beta_sigmoid_in_kernel, + safe_gate=ctx.safe_gate, + gate_lower_bound=ctx.gate_lower_bound, + a_log=a_log, + dt_bias=dt_bias, plan_name=ctx.plan_name, ) # q, k, v, g, beta, w, cu_seqlens, scale, initial_state, # output_final_state, use_qk_l2norm_in_kernel, batch_invariant, - # safe_gate, gate_lower_bound, a_log, dt_bias, - # checkpoint_every_n_tokens, plan_name + # use_beta_sigmoid_in_kernel, safe_gate, gate_lower_bound, a_log, + # dt_bias, checkpoint_every_n_tokens, plan_name return ( dq, dk, @@ -844,7 +947,8 @@ def _gdn2_backward(ctx, dO, dFinal, _dstate_checkpoints): None, None, None, - None, + d_a_log if ctx.safe_gate else None, + d_dt_bias if ctx.safe_gate else None, None, None, ) @@ -875,6 +979,7 @@ def gated_delta_net_v2( output_final_state: bool = False, use_qk_l2norm_in_kernel: bool = False, batch_invariant: bool = False, + use_beta_sigmoid_in_kernel: bool = False, safe_gate: bool = False, gate_lower_bound: Optional[float] = None, a_log: Optional[torch.Tensor] = None, @@ -899,7 +1004,8 @@ def gated_delta_net_v2( Args: g: per-key-channel log-space decay (``alpha = exp(g)``), or raw pre-activation logits when ``safe_gate=True``. - beta: per-key erase gate (io dtype, post-activation). + beta: per-key erase gate (io dtype, post-activation), or io-dtype + logits when ``use_beta_sigmoid_in_kernel=True``. w: per-value write gate (io dtype, post-activation). cu_seqlens: ``[N+1]`` int32 sequence boundaries over the packed tokens. scale: attention scale applied to ``q``. Defaults to ``1 / sqrt(K)``. @@ -912,9 +1018,12 @@ def gated_delta_net_v2( batch_invariant: if ``True``, each sequence's results are bitwise independent of the batch composition (whole-sequence scheduling; disables split-K load balancing). + use_beta_sigmoid_in_kernel: apply ``sigmoid(beta)`` inside the kernel; + the backward returns the raw-logit beta gradient. safe_gate: interpret ``g`` through the safe-gate transform ``gate_lower_bound * sigmoid(exp(a_log) * (g + dt_bias))``. - Requires ``a_log`` and ``dt_bias``. Forward-only. + Requires ``a_log`` and ``dt_bias``; the backward returns the + raw-logit ``g`` gradient plus ``a_log`` / ``dt_bias`` gradients. gate_lower_bound: safe-gate lower bound in log space (default -5.0). a_log: ``[HO]`` float32 safe-gate per-head log-amplitude. dt_bias: ``[HO, K]`` float32 safe-gate channel bias. @@ -949,6 +1058,7 @@ def gated_delta_net_v2( output_final_state=bool(output_final_state), use_qk_l2norm_in_kernel=bool(use_qk_l2norm_in_kernel), batch_invariant=bool(batch_invariant), + use_beta_sigmoid_in_kernel=bool(use_beta_sigmoid_in_kernel), safe_gate=bool(safe_gate), gate_lower_bound=float(gate_lower_bound) if gate_lower_bound is not None else None, a_log=a_log, diff --git a/python/cudnn/linear_attention/ops/kda.py b/python/cudnn/linear_attention/ops/kda.py index 7362071f2..2bdff166e 100644 --- a/python/cudnn/linear_attention/ops/kda.py +++ b/python/cudnn/linear_attention/ops/kda.py @@ -195,6 +195,9 @@ def _make_bprop_cache_key( scale, use_qk_l2norm, batch_invariant, + use_beta_sigmoid, + safe_gate, + gate_lower_bound, device, plan_name, ): @@ -222,6 +225,9 @@ def _make_bprop_cache_key( float(scale), bool(use_qk_l2norm), bool(batch_invariant), + bool(use_beta_sigmoid), + bool(safe_gate), + float(gate_lower_bound) if gate_lower_bound is not None else None, device, plan_name, ) @@ -509,7 +515,26 @@ def _kda_fwd_fake( def _build_bprop_graph( - total, N, H, HK, HV, K, V, io_dtype, g_dtype, beta_dtype, state_dtype, dstate_in_dtype, cu_dtype, ckpt_rows, scale, use_qk_l2norm, batch_invariant + total, + N, + H, + HK, + HV, + K, + V, + io_dtype, + g_dtype, + beta_dtype, + state_dtype, + dstate_in_dtype, + cu_dtype, + ckpt_rows, + scale, + use_qk_l2norm, + batch_invariant, + use_beta_sigmoid=False, + safe_gate=False, + gate_lower_bound=None, ): graph = cudnn.pygraph() HO = max(H, HV) @@ -529,7 +554,12 @@ def _build_bprop_graph( ckpts_t = None if ckpt_rows is not None: ckpts_t = graph.tensor([ckpt_rows, HO, K, V], data_type=io_dtype, name="state_checkpoints") - dQ_t, dK_t, dV_t, dG_t, dBeta_t, dstate0_t = graph.kda_bwd( + a_log_t = None + dt_bias_t = None + if safe_gate: + a_log_t = graph.tensor([HO], data_type=cudnn.data_type.FLOAT, name="a_log") + dt_bias_t = graph.tensor([HO, K], data_type=cudnn.data_type.FLOAT, name="dt_bias") + dQ_t, dK_t, dV_t, dG_t, dBeta_t, dstate0_t, dA_t, dDt_t = graph.kda_bwd( q=q_t, k=k_t, v=v_t, @@ -540,9 +570,14 @@ def _build_bprop_graph( state_checkpoints=ckpts_t, initial_state=state0_t, d_final_state=dfs_t, + a_log=a_log_t, + dt_bias=dt_bias_t, scale=scale, use_qk_l2norm=use_qk_l2norm, batch_invariant=batch_invariant, + use_beta_sigmoid=use_beta_sigmoid or None, + safe_gate=safe_gate or None, + gate_lower_bound=gate_lower_bound, name="kda_bwd", ) return graph, dict( @@ -555,12 +590,16 @@ def _build_bprop_graph( dO=dO_t, state0=state0_t, dfs=dfs_t, + a_log=a_log_t, + dt_bias=dt_bias_t, dQ=dQ_t, dK=dK_t, dV=dV_t, dG=dG_t, dBeta=dBeta_t, dstate0=dstate0_t, + d_a_log=dA_t, + d_dt_bias=dDt_t, ckpts=ckpts_t, ) @@ -585,15 +624,24 @@ def _kda_bwd( state_checkpoints: Optional[torch.Tensor] = None, use_qk_l2norm_in_kernel: bool = False, batch_invariant: bool = False, + use_beta_sigmoid_in_kernel: bool = False, + safe_gate: bool = False, + gate_lower_bound: Optional[float] = None, + a_log: Optional[torch.Tensor] = None, + dt_bias: Optional[torch.Tensor] = None, plan_name: Optional[str] = None, -) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: """KDA backward (internal): a cached single-node KDA_BWD pygraph, THD layout. ``state_checkpoints`` is the forward's per-chunk state series (io dtype, chunk cadence); when given, the engine consumes it instead of running the checkpoint recompute pass. Returns ``(dq, dk, dv, dg, dbeta, - d_initial_state)``; ``d_initial_state`` is a zero-size tensor when - ``initial_state`` is ``None``. + d_initial_state, d_a_log, d_dt_bias)``; ``d_initial_state`` is a + zero-size tensor when ``initial_state`` is ``None``, and ``d_a_log`` / + ``d_dt_bias`` are zero-size tensors unless ``safe_gate``. With + ``safe_gate``, ``g`` is the raw logits and ``dg`` is the raw-logit + gradient; with ``use_beta_sigmoid_in_kernel``, ``beta`` is io-dtype + logits and ``dbeta`` is the raw-logit gradient. """ total, H, K = q.shape # autograd materializes reduction grads as broadcast (stride-0) @@ -613,7 +661,17 @@ def _kda_bwd( raise ValueError(f"kimi_delta_attention: cu_seqlens must be int32 or int64; got {cu_seqlens.dtype}") cu = cu_seqlens _check_dtype("g", g, torch.float32) - _check_dtype("beta", beta, torch.float32) + if use_beta_sigmoid_in_kernel: + _check_dtype("beta", beta, q.dtype) + else: + _check_dtype("beta", beta, torch.float32) + if safe_gate: + if a_log is None or dt_bias is None: + raise ValueError("kimi_delta_attention: safe_gate requires a_log and dt_bias") + _check_dtype("a_log", a_log, torch.float32) + _check_dtype("dt_bias", dt_bias, torch.float32) + elif a_log is not None or dt_bias is not None: + raise ValueError("kimi_delta_attention: a_log/dt_bias require safe_gate=True") if initial_state is not None: _check_dtype("initial_state", initial_state, torch.float32) if initial_state.shape[0] != N: @@ -631,6 +689,8 @@ def _kda_bwd( ("dO", dO), ("d_final_state", d_final_state), ("state_checkpoints", state_checkpoints), + ("a_log", a_log), + ("dt_bias", dt_bias), ): if _t is not None and _t.device != device: raise ValueError(f"kimi_delta_attention: {_name} must be on q's device ({device}); got {_t.device}") @@ -660,6 +720,9 @@ def _kda_bwd( scale, use_qk_l2norm_in_kernel, batch_invariant, + use_beta_sigmoid_in_kernel, + safe_gate, + gate_lower_bound, device, plan_name, ) @@ -682,6 +745,9 @@ def _kda_bwd( float(scale), bool(use_qk_l2norm_in_kernel), bool(batch_invariant), + use_beta_sigmoid=bool(use_beta_sigmoid_in_kernel), + safe_gate=bool(safe_gate), + gate_lower_bound=float(gate_lower_bound) if gate_lower_bound is not None else None, ) select_plan(_bprop_cache[cache_key][0], plan_name) @@ -715,10 +781,19 @@ def _kda_bwd( variant_pack[t["dfs"]] = dstate_in if state_checkpoints is not None: variant_pack[t["ckpts"]] = state_checkpoints + d_a_log = torch.empty(0, dtype=torch.float32, device=device) + d_dt_bias = torch.empty(0, dtype=torch.float32, device=device) + if safe_gate: + variant_pack[t["a_log"]] = a_log + variant_pack[t["dt_bias"]] = dt_bias + d_a_log = torch.empty(HO, dtype=torch.float32, device=device) + d_dt_bias = torch.empty(HO, K, dtype=torch.float32, device=device) + variant_pack[t["d_a_log"]] = d_a_log + variant_pack[t["d_dt_bias"]] = d_dt_bias graph.execute(variant_pack, workspace=_graph_workspace(graph, device), handle=_get_handle(device)) if dstate0 is None: dstate0 = torch.empty(0, dtype=torch.float32, device=device) - return dq, dk, dv, dg, dbeta, dstate0 + return dq, dk, dv, dg, dbeta, dstate0, d_a_log, d_dt_bias @_kda_bwd.register_fake @@ -736,9 +811,18 @@ def _kda_bwd_fake( state_checkpoints=None, use_qk_l2norm_in_kernel=False, batch_invariant=False, + use_beta_sigmoid_in_kernel=False, + safe_gate=False, + gate_lower_bound=None, + a_log=None, + dt_bias=None, plan_name=None, ): + if safe_gate and (a_log is None or dt_bias is None): + raise ValueError("kimi_delta_attention: safe_gate requires a_log and dt_bias") dstate0 = torch.empty_like(initial_state) if initial_state is not None else q.new_empty(0, dtype=torch.float32) + d_a_log = torch.empty_like(a_log) if safe_gate else q.new_empty(0, dtype=torch.float32) + d_dt_bias = torch.empty_like(dt_bias) if safe_gate else q.new_empty(0, dtype=torch.float32) return ( torch.empty_like(q), torch.empty_like(k), @@ -746,6 +830,8 @@ def _kda_bwd_fake( torch.empty_like(g), torch.empty_like(beta), dstate0, + d_a_log, + d_dt_bias, ) @@ -776,10 +862,13 @@ def _kda_setup_context(ctx, inputs, output): plan_name, ) = inputs # save_for_backward cannot hold None; keep initial_state as an attribute. + # g/beta are saved as passed: raw logits under safe_gate / use_beta_sigmoid. saved = [q, k, v, g, beta, cu_seqlens] ctx.ckpt_reuse = checkpoint_every_n_tokens == 16 and output[2].numel() > 0 if ctx.ckpt_reuse: saved.append(output[2]) + if safe_gate: + saved.extend([a_log, dt_bias]) ctx.save_for_backward(*saved) ctx.initial_state = initial_state ctx.scale = scale @@ -788,24 +877,26 @@ def _kda_setup_context(ctx, inputs, output): ctx.plan_name = plan_name ctx.use_beta_sigmoid_in_kernel = bool(use_beta_sigmoid_in_kernel) ctx.safe_gate = bool(safe_gate) + ctx.gate_lower_bound = gate_lower_bound ctx.set_materialize_grads(False) ctx.mark_non_differentiable(output[2]) def _kda_backward(ctx, dO, dFinal, _dstate_checkpoints): - if ctx.use_beta_sigmoid_in_kernel or ctx.safe_gate: - raise NotImplementedError("kimi_delta_attention: safe_gate/use_beta_sigmoid_in_kernel are forward-only (KDA_BWD takes post-activation gates)") + a_log = dt_bias = None + if ctx.safe_gate: + a_log, dt_bias = ctx.saved_tensors[-2:] if ctx.ckpt_reuse: - q, k, v, g, beta, cu_seqlens, state_checkpoints = ctx.saved_tensors + q, k, v, g, beta, cu_seqlens, state_checkpoints = ctx.saved_tensors[:7] else: - q, k, v, g, beta, cu_seqlens = ctx.saved_tensors + q, k, v, g, beta, cu_seqlens = ctx.saved_tensors[:6] state_checkpoints = None initial_state = ctx.initial_state if dO is None: dO = torch.zeros(q.shape[0], max(q.shape[1], v.shape[1]), v.shape[2], dtype=q.dtype, device=q.device) dstate_in = dFinal if (dFinal is not None and dFinal.numel() > 0) else None - dq, dk, dv, dg, dbeta, dstate0 = torch.ops.cudnn.kimi_delta_attention_bwd( + dq, dk, dv, dg, dbeta, dstate0, d_a_log, d_dt_bias = torch.ops.cudnn.kimi_delta_attention_bwd( dO, q, k, @@ -819,6 +910,11 @@ def _kda_backward(ctx, dO, dFinal, _dstate_checkpoints): state_checkpoints=state_checkpoints, use_qk_l2norm_in_kernel=ctx.use_qk_l2norm_in_kernel, batch_invariant=ctx.batch_invariant, + use_beta_sigmoid_in_kernel=ctx.use_beta_sigmoid_in_kernel, + safe_gate=ctx.safe_gate, + gate_lower_bound=ctx.gate_lower_bound, + a_log=a_log, + dt_bias=dt_bias, plan_name=ctx.plan_name, ) # q, k, v, g, beta, cu_seqlens, scale, initial_state, output_final_state, @@ -840,8 +936,8 @@ def _kda_backward(ctx, dO, dFinal, _dstate_checkpoints): None, None, None, - None, - None, + d_a_log if ctx.safe_gate else None, + d_dt_bias if ctx.safe_gate else None, None, None, ) @@ -893,9 +989,10 @@ def kimi_delta_attention( A dense batch of N equal-length sequences is expressed as ``cu_seqlens = [0, T, 2T, ...]`` over the flattened tokens. - Dtypes are kernel-native and strict (callers convert): ``g``, ``beta`` - and the states are float32; ``final_state``, ``dG``, ``dBeta`` and - ``d_initial_state`` are returned in float32. + Dtypes are kernel-native and strict (callers convert): ``g`` and the + states are float32; ``final_state``, ``dG`` and ``d_initial_state`` are + returned in float32. ``beta`` and ``dBeta`` are float32, or io dtype + under ``use_beta_sigmoid_in_kernel``. Args: g: per-key-channel log-space decay (``alpha = exp(g) in (0, 1]^K``), @@ -913,11 +1010,12 @@ def kimi_delta_attention( batch_invariant: if ``True``, each sequence's results are bitwise independent of the batch composition (whole-sequence scheduling; disables split-K load balancing). - use_beta_sigmoid_in_kernel: apply ``sigmoid(beta)`` inside the kernel. - Forward-only. + use_beta_sigmoid_in_kernel: apply ``sigmoid(beta)`` inside the kernel; + the backward returns the raw-logit beta gradient. safe_gate: interpret ``g`` through the safe-gate transform ``gate_lower_bound * sigmoid(exp(a_log) * (g + dt_bias))``. - Requires ``a_log`` and ``dt_bias``. Forward-only. + Requires ``a_log`` and ``dt_bias``; the backward returns the + raw-logit ``g`` gradient plus ``a_log`` / ``dt_bias`` gradients. gate_lower_bound: safe-gate lower bound in log space (default -5.0). a_log: ``[HO]`` float32 safe-gate per-head log-amplitude. dt_bias: ``[HO, K]`` float32 safe-gate channel bias. diff --git a/test/python/linear_attention/frost/examples/02_gdn_backward.py b/test/python/linear_attention/frost/examples/02_gdn_backward.py index 6cc05443b..2f4338758 100644 --- a/test/python/linear_attention/frost/examples/02_gdn_backward.py +++ b/test/python/linear_attention/frost/examples/02_gdn_backward.py @@ -69,7 +69,7 @@ def main(seq_lens=(192, 320), H: int = 2, D: int = 128) -> None: beta_t = g.tensor([total, H], data_type=cudnn.data_type.FLOAT, name="beta") cu_t = g.tensor([num_seqs + 1], data_type=cudnn.data_type.INT32, name="cu_seqlens") do_t = g.tensor([total, H, D], data_type=cudnn.data_type.BFLOAT16, name="dO") - dQ_t, dK_t, dV_t, dG_t, dBeta_t, _dS0_t = g.gdn_bwd( + dQ_t, dK_t, dV_t, dG_t, dBeta_t, _dS0_t, _dA_t, _dDt_t = g.gdn_bwd( q=q_t, k=k_t, v=v_t, diff --git a/test/python/linear_attention/reference_gdn.py b/test/python/linear_attention/reference_gdn.py index 4df5f1120..3b0b9f6c2 100644 --- a/test/python/linear_attention/reference_gdn.py +++ b/test/python/linear_attention/reference_gdn.py @@ -60,6 +60,10 @@ def gdn_reference( scale: Optional[float] = None, initial_state: Optional[torch.Tensor] = None, cu_seqlens: Optional[torch.Tensor] = None, + safe_gate: bool = False, + a_log: Optional[torch.Tensor] = None, + dt_bias: Optional[torch.Tensor] = None, + use_beta_sigmoid: bool = False, ) -> Tuple[torch.Tensor, torch.Tensor]: """GDN reference. @@ -69,6 +73,10 @@ def gdn_reference( scale: applied to q; defaults to ``1/sqrt(K)``. initial_state: ``[B, HO, K, V]`` (or ``[N, HO, K, V]`` with cu_seqlens). cu_seqlens: packed varlen boundaries (requires B == 1). + safe_gate: treat ``g`` as raw logits and use the log decay + ``-exp(a_log) * softplus(g + dt_bias)`` (differentiable; a_log / + dt_bias are per-head ``[Hg]``). + use_beta_sigmoid: treat ``beta`` as raw logits; apply ``sigmoid``. Returns: ``(o, final_state)`` in fp64: o ``[B, T, HO, V]``, final_state @@ -83,8 +91,13 @@ def gdn_reference( qf = q.double() * scale kf = k.double() vf = v.double() - alphaf = g.double().exp() + gf = g.double() + if safe_gate: + gf = -a_log.double().exp() * torch.nn.functional.softplus(gf + dt_bias.double()) + alphaf = gf.exp() betaf = beta.double() + if use_beta_sigmoid: + betaf = betaf.sigmoid() # expand tensors for grouped heads (view, no copy), as in the sdpa references if q.shape[2] != HO: qf = qf.unsqueeze(3).expand(-1, -1, -1, HO // q.shape[2], -1).reshape(q.shape[0], q.shape[1], HO, -1) diff --git a/test/python/linear_attention/reference_gdn2.py b/test/python/linear_attention/reference_gdn2.py index f91e599a4..1bac23eab 100644 --- a/test/python/linear_attention/reference_gdn2.py +++ b/test/python/linear_attention/reference_gdn2.py @@ -78,6 +78,11 @@ def gdn2_reference( scale: Optional[float] = None, initial_state: Optional[torch.Tensor] = None, cu_seqlens: Optional[torch.Tensor] = None, + safe_gate: bool = False, + gate_lower_bound: Optional[float] = None, + a_log: Optional[torch.Tensor] = None, + dt_bias: Optional[torch.Tensor] = None, + use_beta_sigmoid: bool = False, ) -> Tuple[torch.Tensor, torch.Tensor]: """GDN-2 reference. @@ -90,6 +95,11 @@ def gdn2_reference( initial_state: ``[B, HO, K, V]`` (or ``[N, HO, K, V]`` with cu_seqlens), K-major. cu_seqlens: packed varlen boundaries (requires B == 1). + safe_gate: treat ``g`` as raw logits and use the log decay + ``gate_lower_bound * sigmoid(exp(a_log) * (g + dt_bias))`` + (differentiable; a_log ``[Hg]``, dt_bias ``[Hg, K]``). + gate_lower_bound: safe-gate lower bound in log space (default -5.0). + use_beta_sigmoid: treat ``beta`` as raw logits; apply ``sigmoid``. Returns: ``(o, final_state)`` in fp64: o ``[B, T, HO, V]``, final_state @@ -104,8 +114,14 @@ def gdn2_reference( qf = q.double() * scale kf = k.double() vf = v.double() - alphaf = g.double().exp() # [B, T, HO, K] + gf = g.double() + if safe_gate: + lb = -5.0 if gate_lower_bound is None else float(gate_lower_bound) + gf = lb * torch.sigmoid(a_log.double().exp()[:, None] * (gf + dt_bias.double())) + alphaf = gf.exp() # [B, T, HO, K] betaf = beta.double() # [B, T, HO, K] + if use_beta_sigmoid: + betaf = betaf.sigmoid() wf = w.double() # [B, T, HO, V] # expand tensors for grouped heads (view, no copy), as in the sdpa references if q.shape[2] != HO: diff --git a/test/python/linear_attention/reference_kda.py b/test/python/linear_attention/reference_kda.py index 28ed0e9db..1f0894a4c 100644 --- a/test/python/linear_attention/reference_kda.py +++ b/test/python/linear_attention/reference_kda.py @@ -68,6 +68,11 @@ def kda_reference( scale: Optional[float] = None, initial_state: Optional[torch.Tensor] = None, cu_seqlens: Optional[torch.Tensor] = None, + safe_gate: bool = False, + gate_lower_bound: Optional[float] = None, + a_log: Optional[torch.Tensor] = None, + dt_bias: Optional[torch.Tensor] = None, + use_beta_sigmoid: bool = False, ) -> Tuple[torch.Tensor, torch.Tensor]: """KDA reference. @@ -78,6 +83,11 @@ def kda_reference( scale: applied to q; defaults to ``1/sqrt(K)``. initial_state: ``[B, HO, K, V]`` (or ``[N, HO, K, V]`` with cu_seqlens). cu_seqlens: packed varlen boundaries (requires B == 1). + safe_gate: treat ``g`` as raw logits and use the log decay + ``gate_lower_bound * sigmoid(exp(a_log) * (g + dt_bias))`` + (differentiable; a_log ``[Hg]``, dt_bias ``[Hg, K]``). + gate_lower_bound: safe-gate lower bound in log space (default -5.0). + use_beta_sigmoid: treat ``beta`` as raw logits; apply ``sigmoid``. Returns: ``(o, final_state)`` in fp64: o ``[B, T, HO, V]``, final_state @@ -92,8 +102,14 @@ def kda_reference( qf = q.double() * scale kf = k.double() vf = v.double() - alphaf = g.double().exp() # [B, T, HO, K] + gf = g.double() + if safe_gate: + lb = -5.0 if gate_lower_bound is None else float(gate_lower_bound) + gf = lb * torch.sigmoid(a_log.double().exp()[:, None] * (gf + dt_bias.double())) + alphaf = gf.exp() # [B, T, HO, K] betaf = beta.double() + if use_beta_sigmoid: + betaf = betaf.sigmoid() # expand tensors for grouped heads (view, no copy), as in the sdpa references if q.shape[2] != HO: qf = qf.unsqueeze(3).expand(-1, -1, -1, HO // q.shape[2], -1).reshape(q.shape[0], q.shape[1], HO, -1) diff --git a/test/python/linear_attention/test_la.py b/test/python/linear_attention/test_la.py index 271f8a92a..ebd086a38 100644 --- a/test/python/linear_attention/test_la.py +++ b/test/python/linear_attention/test_la.py @@ -787,55 +787,84 @@ def test_checkpoints_coarse_cadence(backend, variant, ckpt_mult): def safe_gate_case(variant, T=256, H=2, K=128, V=128, seed=SEED + 7): case = make_case(variant, torch.bfloat16, T=T, H=H, K=K, V=V, seed=seed) set_seed(seed + 1) - graw = torch.randn(1, T, case.HO, K, device="cuda", dtype=torch.float32) + if variant == "gdn": + # scalar per-head gate: -exp(a_log[h]) * softplus(g + dt_bias[h]) + graw = torch.randn(1, T, case.HO, device="cuda", dtype=torch.float32) + dt_bias = torch.zeros(case.HO, dtype=torch.float32, device="cuda") + else: + graw = torch.randn(1, T, case.HO, K, device="cuda", dtype=torch.float32) + dt_bias = torch.zeros(case.HO, K, dtype=torch.float32, device="cuda") a_log = torch.zeros(case.HO, dtype=torch.float32, device="cuda") - dt_bias = torch.zeros(case.HO, K, dtype=torch.float32, device="cuda") return case, graw, a_log, dt_bias -@pytest.mark.parametrize("variant", ["kda", "gdn2"]) +@pytest.mark.parametrize("variant", VARIANTS) def test_safe_gate_forward_parity(backend, variant): """Raw logits with a_log = 0 / dt_bias = 0 match the post-activation path - fed the host-side transform ``lb * sigmoid(g)``.""" + fed the host-side transform (``lb * sigmoid(g)``, or ``-softplus(g)`` for + GDN's scalar gate).""" lb = -5.0 case, graw, a_log, dt_bias = safe_gate_case(variant) kw = dict(output_final_state=True, use_qk_l2norm_in_kernel=True) - raw_kw = dict(kw, safe_gate=True, gate_lower_bound=lb, a_log=a_log, dt_bias=dt_bias) + raw_kw = dict(kw, safe_gate=True, a_log=a_log, dt_bias=dt_bias) + if variant != "gdn": + raw_kw["gate_lower_bound"] = lb raw_gates = dict(case.gates, g=graw) - if variant == "kda": + if variant in ("gdn", "kda"): braw = torch.randn(1, case.T, case.HO, device="cuda").to(case.dtype) raw_gates["beta"] = braw raw_kw["use_beta_sigmoid_in_kernel"] = True eff_beta = braw.float().sigmoid() else: eff_beta = case.gates["beta"] + g_eff = -F.softplus(graw) if variant == "gdn" else lb * torch.sigmoid(graw) raw_case = case.clone(gates=raw_gates) - eff_case = case.clone(gates=dict(case.gates, g=lb * torch.sigmoid(graw), beta=eff_beta)) + eff_case = case.clone(gates=dict(case.gates, g=g_eff, beta=eff_beta)) o_raw, fs_raw = run_fwd(backend, raw_case, **raw_kw) o_eff, fs_eff = run_fwd(backend, eff_case, **kw) check("o", o_raw, o_eff.double(), 2e-2) assert rms_ratio(fs_raw, fs_eff) < 2e-2 -@pytest.mark.parametrize("variant", ["kda", "gdn2"]) -def test_safe_gate_backward_raises(backend, variant): - """Raw-logit gate modes are forward-only by contract.""" +@pytest.mark.parametrize("variant", VARIANTS) +def test_safe_gate_backward(backend, variant): + """Fused-gate training: dG comes back in raw-logit space and the + parameter gradients satisfy their exact identities over dG + (d_dt_bias = sum dg_raw; d_a_log = sum dg_raw * (g + dt_bias) per-channel, + or sum dg_raw * softplus(y) / sigmoid(y) for GDN's scalar gate).""" lb = -5.0 case, graw, a_log, dt_bias = safe_gate_case(variant, T=128) + set_seed(SEED + 9) + a_leaf = (torch.randn_like(a_log) * 0.3).requires_grad_(True) + dt_leaf = (torch.randn_like(dt_bias) * 0.3).requires_grad_(True) raw_gates = dict(case.gates, g=graw) - kw = dict(safe_gate=True, gate_lower_bound=lb, a_log=a_log, dt_bias=dt_bias, use_qk_l2norm_in_kernel=True) - if variant == "kda": + kw = dict(safe_gate=True, a_log=a_leaf, dt_bias=dt_leaf, use_qk_l2norm_in_kernel=True) + if variant != "gdn": + kw["gate_lower_bound"] = lb + if variant in ("gdn", "kda"): raw_gates["beta"] = torch.randn(1, case.T, case.HO, device="cuda").to(case.dtype) kw["use_beta_sigmoid_in_kernel"] = True raw_case = case.clone(gates=raw_gates) g_leaf = to_thd(raw_gates["g"]).detach().clone().requires_grad_(True) - args = [to_thd(raw_case.q).detach().clone().requires_grad_(True), to_thd(raw_case.k), to_thd(raw_case.v), g_leaf, to_thd(raw_gates["beta"])] + beta_leaf = to_thd(raw_gates["beta"]).detach().clone().requires_grad_(True) + args = [to_thd(raw_case.q).detach().clone().requires_grad_(True), to_thd(raw_case.k), to_thd(raw_case.v), g_leaf, beta_leaf] if variant == "gdn2": args.append(to_thd(raw_gates["w"])) with waive_unsupported(backend, variant): o, _ = pinned_op(backend, variant)(*args, case.cu, **kw) - with pytest.raises(NotImplementedError, match="forward-only"): o.sum().backward() + dg_raw = g_leaf.grad.double() + ddt_id = dg_raw.sum(0) + if variant == "gdn": + y = g_leaf.detach().double() + dt_leaf.detach().double()[None] + da_id = (dg_raw * (F.softplus(y) / torch.sigmoid(y))).sum(0) + else: + da_id = (dg_raw * (g_leaf.detach().double() + dt_leaf.detach().double()[None])).sum(dim=(0, 2)) + for name, got, ident in (("d_dt_bias", dt_leaf.grad.double(), ddt_id), ("d_a_log", a_leaf.grad.double(), da_id)): + scale = max(ident.abs().max().item(), 1e-6) + assert (got - ident).abs().max().item() / scale < 1e-4, name + for name, leaf in (("dq", args[0]), ("dbeta", beta_leaf)): + assert leaf.grad is not None and bool(torch.isfinite(leaf.grad).all()), name def test_beta_sigmoid_in_kernel(backend): @@ -852,6 +881,57 @@ def test_beta_sigmoid_in_kernel(backend): assert rms_ratio(fs_raw, fs_eff) < 2e-2 +@pytest.mark.parametrize("variant", VARIANTS) +def test_beta_sigmoid_backward(backend, variant): + """The in-kernel Beta sigmoid returns the gradient wrt the raw logit, so + dbeta must equal the post-activation path's dbeta times s * (1 - s) at the + io-rounded s the forward stores.""" + case = make_case(variant, torch.bfloat16, T=256) + set_seed(SEED + 13) + braw = torch.randn_like(case.gates["beta"].float()).to(case.dtype) + s_io = torch.sigmoid(braw.float()).to(case.dtype) + + def dbeta(beta, **kw): + leaf = to_thd(beta).detach().clone().requires_grad_(True) + args = [to_thd(case.q), to_thd(case.k), to_thd(case.v), to_thd(case.gates["g"]), leaf] + if variant == "gdn2": + args.append(to_thd(case.gates["w"])) + with waive_unsupported(backend, variant): + o, _ = pinned_op(backend, variant)(*args, case.cu, **kw) + o.sum().backward() + return leaf.grad.double() + + got = dbeta(braw, use_beta_sigmoid_in_kernel=True) + s = to_thd(s_io).double() + ident = dbeta(s_io.to(case.gates["beta"].dtype)) * s * (1 - s) + scale = ident.abs().max().item() + assert scale > 1e-3, "dbeta is ~0, the comparison would be vacuous" + assert (got - ident).abs().max().item() / scale < 2e-2 + + +@pytest.mark.parametrize("H", (40, 160)) +def test_scalar_gate_head_tiling(backend, H): + """GDN's scalar gate-parameter reduction tiles heads; the dA_log / + ddt_bias identities must hold past one tile and past 128 heads.""" + case, graw, a_log, dt_bias = safe_gate_case("gdn", T=128, H=H) + set_seed(SEED + 15) + a_leaf = (torch.randn_like(a_log) * 0.3).requires_grad_(True) + dt_leaf = (torch.randn_like(dt_bias) * 0.3).requires_grad_(True) + g_leaf = to_thd(graw).detach().clone().requires_grad_(True) + args = [to_thd(case.q), to_thd(case.k), to_thd(case.v), g_leaf, to_thd(case.gates["beta"])] + with waive_unsupported(backend, "gdn"): + o, _ = pinned_op(backend, "gdn")(*args, case.cu, safe_gate=True, a_log=a_leaf, dt_bias=dt_leaf, use_qk_l2norm_in_kernel=True) + o.sum().backward() + dg_raw = g_leaf.grad.double() + y = g_leaf.detach().double() + dt_leaf.detach().double()[None] + for name, got, ident in ( + ("d_dt_bias", dt_leaf.grad.double(), dg_raw.sum(0)), + ("d_a_log", a_leaf.grad.double(), (dg_raw * (F.softplus(y) / torch.sigmoid(y))).sum(0)), + ): + scale = max(ident.abs().max().item(), 1e-6) + assert (got - ident).abs().max().item() / scale < 1e-4, name + + # --------------------------------------------------------------------------- # torch.compile # ---------------------------------------------------------------------------