# How FlashKDA Handles Variable-Length Sequences Using `cu_seqlens`

> Learn how FlashKDA efficiently manages variable-length sequences with cu_seqlens. Discover how it avoids padding memory overhead for faster processing.

- Repository: [Moonshot AI/FlashKDA](https://github.com/MoonshotAI/FlashKDA)
- Tags: internals
- Published: 2026-07-31

---

**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:

```python
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:

```python
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`](https://github.com/MoonshotAI/FlashKDA/blob/main/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`](https://github.com/MoonshotAI/FlashKDA/blob/main/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.

```python
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:

```python
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`](https://github.com/MoonshotAI/FlashKDA/blob/main/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`.