CUDA kernels for linear attention variants, written in CuTe DSL and CUTLASS C++.
Python
553
66 commits
updated Sep 20, 2026
High-performance CUDA kernels for linear attention variants, written in CuTe DSL and CUTLASS C++.
Linear attention mechanisms reformulate standard attention to use linear-time state updates instead of quadratic pairwise interactions, making them well suited for long-context LLM workloads. Recent variants such as GLA, KDA, GDN, and Lightning Attention further improve expressiveness with gating, delta-style updates, and chunkwise decomposition.
cuLA provides hand-tuned CUDA implementations of these linear attention variants, targeting NVIDIA Blackwell (SM10X) and Hopper (SM90) GPUs. It is designed as a submodule of flash-linear-attention (FLA), sharing the same interface — adopting cuLA requires only a one-line import change. For ease of maintenance, cuLA is currently developed as a standalone library; the end goal is for users to seamlessly access these kernels through FLA. Since FLA already has a kernel dispatch mechanism in place, integration will be ready soon.
⚠️ Early Stage: cuLA is in its early development phase. Many kernels still have significant room for optimization, and the API may evolve. We warmly welcome contributions from the community — whether it's performance tuning, new algorithm support, bug fixes, or architectural improvements. Every contribution helps push the boundaries of linear attention on modern GPUs!
cuLA supports both Hopper (SM90) and Blackwell (SM10X) GPUs.
Requirements (Hopper & Blackwell): Python 3.12+, CUDA Toolkit 12.9+ (SM10X support), NVCC 12.9+, PyTorch 2.9.1+
Note: The PyTorch CUDA version must match your system CUDA Toolkit version. Check with
nvcc --versionandpython -c "import torch; print(torch.version.cuda)".
Pre-built fat-binary wheels (SM90 + SM100 + SM103) are available on GitHub Releases. Linux wheels target manylinux_2_28 and require glibc 2.28 or newer:
pip install "cuda-linear-attention==<VERSION>+<CUDA_TAG>" -f https://github.com/inclusionAI/cuLA/releases/expanded_assets/<TAG>
Replace <TAG> with the release tag (e.g., v0.2.0), <VERSION> with the base version (e.g., 0.2.0), and <CUDA_TAG> with your PyTorch CUDA build tag (e.g., cu129 or cu130). Or download the .whl file directly from the Releases page and install it with pip install <filename>.whl.
Clone cuLA & dependencies:
git clone https://github.com/inclusionAI/cuLA.git
git submodule update --init --recursive
Install PyTorch:
pip install torch==2.9.1 --index-url https://download.pytorch.org/whl/cu129
Install cuLA & dependencies:
# Install flash-linear-attention for benchmark repro
pip install -e third_party/flash-linear-attention
# Install cuLA
pip install -e . --no-build-isolation
Build fat wheel (SM90 + SM100 + SM103):
CULA_BUILD_ALL_ARCHS=1 python -m build --wheel --no-isolation
import torch
from cula.lightning import lightning_attn_fwd
B, T, H, D = 2, 4096, 64, 128
q = torch.randn(B, T, H, D, device="cuda", dtype=torch.bfloat16)
k = torch.randn_like(q)
v = torch.randn_like(q)
decay = torch.linspace(0.0, 0.04, H, device="cuda", dtype=torch.float32)
output, final_state = lightning_attn_fwd(
q,
k,
v,
decay,
scale=D**-0.5,
output_final_state=True,
)
The SM90 CuTe DSL backend supports fixed and packed variable-length prefill, optional recurrent state, GVA head mapping, and persistent packed scheduling. It uses BF16 Q/K/V, FP32 decay/state, head dimension 128, and chunk size 64. See the SM90 Lightning prefill pipeline for the warp-group roles, WGMMA dataflow, pipeline stages, recurrent-state placement, and fixed/packed scheduling modes.
Just change the import:
import torch
from cula.kda import chunk_kda # <-- one-line change from fla.ops.kda
B, T, H, K, V = 2, 2048, 32, 128, 128
device = 'cuda'
q = torch.randn(B, T, H, K, device=device, dtype=torch.bfloat16, requires_grad=True)
k = torch.randn(B, T, H, K, device=device, dtype=torch.bfloat16, requires_grad=True)
v = torch.randn(B, T, H, V, device=device, dtype=torch.bfloat16, requires_grad=True)
g = torch.randn(B, T, H, K, device=device, dtype=torch.bfloat16) * 0.1 # gate (log space)
beta = torch.randn(B, T, H, device=device, dtype=torch.bfloat16).sigmoid()
A_log = torch.randn(H, device=device, dtype=torch.float32) * 0.01
dt_bias = torch.zeros(H * K, device=device, dtype=torch.float32)
init_state = torch.zeros(B, H, K, V, device=device, dtype=torch.float32)
# Forward
o, final_state = chunk_kda(
q=q, k=k, v=v, g=g, beta=beta,
A_log=A_log,
dt_bias=dt_bias,
initial_state=init_state,
output_final_state=True,
use_qk_l2norm_in_kernel=True,
use_gate_in_kernel=True,
safe_gate=True,
lower_bound=-5.0,
)
# Backward
do = torch.randn_like(o)
o.backward(do)
print(f'Output shape: {o.shape}') # [2, 2048, 32, 128]
print(f'Final state shape: {final_state.shape}') # [2, 32, 128, 128]
Notes:
safe_gate=True is required to leverage TensorCore acceleration.beta supports both float32 and bfloat16; initial_state must be float32.cu_seqlens (for variable-length sequences) must be int32.See USAGE.md for detailed usage examples and notes.
Benchmarks run on a single NVIDIA GB200/H200 GPU with PyTorch 2.9.1, Triton 3.5.1.
FLA baseline: flash-linear-attention v0.5.0.
Blackwell (SM10X)
See BENCHMARK_GB200_CUDA_130.md tested with CUDA 13.0 for detailed results.
Hopper (SM90)
See BENCHMARK_H200.md for CuTe DSL FlashKDA results on an H200 141GB with CUDA 12.9.
Highlights:
To reproduce the benchmark suites directly:
# Blackwell (SM10X)
python benchmarks/bench_kda.py --mode both
python benchmarks/bench_lightning_attn_prefill.py --modes no_state varlen
python benchmarks/bench_la_decode_vs_fla.py --heads 64 --head-dim 128
# Hopper (SM90)
python benchmarks/bench_kda_sm90_prefill.py --mode both
python benchmarks/bench_kda_sm90_cp.py
python benchmarks/bench_kda_bwd_wy_dqkg_sm90.py --mode both --heads 32
# Tests for modular KDA forward against FLA Triton implementation
python -m pytest tests/test_kda_sm100_chunk_vs_fla.py -v
# Tests for modular KDA forward against naive KDA reference
python -m pytest tests/test_kda_sm100_chunk_vs_naive.py -v
# Tests for the SM90 CuTeDSL two-kernel prefill + intracard CP (vs FLA)
python -m pytest tests/test_kda_sm90_prefill_vs_fla.py tests/test_kda_sm90_intracard_cp.py -v
# Tests for Lightning Attention prefill on SM100
python tests/test_lightning_sm100_prefill.py
# Tests for the SM90 Lightning public dispatch, semantics, and kernel structure
python -m pytest tests/test_lightning_attn_prefill_dispatch.py tests/test_lightning_attn_prefill_sm90.py -v
# Tests for Lightning Attention decode
python -m pytest tests/test_lightning_decode.py -v
# test_kda_sm100_chunk_vs_naive.py and test_kda_sm100_chunk_vs_fla.py support a fast/slow split.
# Fast (default) — representative correctness paths for default CI and local iteration
python -m pytest tests/test_kda_sm100_chunk_vs_naive.py tests/test_kda_sm100_chunk_vs_fla.py -v
# Slow — broader stress coverage for nightly or manual runs
python -m pytest -m kda_slow tests/test_kda_sm100_chunk_vs_naive.py tests/test_kda_sm100_chunk_vs_fla.py -v
# Full sweep (fast + slow) — run before submitting a PR
python -m pytest -m kda_full tests/test_kda_sm100_chunk_vs_naive.py tests/test_kda_sm100_chunk_vs_fla.py -v
tests/test_kda_e2e_compare_fla.py::test_safe_gate_chunk[B1-T63-H1-D128-...] PASSED
tests/test_kda_e2e_compare_fla.py::test_safe_gate_chunk[B2-T500-H3-D128-...] PASSED
tests/test_kda_e2e_compare_fla.py::test_safe_gate_chunk[B2-T1000-H3-D128-...] PASSED
tests/test_kda_e2e_compare_fla.py::test_safe_gate_chunk[B3-T1024-H4-D128-...] PASSED
tests/test_kda_e2e_compare_fla.py::test_safe_gate_chunk[B4-T1024-H4-D128-...] PASSED
tests/test_kda_e2e_compare_fla.py::test_safe_gate_chunk[B4-T2048-H8-D128-...] PASSED
tests/test_kda_e2e_compare_fla.py::test_safe_gate_chunk_varlen[...] PASSED
...
======================= 17 passed in 40.95s =======================
CUDA kernel tuning is significantly more labor-intensive than Triton — contributions from the open-source community are warmly welcomed!
See REPO_LAYOUT.md for the full directory structure and a summary of each component.
Train
Modular KDA Forward (SM10X, compatible with Kimi CP)
Modular GDN Forward / Backward Kernels (compatible with Kimi CP)
Backward pass optimizations.
Kernel-level compute-communication overlapping for CP linear attention kernels (via nvshmem)
Inference
Lightning prefill kernel (SM90 & SM10X)
Lightning decode kernel (SM90 & SM10X)
Fused KDA prefill kernel (SM90)
Fused KDA prefill kernel (SM10X)
MTP support
More aggressive fusion of small neighboring kernels like cumsum for inference scenarios.
This project is inspired by flash-linear-attention, CUTLASS, CuTe DSL, FlashInfer, Flash-Attention, and FlashMLA. We thank FLA-org and NVIDIA for their great work.
If you find cuLA useful, please cite it using the metadata in our CITATION.cff file:
@software{cula2026,
title = {cuLA: CUDA Linear Attention},
author = {Chaofan Yu, Bowen Zeng, Hao Chen, Zhe Yang, Zhiqiang Zhang, Huan Li and Jun Zhou},
year = {2026},
url = {https://github.com/InclusionAI/cuLA}
}
2,123 followers · starred Apr 2026
32 followers · starred Apr 2026
26 followers · starred Apr 2026
23 followers · starred Apr 2026
Python
84.4%
C++
13.7%
Cuda
1.9%
CUDA kernels for linear attention variants, written in CuTe DSL and CUTLASS C++.
Python
553
66 commits
updated Sep 20, 2026
High-performance CUDA kernels for linear attention variants, written in CuTe DSL and CUTLASS C++.
Linear attention mechanisms reformulate standard attention to use linear-time state updates instead of quadratic pairwise interactions, making them well suited for long-context LLM workloads. Recent variants such as GLA, KDA, GDN, and Lightning Attention further improve expressiveness with gating, delta-style updates, and chunkwise decomposition.
cuLA provides hand-tuned CUDA implementations of these linear attention variants, targeting NVIDIA Blackwell (SM10X) and Hopper (SM90) GPUs. It is designed as a submodule of flash-linear-attention (FLA), sharing the same interface — adopting cuLA requires only a one-line import change. For ease of maintenance, cuLA is currently developed as a standalone library; the end goal is for users to seamlessly access these kernels through FLA. Since FLA already has a kernel dispatch mechanism in place, integration will be ready soon.
⚠️ Early Stage: cuLA is in its early development phase. Many kernels still have significant room for optimization, and the API may evolve. We warmly welcome contributions from the community — whether it's performance tuning, new algorithm support, bug fixes, or architectural improvements. Every contribution helps push the boundaries of linear attention on modern GPUs!
cuLA supports both Hopper (SM90) and Blackwell (SM10X) GPUs.
Requirements (Hopper & Blackwell): Python 3.12+, CUDA Toolkit 12.9+ (SM10X support), NVCC 12.9+, PyTorch 2.9.1+
Note: The PyTorch CUDA version must match your system CUDA Toolkit version. Check with
nvcc --versionandpython -c "import torch; print(torch.version.cuda)".
Pre-built fat-binary wheels (SM90 + SM100 + SM103) are available on GitHub Releases. Linux wheels target manylinux_2_28 and require glibc 2.28 or newer:
pip install "cuda-linear-attention==<VERSION>+<CUDA_TAG>" -f https://github.com/inclusionAI/cuLA/releases/expanded_assets/<TAG>
Replace <TAG> with the release tag (e.g., v0.2.0), <VERSION> with the base version (e.g., 0.2.0), and <CUDA_TAG> with your PyTorch CUDA build tag (e.g., cu129 or cu130). Or download the .whl file directly from the Releases page and install it with pip install <filename>.whl.
Clone cuLA & dependencies:
git clone https://github.com/inclusionAI/cuLA.git
git submodule update --init --recursive
Install PyTorch:
pip install torch==2.9.1 --index-url https://download.pytorch.org/whl/cu129
Install cuLA & dependencies:
# Install flash-linear-attention for benchmark repro
pip install -e third_party/flash-linear-attention
# Install cuLA
pip install -e . --no-build-isolation
Build fat wheel (SM90 + SM100 + SM103):
CULA_BUILD_ALL_ARCHS=1 python -m build --wheel --no-isolation
import torch
from cula.lightning import lightning_attn_fwd
B, T, H, D = 2, 4096, 64, 128
q = torch.randn(B, T, H, D, device="cuda", dtype=torch.bfloat16)
k = torch.randn_like(q)
v = torch.randn_like(q)
decay = torch.linspace(0.0, 0.04, H, device="cuda", dtype=torch.float32)
output, final_state = lightning_attn_fwd(
q,
k,
v,
decay,
scale=D**-0.5,
output_final_state=True,
)
The SM90 CuTe DSL backend supports fixed and packed variable-length prefill, optional recurrent state, GVA head mapping, and persistent packed scheduling. It uses BF16 Q/K/V, FP32 decay/state, head dimension 128, and chunk size 64. See the SM90 Lightning prefill pipeline for the warp-group roles, WGMMA dataflow, pipeline stages, recurrent-state placement, and fixed/packed scheduling modes.
Just change the import:
import torch
from cula.kda import chunk_kda # <-- one-line change from fla.ops.kda
B, T, H, K, V = 2, 2048, 32, 128, 128
device = 'cuda'
q = torch.randn(B, T, H, K, device=device, dtype=torch.bfloat16, requires_grad=True)
k = torch.randn(B, T, H, K, device=device, dtype=torch.bfloat16, requires_grad=True)
v = torch.randn(B, T, H, V, device=device, dtype=torch.bfloat16, requires_grad=True)
g = torch.randn(B, T, H, K, device=device, dtype=torch.bfloat16) * 0.1 # gate (log space)
beta = torch.randn(B, T, H, device=device, dtype=torch.bfloat16).sigmoid()
A_log = torch.randn(H, device=device, dtype=torch.float32) * 0.01
dt_bias = torch.zeros(H * K, device=device, dtype=torch.float32)
init_state = torch.zeros(B, H, K, V, device=device, dtype=torch.float32)
# Forward
o, final_state = chunk_kda(
q=q, k=k, v=v, g=g, beta=beta,
A_log=A_log,
dt_bias=dt_bias,
initial_state=init_state,
output_final_state=True,
use_qk_l2norm_in_kernel=True,
use_gate_in_kernel=True,
safe_gate=True,
lower_bound=-5.0,
)
# Backward
do = torch.randn_like(o)
o.backward(do)
print(f'Output shape: {o.shape}') # [2, 2048, 32, 128]
print(f'Final state shape: {final_state.shape}') # [2, 32, 128, 128]
Notes:
safe_gate=True is required to leverage TensorCore acceleration.beta supports both float32 and bfloat16; initial_state must be float32.cu_seqlens (for variable-length sequences) must be int32.See USAGE.md for detailed usage examples and notes.
Benchmarks run on a single NVIDIA GB200/H200 GPU with PyTorch 2.9.1, Triton 3.5.1.
FLA baseline: flash-linear-attention v0.5.0.
Blackwell (SM10X)
See BENCHMARK_GB200_CUDA_130.md tested with CUDA 13.0 for detailed results.
Hopper (SM90)
See BENCHMARK_H200.md for CuTe DSL FlashKDA results on an H200 141GB with CUDA 12.9.
Highlights:
To reproduce the benchmark suites directly:
# Blackwell (SM10X)
python benchmarks/bench_kda.py --mode both
python benchmarks/bench_lightning_attn_prefill.py --modes no_state varlen
python benchmarks/bench_la_decode_vs_fla.py --heads 64 --head-dim 128
# Hopper (SM90)
python benchmarks/bench_kda_sm90_prefill.py --mode both
python benchmarks/bench_kda_sm90_cp.py
python benchmarks/bench_kda_bwd_wy_dqkg_sm90.py --mode both --heads 32
# Tests for modular KDA forward against FLA Triton implementation
python -m pytest tests/test_kda_sm100_chunk_vs_fla.py -v
# Tests for modular KDA forward against naive KDA reference
python -m pytest tests/test_kda_sm100_chunk_vs_naive.py -v
# Tests for the SM90 CuTeDSL two-kernel prefill + intracard CP (vs FLA)
python -m pytest tests/test_kda_sm90_prefill_vs_fla.py tests/test_kda_sm90_intracard_cp.py -v
# Tests for Lightning Attention prefill on SM100
python tests/test_lightning_sm100_prefill.py
# Tests for the SM90 Lightning public dispatch, semantics, and kernel structure
python -m pytest tests/test_lightning_attn_prefill_dispatch.py tests/test_lightning_attn_prefill_sm90.py -v
# Tests for Lightning Attention decode
python -m pytest tests/test_lightning_decode.py -v
# test_kda_sm100_chunk_vs_naive.py and test_kda_sm100_chunk_vs_fla.py support a fast/slow split.
# Fast (default) — representative correctness paths for default CI and local iteration
python -m pytest tests/test_kda_sm100_chunk_vs_naive.py tests/test_kda_sm100_chunk_vs_fla.py -v
# Slow — broader stress coverage for nightly or manual runs
python -m pytest -m kda_slow tests/test_kda_sm100_chunk_vs_naive.py tests/test_kda_sm100_chunk_vs_fla.py -v
# Full sweep (fast + slow) — run before submitting a PR
python -m pytest -m kda_full tests/test_kda_sm100_chunk_vs_naive.py tests/test_kda_sm100_chunk_vs_fla.py -v
tests/test_kda_e2e_compare_fla.py::test_safe_gate_chunk[B1-T63-H1-D128-...] PASSED
tests/test_kda_e2e_compare_fla.py::test_safe_gate_chunk[B2-T500-H3-D128-...] PASSED
tests/test_kda_e2e_compare_fla.py::test_safe_gate_chunk[B2-T1000-H3-D128-...] PASSED
tests/test_kda_e2e_compare_fla.py::test_safe_gate_chunk[B3-T1024-H4-D128-...] PASSED
tests/test_kda_e2e_compare_fla.py::test_safe_gate_chunk[B4-T1024-H4-D128-...] PASSED
tests/test_kda_e2e_compare_fla.py::test_safe_gate_chunk[B4-T2048-H8-D128-...] PASSED
tests/test_kda_e2e_compare_fla.py::test_safe_gate_chunk_varlen[...] PASSED
...
======================= 17 passed in 40.95s =======================
CUDA kernel tuning is significantly more labor-intensive than Triton — contributions from the open-source community are warmly welcomed!
See REPO_LAYOUT.md for the full directory structure and a summary of each component.
Train
Modular KDA Forward (SM10X, compatible with Kimi CP)
Modular GDN Forward / Backward Kernels (compatible with Kimi CP)
Backward pass optimizations.
Kernel-level compute-communication overlapping for CP linear attention kernels (via nvshmem)
Inference
Lightning prefill kernel (SM90 & SM10X)
Lightning decode kernel (SM90 & SM10X)
Fused KDA prefill kernel (SM90)
Fused KDA prefill kernel (SM10X)
MTP support
More aggressive fusion of small neighboring kernels like cumsum for inference scenarios.
This project is inspired by flash-linear-attention, CUTLASS, CuTe DSL, FlashInfer, Flash-Attention, and FlashMLA. We thank FLA-org and NVIDIA for their great work.
If you find cuLA useful, please cite it using the metadata in our CITATION.cff file:
@software{cula2026,
title = {cuLA: CUDA Linear Attention},
author = {Chaofan Yu, Bowen Zeng, Hao Chen, Zhe Yang, Zhiqiang Zhang, Huan Li and Jun Zhou},
year = {2026},
url = {https://github.com/InclusionAI/cuLA}
}
2,123 followers · starred Apr 2026
32 followers · starred Apr 2026
26 followers · starred Apr 2026
23 followers · starred Apr 2026
Python
84.4%
C++
13.7%
Cuda
1.9%