flashinfer.gdn_prefill.chunk_gated_delta_rule¶
- flashinfer.gdn_prefill.chunk_gated_delta_rule(q: Tensor, k: Tensor, v: Tensor, g: Tensor | None = None, beta: Tensor | None = None, scale: float | None = None, initial_state: Tensor | None = None, output_final_state: bool = False, cu_seqlens: Tensor | None = None, use_qk_l2norm_in_kernel: bool = False, output: Tensor | None = None, output_state: Tensor | None = None, state_checkpoints: Tensor | None = None, checkpoint_cu_starts: Tensor | None = None, checkpoint_every_n_tokens: int = 0, use_cp: Literal['auto'] | bool = 'auto', state_indices: Tensor | None = None, _cp_chunk_len: int | None = None, backend: Literal['auto', 'flashinfer', 'cake_gdn', 'cudnn'] = 'auto', max_seqlen: int | None = None) Tensor | Tuple[Tensor, Tensor]¶
Chunked Gated Delta Rule (GDN) attention for prefill.
Implements the gated delta rule linear attention mechanism for efficient training and inference. Supports both GQA (grouped query attention) and GVA (grouped value attention) configurations.
- Parameters:
q (torch.Tensor) – Queries of shape
[total_seq_len, num_q_heads, head_size]. Must be contiguous and on CUDA.k (torch.Tensor) – Keys of shape
[total_seq_len, num_k_heads, head_size]. Must be contiguous and on CUDA.v (torch.Tensor) – Values of shape
[total_seq_len, num_v_heads, head_size]. Must be contiguous and on CUDA.g (torch.Tensor, optional) – Forget gate (alpha) of shape
[total_seq_len, num_sab_heads]wherenum_sab_heads = max(num_q_heads, num_v_heads). Must be float32. Defaults to all ones whenNone.beta (torch.Tensor, optional) – Update gate (beta) of shape
[total_seq_len, num_sab_heads]. Must be float32. Defaults to all ones whenNone.scale (float, optional) – Scale factor for the attention scores. Defaults to
1 / sqrt(head_size)whenNone.initial_state (torch.Tensor, optional) – Initial KV state. Packed, sequence-ordered shape
[num_seqs, num_sab_heads, head_size, head_size]. Must be float32, bfloat16, float16, float8_e4m3fn, or float8_e5m2. Starts from zero state whenNone. Whenstate_indicesis given (SM90/SM100/SM103/SM120), this is instead the state pool[N_pool, num_sab_heads, head_size, head_size]and sequenceireads its initial state from rowstate_indices[i]; on SM100/SM103 every rank-4 positive, non-overlapping stride layout is accepted.output_final_state (bool) – Whether to output the final state. Default:
False.cu_seqlens (torch.Tensor) – Cumulative sequence lengths of shape
[num_seqs + 1], integer int32 or int64 dtype on the same CUDA device asq. Required for variable-length sequences (varlen mode); must not beNone. Repeated adjacent offsets represent legal zero-length sequences.use_qk_l2norm_in_kernel (bool) – Whether to L2-normalize each Q/K head with epsilon
1e-6. Normalization accumulates in float32 and rounds back to the input dtype before the chunked kernel. Default:False.output (torch.Tensor, optional) – Pre-allocated output tensor of shape
[total_seq_len, num_o_heads, head_size]wherenum_o_heads = max(num_q_heads, num_v_heads). Allocated automatically whenNone.output_state (torch.Tensor, optional) – Pre-allocated output state tensor. Packed, sequence-ordered shape
[num_seqs, num_sab_heads, head_size, head_size]. May be float32, bfloat16, float16, float8_e4m3fn, or float8_e5m2. Required whenoutput_final_state=True. Whenstate_indicesis given it is instead the output state pool[N_pool, ...]and sequencei’s final state is written to rowstate_indices[i](in place whenoutput_state is initial_state); it must be provided by the caller (auto-allocation is rejected, since a compact[num_seqs, ...]buffer would be indexed out of bounds by the pool slot ids). SM100/SM103 accepts the same positive, non-overlapping rank-4 stride layouts asinitial_state.state_checkpoints (torch.Tensor, optional) – Pre-allocated checkpoint tensor of shape
[total_checkpoints, num_sab_heads, head_size, head_size]. May be float32, bfloat16, float16, float8_e4m3fn, or float8_e5m2. Required whencheckpoint_every_n_tokens > 0. Context-parallel checkpointing is currently supported on SM90, SM100, and SM120.checkpoint_cu_starts (torch.Tensor, optional) – Cumulative checkpoint counts of shape
[num_seqs + 1], int32 or int64 on the same CUDA device asq.checkpoint_cu_starts[i+1] - checkpoint_cu_starts[i]is the number of checkpoints for sequencei(=seq_len_i // checkpoint_every_n_tokens). Required whencheckpoint_every_n_tokens > 0. The values must be monotonic and consistent withcu_seqlens; this caller precondition is not checked at launch to avoid a device-to-host synchronization.checkpoint_every_n_tokens (int) – Store intermediate state every N tokens. Must be a multiple of the chunk size (64).
0disables checkpointing (default).use_cp (Literal["auto"] | bool, optional:) – Whether to use context parallelism when low-parallelism heuristics match. SM100/SM103 uses the generated GDN CP-only four-stage implementation for structurally supported shapes. Other legal configurations retain the CuTe-DSL implementation.
"auto"enables conservative routing,Truerequires CP support, andFalsedisables CP. Default:"auto".state_indices (torch.Tensor, optional) –
Int32 or int64 tensor of shape
[num_seqs](SM90/SM100/SM103/SM120). When provided,initial_stateandoutput_stateare treated as a state pool whose first dimension is indexed by these slot ids rather than laid out in sequence order: sequenceireads its initial state from rowstate_indices[i]and writes its final state back to the same row (in place whenoutput_state is initial_state). This lets callers that keep a paged/indexed state pool avoid gathering the active rows into a packed buffer and scattering the result back. The pool may be any positive, non-overlapping rank-4 stride layout on SM100/SM103.None(default) keeps sequence-ordered row mapping without requiring the physical view itself to be contiguous.The ids must be unique: as with any indexed scatter, two sequences sharing a slot id would concurrently write the same pool row across work tiles, leaving that row’s final state nondeterministic. Uniqueness is a caller precondition (not checked at launch, to avoid a per-call host sync); the caller’s slot allocator is expected to guarantee it.
_cp_chunk_len (int, optional) – Internal context-parallel chunk-length override used for testing and tuning.
Nonelets the CP backend select the length automatically; an explicit value must be a multiple of 64.backend ({"auto", "flashinfer", "cake_gdn", "cudnn"}) –
autouses the same SM90/SM100/SM120 kernels and context-parallel routing asflashinfer. Cake kernels require an explicitcake_gdnrequest. Usebackend="cake_gdn", use_cp=Truefor Cake CP on SM100/SM103;use_cp=Falseor"auto"retains Cake non-CP. Explicit Cake requests fail for unsupported inputs without falling back to another backend.cudnnruns cuDNN’s fused SM100 linear-attention engine throughflashinfer.cudnn.cudnn_chunk_gated_delta_rule().max_seqlen (int, optional) – Safe upper bound on the maximum logical sequence length, no larger than
total_seq_len. CP kernels use this host-side hint to bound their per-sequence launch grids without readingcu_seqlensback from the GPU. When omitted, CP usestotal_seq_len, which is correct for any batch; passing the exact maximum of a batched call lets CP launch smaller grids.
- Returns:
When
output_final_state=False, the output tensor of shape[total_seq_len, num_o_heads, head_size]. Otherwise a tuple(output, final_state)wherefinal_statehas shape[num_seqs, num_sab_heads, head_size, head_size]— or, whenstate_indicesis given, the state pool[N_pool, ...]itself (i.e.output_state), whose rows named bystate_indicesnow hold the updated final states.- Return type:
torch.Tensor or Tuple[torch.Tensor, torch.Tensor]
Notes
Supports GQA (
num_q_heads > num_k_heads = num_v_heads) and GVA (num_v_heads > num_q_heads = num_k_heads).The final state layout is
[N, H, V, K].Requires SM90 (Hopper) or SM100 (Blackwell) architecture. The SM100 path requires
head_size == 128. On SM100/SM103,gdn_cpsupports structurally legal equal-head, GQA, and GVA shapes on CUDA 12.8, CUDA 12.9, and CUDA 13. Other SM100 CP DSL routes require CUDA 13 andnvidia-cutlass-dsl[cu13]>=4.4.2(pip install flashinfer-python[cu13]).