Native Sparse Attention (NSA)
Native Sparse Attention (NSA)
This is an experimental API and subject to change.
Overview
The Native Sparse Attention (NSA) module implements the sparse attention mechanism described in Native Sparse Attention: Hardware-Aligned and Natively Trainable Sparse Attention. NSA provides high-performance sparse attention kernels optimized for Blackwell (SM100+) GPUs. The Selection, Compression, and Top-K components are implemented with CUTLASS/CUTE. Sliding Window Attention currently utilizes cuDNN backend.
NSA combines multiple attention strategies to efficiently process long sequences:
- Selection Attention: Attends to dynamically selected important blocks across the full context
- Compression Attention: Attends to compressed key-value representations for global context
- Sliding Window Attention: Attends to a local sliding window for fine-grained local context
- Top-K Reduction: Identifies the most important key-value blocks for selection attention
Architecture
Each component can be used independently or combined for the full NSA pipeline.
Installation
Install the cuDNN Frontend package; the CuTe DSL runtime it JITs through is a required dependency and comes with it:
API Usage
NSA Namespace
All NSA components are accessible through the NSA namespace:
Components
1. Selection Attention
Selection Attention performs attention on dynamically selected key-value blocks. Given pre-computed block indices (typically from Top-K Reduction), it efficiently attends only to the most relevant parts of the context.
Shapes
-
Inputs
Q(Query):(T, H_q, D)whereTis total sequence length,H_qis number of query heads,Dis head dimensionK(Key):(T, H_kv, D)whereH_kvis number of key-value headsV(Value):(T, H_kv, D_v)whereD_vis value dimensionblock_indices:(T, H_kv, K)– indices of selected blocks for each query positionblock_counts:(T, H_kv)– number of valid blocks per query positioncum_seqlen_q:(batch_size + 1,)– cumulative sequence lengths for queriescum_seqlen_k:(batch_size + 1,)– cumulative sequence lengths for keys (must equalcum_seqlen_q)
-
Outputs
O(Output):(T, H_q, D_v)L(LogSumExp):(T, H_q)M(Max):(T, H_q)
Equation
For each query position attending to selected blocks :
High-level Wrapper
Class API
Parameters
Constraints
- Input dtype must be
float16orbfloat16 H_qmust be divisible byH_kv(supports GQA/MQA)- Currently only supports
T,H,Dlayout (variable-length batched sequences) cum_seqlen_qandcum_seqlen_kmust be identical- Requires SM90+ (Hopper or newer)
2. Compression Attention
Compression Attention performs attention over compressed key-value sequences. This is useful for maintaining a global view of the context with reduced memory and computation.
Shapes
-
Inputs
Q(Query):(B, H_q, S_q, D)or(T, H_q, D)K(Key):(B, H_kv, S_kv, D)or(T_kv, H_kv, D)– compressed KV sequenceV(Value):(B, H_kv, S_kv, D_v)or(T_kv, H_kv, D_v)cum_seqlen_q:(batch_size + 1,)– cumulative sequence lengths for queries (T,H,D layout only)cum_seqlen_k:(batch_size + 1,)– cumulative sequence lengths for compressed keys (T,H,D layout only)
-
Outputs
O(Output): Same shape asQLSE(LogSumExp, optional):(B, H_q, S_q)or(T, H_q)
Equation
Standard scaled dot-product attention with compressed causal masking:
where , , , are optional scaling factors for quantized inputs.
High-level Wrapper
Class API
Parameters
Constraints
- Input dtype must be
float16,bfloat16, orfloat8_e4m3fn - Output dtype must be
float16,bfloat16, orfloat8_e4m3fn - Head dimension
Dmust be one of{32, 64, 128} H_qmust be divisible byH_kv(supports GQA/MQA)- Requires SM100+ (Blackwell or newer)
3. Sliding Window Attention
Sliding Window Attention performs attention within a local window around each query position. This captures fine-grained local dependencies efficiently. This implementation is a wrapper around cudnn backend (and is not strictly open source).
Shapes
-
Inputs
Q(Query):(B, H_q, S_q, D)or(T, H_q, D)K(Key):(B, H_kv, S_kv, D)or(T, H_kv, D)V(Value):(B, H_kv, S_kv, D_v)or(T, H_kv, D_v)seq_len_q:(B, 1, 1, 1)– sequence lengths for queries (T,H,D layout only)seq_len_kv:(B, 1, 1, 1)– sequence lengths for keys/values (T,H,D layout only)
-
Outputs
O(Output): Same shape asQStats(optional):(B, H_q, S_q, 1)or(T, H_q, 1)– softmax statistics for training
Equation
For each query position , attention is restricted to key positions within the window:
where is left_bound and is right_bound.
High-level Wrapper
Class API
Parameters
Constraints
- Supports both
B,H,S,D(batched) andT,H,D(variable-length) layouts - For
T,H,Dlayout, requiresseq_len_qandseq_len_kv(and optionally ragged offset tensors, otherwise fully packed layout is assumed) cudnn_handleshould be reused across calls for performance
4. Top-K Reduction
Top-K Reduction identifies the most important key-value blocks for each query position based on attention scores. This is used to generate block indices for Selection Attention.
Shapes
-
Inputs
Q(Query):(B, H_q, S_q, D)or(T, H_q, D)K(Key):(B, H_kv, S_kv, D)or(T, H_kv, D)LSE(LogSumExp):(B, H_q, S_q)or(T, H_q)– from a prior attention passcum_seqlen_q:(batch_size + 1,)– cumulative sequence lengths (T,H,D layout only)cum_seqlen_k:(batch_size + 1,)– cumulative sequence lengths (T,H,D layout only)
-
Outputs
topk_scores:(B, H_kv, S_q, K)or(T, H_kv, K)– top-K attention scorestopk_indices:(B, H_kv, S_q, K)or(T, H_kv, K)– indices of top-K blocks
Equation
For each query position, compute block-level attention scores and select the top blocks:
High-level Wrapper
Class API
Parameters
Constraints
- Input dtype for Q/K must match
- LSE dtype must match
acc_dtype topk_indicesmust beint32- Requires SM100+ (Blackwell or newer)
Note: The returned values exclude the first block and neighboring blocks from the reduction. Rows with all -inf scores and -1 indices are expected for positions near the beginning of sequences.
Tensor Formats
Supported Layouts
T,H,D Format (Variable-Length Batched)
Used for sequences of varying lengths packed into a single tensor:
- Q/K/V:
(T, H, D)whereT = sum(seq_lengths) - cum_seqlen:
(batch_size + 1,)– cumulative sequence lengths, e.g.,[0, 128, 320, 512]for 3 sequences of lengths 128, 192, 192
B,H,S,D Format (Fixed-Length Batched)
Traditional batched format with padding:
- Q/K/V:
(B, H, S, D)whereBis batch size,Sis padded sequence length
Data Types
Supported Input/Output Types
Accumulator Types
All components require float32 accumulator dtype for numerical stability.
Hardware Requirements
Usage Examples
For complete usage examples and tests, see:
test/python/fe_api/nsa/test_NSA_selection_attention.pytest/python/fe_api/nsa/test_NSA_compression_attention.pytest/python/fe_api/nsa/test_NSA_swa.pytest/python/fe_api/nsa/test_NSA_topk_reduction.py