How FlashKDA Calculates and Uses Workspace Size for CUDA Kernel Execution

FlashKDA computes the required GPU workspace size via the C++ helper get_workspace_size in 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 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:

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:

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:

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

Finally, it returns the total workspace size:

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, the wrapper computes dimensions and allocates the buffer immediately before kernel launch:

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:

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 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 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.

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 →