How FlashKDA Handles Variable-Length Sequences Using `cu_seqlens`

FlashKDA processes batches containing sequences of different lengths by accepting a cumulative-sequence-lengths tensor (cu_seqlens) that tells the underlying CUDA kernels exactly where each logical sequence starts and ends in a flattened tensor, eliminating the memory overhead of padding.

The FlashKDA library — developed by MoonshotAI — extends Flash Attention to support packed, variable-length inputs without requiring uniform sequence lengths. By leveraging the cu_seqlens tensor, the library avoids the costly memory and compute waste associated with padding every sequence to the maximum length in a batch.

What Is cu_seqlens and Why It Matters

cu_seqlens is a one-dimensional cumulative sequence lengths tensor that encodes the boundaries of variable-length sequences within a flattened batch.

For a batch of $B$ sequences with lengths $len_1 \dots len_B$, the tensor stores the running total:

cu_seqlens[i] = sum(len_j for j in range(i + 1))  # i = 0 … B-1

The final entry equals the total number of tokens $N$ in the batch. This representation allows CUDA kernels to compute the start and end indices of any sequence $i$ with simple arithmetic:

start_i = 0 if i == 0 else cu_seqlens[i-1]
end_i   = cu_seqlens[i]

Because the kernels know these boundaries, they can apply per-sequence attention masks and skip padded regions entirely, achieving $O(N \cdot d)$ compute complexity where $N$ is the actual token count rather than $B \times max_length$.

How FlashKDA Uses cu_seqlens in the Forward Pass

In flash_kda/__init__.py, the public function flash_attn_forward exposes an optional cuseqlens argument. When provided, the Python wrapper passes this tensor directly to the underlying C++/CUDA implementation.

Inside the kernel, the cumulative lengths are converted into per-token offsets that drive the packed matrix multiplications used by Flash Attention. This mechanism ensures that attention computations respect each sequence’s actual length while maintaining the memory-efficient tiling strategy that makes Flash Attention fast.

Key Implementation Details

  • Location: flash_kda/__init__.py — exposes flash_attn_forward
  • Parameter: cuseqlens (1-D torch.int64 tensor)
  • Requirement: Must be sorted in strictly ascending order
  • Validation: The final entry must match the size of the flattened query/key/value tensors to prevent out-of-bounds memory accesses

If cuseqlens is omitted, FlashKDA assumes a regular padded batch where all sequences share the same length and behave like a standard rectangular tensor.

Computing and Passing cu_seqlens for Variable-Length Batches

The following example demonstrates how to construct the cu_seqlens tensor for a batch of three sequences with lengths 7, 4, and 9, then pass it to FlashKDA.

import torch
from flash_kda import flash_attn_forward

# ----------------------------------------------------------------------

# 1. Define variable-length sequences

# ----------------------------------------------------------------------

seq_lengths = torch.tensor([7, 4, 9], dtype=torch.int64)  # 3 sequences

total_tokens = seq_lengths.sum().item()  # 20

# Create flattened Q/K/V tensors (no padding)

Q = torch.randn(total_tokens, 64, device='cuda')
K = torch.randn(total_tokens, 64, device='cuda')
V = torch.randn(total_tokens, 64, device='cuda')

# ----------------------------------------------------------------------

# 2. Compute cumulative sequence lengths

# ----------------------------------------------------------------------

cu_seqlens = torch.cumsum(seq_lengths, dim=0)  # tensor([ 7, 11, 20])

# ----------------------------------------------------------------------

# 3. Execute FlashKDA forward pass

# ----------------------------------------------------------------------

output = flash_attn_forward(
    q=Q, k=K, v=V,
    cuseqlens=cu_seqlens,   # Enables variable-length processing

    dropout=0.0,
    causal=False
)

print(output.shape)  # torch.Size([20, 64])

Using the High-Level Wrapper

The FlashKDA class can also handle the conversion internally when you provide raw lengths:

from flash_kda import FlashKDA

model = FlashKDA(dim=64, causal=False).cuda()
output = model(
    q=Q, k=K, v=V,
    seq_lens=seq_lengths  # Wrapper builds cu_seqlens internally

)

Edge Cases and Validation

FlashKDA enforces several constraints on cu_seqlens to ensure memory safety:

  • Data Type: Must be torch.int64 (64-bit integer)
  • Shape: Must be a 1-D vector of length $B$ (number of sequences)
  • Ordering: Values must be strictly ascending; duplicate or out-of-order values raise an error
  • Final Entry: Must exactly equal the first dimension of the query/key/value tensors

If these conditions are not met, the kernel raises a clear error before launching the CUDA operation, preventing silent memory corruption.

Summary

  • FlashKDA uses cu_seqlens — a cumulative sum of sequence lengths — to locate the start and end of each variable-length sequence within a flattened batch tensor.
  • The flash_attn_forward function in flash_kda/__init__.py accepts this tensor and passes it to CUDA kernels that compute attention without padding.
  • When cu_seqlens is provided, FlashKDA achieves the same $O(N \cdot d)$ complexity on variable-length inputs as standard Flash Attention does on fixed-length inputs.
  • Omitting cu_seqlens causes FlashKDA to fall back to standard padded batch processing.
  • The tensor must be int64, sorted ascending, and its last element must match the total token count.

Frequently Asked Questions

What happens if I don't provide cu_seqlens to FlashKDA?

FlashKDA assumes a regular, fully-padded batch where every sequence has the same length. The kernels will treat the input as a dense $B \times N \times D$ tensor, which wastes compute and memory if your actual sequences have varying lengths.

What data type and shape does cu_seqlens require?

cu_seqlens must be a 1-D torch.int64 tensor with length equal to the number of sequences in the batch. The values must represent strictly increasing cumulative lengths, with the final value equal to the total number of tokens across all sequences.

How does FlashKDA prevent out-of-bounds memory accesses with variable-length inputs?

Before launching CUDA kernels, FlashKDA validates that the final entry of cu_seqlens matches the size of the flattened input tensors (queries, keys, values). This check ensures that no kernel thread attempts to read or write beyond the allocated memory bounds, raising a runtime error if the mismatch is detected.

Can I use cu_seqlens with causal attention masking?

Yes. When causal=True is passed to flash_attn_forward, the CUDA kernels use the cu_seqlens boundaries to ensure that causality is enforced independently within each sequence. Each sequence attends only to previous positions within its own logical span, respecting the variable-length boundaries defined by cu_seqlens.

Have a question about this repo?

These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →