# How FlashKDA Calculates and Uses Workspace Size for CUDA Kernel Execution

> Learn how FlashKDA calculates workspace size for CUDA kernel execution. Discover the linear scaling with attention heads and sequence tiles for intermediate computations.

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

---

**FlashKDA computes the required GPU workspace size via the C++ helper `get_workspace_size` in [`csrc/flash_kda.cpp`](https://github.com/MoonshotAI/FlashKDA/blob/main/csrc/flash_kda.cpp), allocating a temporary byte buffer in Python that scales linearly with attention heads and sequence tiles to support intermediate kernel computations.**

The MoonshotAI/FlashKDA repository implements a high-performance linear attention mechanism that requires temporary GPU memory for intermediate results during CUDA kernel execution. Understanding how FlashKDA calculates workspace size ensures efficient memory planning for large batch processing and variable-length sequences. This analysis examines the exact byte-counting logic implemented in the C++ source and how the Python wrapper manages this buffer during the forward pass.

## How Workspace Size is Calculated in FlashKDA

The workspace calculation logic resides in [`csrc/flash_kda.cpp`](https://github.com/MoonshotAI/FlashKDA/blob/main/csrc/flash_kda.cpp) within the `get_workspace_size` function (lines 5-26). This function performs a deterministic upper-bound calculation based on tiling constants and batch dimensions.

### The `get_workspace_size` Function Signature

The helper accepts three parameters that describe the problem size:

- **`T_total`**: Total number of timesteps across the batch, computed as `B * T_seq`.
- **`H`**: Number of attention heads.
- **`N`**: Batch size when using padded sequences, or the number of variable-length sequences when `cu_seqlens` is provided (defaulting to 1).

### Step-by-Step Calculation Logic

The function derives the total bytes through a five-step process using fixed constants `CHUNK = 16` (tile size) and `D = 128` (head dimension).

First, it calculates the upper-bound number of tiles:

```cpp
int64_t total_tiles = (T_total + CHUNK - 1) / CHUNK + N;

```

Each sequence may contribute at most one extra tile beyond simple floor division.

Second, it computes the per-tile memory footprint, accounting for all intermediate buffers aligned to 128-byte boundaries:

```cpp
int64_t per_tile_bytes = 3 * CHUNK * D * 2   // q/k/v decayed values
                        + D * 4            // g_total
                        + 2 * CHUNK * CHUNK * 2; // INV/Mqk

```

Third, it calculates the trailing prefix-sum buffer needed for tile accumulation:

```cpp
int64_t tile_prefix_bytes = ((N + 1) * 4 + 127) / 128 * 128;

```

Finally, it returns the total workspace size:

```cpp
return H * total_tiles * per_tile_bytes + tile_prefix_bytes;

```

### Memory Scaling Factors

The workspace scales linearly with three dimensions:

- **Head count (`H`)**: Each attention head requires isolated temporary storage.
- **Tile count**: Determined by sequence length divided by chunk size (16), plus one tile per sequence.
- **Fixed per-tile overhead**: Approximately 12,800 bytes per tile (derived from the constants above).

## How the Workspace Buffer is Used During Execution

After calculation, the workspace moves through the Python-C++ boundary to serve as scratch memory during kernel execution.

### Python Wrapper Allocation

In [`flash_kda/__init__.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/flash_kda/__init__.py), the wrapper computes dimensions and allocates the buffer immediately before kernel launch:

```python
B, T_seq, H = q.shape[0], q.shape[1], q.shape[2]
T_total = B * T_seq
N = cu_seqlens.numel() - 1 if cu_seqlens is not None else B

workspace = torch.empty(
    get_workspace_size(T_total, H, N),
    dtype=torch.uint8,
    device=q.device,
)

```

The `torch.uint8` tensor provides a raw byte pointer without type constraints, allowing the C++ layer to reinterpret the memory as needed.

### CUDA Kernel Consumption

The allocated workspace passes unchanged to the C++ `fwd` function, where it is interpreted as `workspace.data_ptr()` and supplied to `launch_fwd`. The CUDA kernels read from and write to this buffer for temporary storage of decayed query/key/value states, inverse matrices, and cumulative sums. The buffer contents are ephemeral—valid only during the kernel execution—and are automatically discarded when the operation completes.

## Practical Example: Computing Workspace Requirements

You can inspect the workspace requirements before running the forward pass:

```python
import torch
from flash_kda import fwd, get_workspace_size

# Configuration: B=2, T=64, H=8, D=128

B, T, H, D = 2, 64, 8, 128
q = torch.randn(B, T, H, D, dtype=torch.bfloat16, device='cuda')
k = v = q.clone()
g = torch.randn_like(q)
beta = torch.randn(B, T, H, dtype=torch.bfloat16, device='cuda')
out = torch.empty_like(q)
A_log = torch.randn(H, dtype=torch.float32, device='cuda')
dt_bias = torch.randn(H, D, dtype=torch.float32, device='cuda')
lower_bound = -2.0

# Calculate required workspace

workspace_bytes = get_workspace_size(T_total=B*T, H=H, N=B)
print(f"Workspace required: {workspace_bytes / 1024:.1f} KB")

# Execute forward pass (handles allocation internally)

fwd(q, k, v, g, beta, scale=1.0, out=out,
    A_log=A_log, dt_bias=dt_bias, lower_bound=lower_bound)

```

This example demonstrates how the workspace scales with batch and head dimensions, requiring approximately 1.6 MB for this specific configuration.

## Summary

- **FlashKDA workspace size** is computed by `get_workspace_size` in [`csrc/flash_kda.cpp`](https://github.com/MoonshotAI/FlashKDA/blob/main/csrc/flash_kda.cpp) using a deterministic formula based on `CHUNK=16` and `D=128` constants.
- The calculation accounts for per-tile buffers (query/key/value decay, inverse matrices) plus a 128-byte aligned prefix-sum buffer.
- The Python wrapper in [`flash_kda/__init__.py`](https://github.com/MoonshotAI/FlashKDA/blob/main/flash_kda/__init__.py) allocates the workspace as a `torch.uint8` tensor on the input device before passing it to CUDA kernels.
- Workspace scales linearly with the number of attention heads (`H`) and the upper-bound tile count derived from total timesteps and batch size.
- Users do not manually manage workspace contents; the library handles allocation and deallocation automatically during the forward pass.

## Frequently Asked Questions

### What is the purpose of the workspace buffer in FlashKDA?

The workspace buffer provides temporary GPU memory for intermediate computations during the linear attention forward pass. It stores decayed query/key/value states, cumulative sums, and matrix inverses that the CUDA kernels generate and consume during execution, preventing the need for separate allocations inside the kernel launch.

### How does variable-length sequence handling affect workspace size?

When using `cu_seqlens` for variable-length sequences, the `N` parameter equals `cu_seqlens.numel() - 1` rather than the batch size `B`. This increases the tile count upper-bound (`total_tiles`) by adding one tile per sequence rather than per batch, and adjusts the prefix-sum buffer size `tile_prefix_bytes` accordingly to accommodate the irregular memory access patterns.

### Why is the tile prefix-sum buffer padded to 128 bytes?

The prefix-sum buffer undergoes 128-byte alignment to ensure coalesced memory access patterns when CUDA threads perform parallel scans across tiles. The calculation `((N + 1) * 4 + 127) / 128 * 128` rounds the required `int32` elements up to the nearest cache line boundary, optimizing memory throughput for the cumulative sum operations.

### Can users manually specify a smaller workspace size than calculated?

No, users cannot safely specify a smaller workspace size. The `get_workspace_size` function computes a hard upper bound based on worst-case tiling scenarios; providing less memory would result in buffer overruns and undefined behavior during kernel execution. The Python wrapper always allocates exactly the computed size to ensure correctness across all supported input configurations.