Files
rustytorch/docs/superpowers/plans/2026-06-27-varlen-flash-attention.md
T
Omar SobhandClaude Sonnet 4.6 80d7c7fb6c feat(flash-attention): add varlen packed-sequence support
Implements variable-length (varlen) FlashAttention that processes
mixed-length batches without padding waste:

- New CUDA kernel flash_attention_varlen_forward with BLOCK_Q=64 /
  BLOCK_K=64 tiling; grid=(ceil(max_seqlen_q/64), num_heads, 1).
  Each block uses a linear scan over cu_seqlens_q to identify its
  owning sequence and exits early when past sequence end.
- New Rust module flash_varlen_forward: always-compiled CPU simulation
  (varlen_attention_cpu) for testing + #[cfg(cuda)] FlashVarlenKernel.
- SdpaBackend::VarLen variant added to backend_selector.
- 8 new CPU-only tests; total test count: 50.

Co-Authored-By: Claude Sonnet 4.6 <[email protected]>
2026-06-27 01:09:39 +00:00

1221 lines
47 KiB
Markdown
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# Varlen Flash Attention Implementation Plan
> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking.
**Goal:** Add variable-length (varlen) packed-sequence support to `rtx-flash-attention` so mixed-length batches can be processed without padding waste, using a `cu_seqlens` prefix-sum descriptor.
**Architecture:** A new Rust file `flash_varlen_forward.rs` provides a pure-Rust CPU simulation (`varlen_attention_cpu`) that is always compiled, plus a `#[cfg(feature = "cuda")]`-gated `FlashVarlenKernel` struct that loads a new CUDA kernel via NVRTC. A new `SdpaBackend::VarLen` variant is added to `backend_selector.rs`. The CUDA kernel maps each grid block to a (q-tile, head) pair using a linear scan over `cu_seqlens_q`, then runs tiled online-softmax attention with BLOCK_Q=64 / BLOCK_K=64 tiles.
**Tech Stack:** Rust 2021, cudarc 0.16 (CUDA feature), half::bf16, existing FlashError/FlashResult error types.
## Global Constraints
- Only modify files inside `crates/training/rtx-flash-attention/`
- `cargo check -p rtx-flash-attention` must produce zero errors/warnings after each task
- `cargo test -p rtx-flash-attention --lib` must pass (42 existing + 8 new = 50 tests total after Task 3)
- No new Cargo dependencies
- No GPU hardware required to run tests (all 8 new tests use `varlen_attention_cpu`)
- Clippy-clean: no `#[allow(clippy::...)]` suppressions without a comment explaining why
---
## File Map
| Status | Path | Role |
|--------|------|------|
| **CREATE** | `src/kernels/cuda/flash_attention_varlen.cu` | CUDA kernel — varlen forward pass |
| **CREATE** | `src/kernels/flash_varlen_forward.rs` | Rust wrapper + CPU simulation |
| **MODIFY** | `src/kernels/mod.rs` | Export the new module and re-exports |
| **MODIFY** | `src/backend_selector.rs` | Add `SdpaBackend::VarLen` variant |
---
## Task 1: CUDA Kernel Source — `flash_attention_varlen.cu`
**Files:**
- Create: `crates/training/rtx-flash-attention/src/kernels/cuda/flash_attention_varlen.cu`
**Interfaces:**
- Produces: CUDA `__global__` function `flash_attention_varlen_forward` with the signature shown in Step 1.
- [ ] **Step 1: Write the CUDA kernel file**
Create the file with the following content exactly. The key design decisions are:
- `BLOCK_Q = 64`, `BLOCK_K = 64` — tile dimensions (match spec)
- Grid: `(ceil(max_seqlen_q / BLOCK_Q), num_heads, 1)` — blocks identified by `(q_tile_idx, head_idx)`
- Sequence identification: linear scan over `cu_seqlens_q` in the block prologue
- Accumulation in `f32`; output stored as `__nv_bfloat16`
- Online softmax (Dao FA2 algorithm): running `m` (row max) and `l` (denominator)
```c
// Variable-length FlashAttention forward kernel (packed sequences)
//
// Preconditions (caller-enforced):
// q, k, v are packed [total_tokens, num_heads, head_dim] in row-major, dtype __nv_bfloat16
// out is same shape, pre-allocated
// lse is [total_tokens, num_heads] in f32
// cu_seqlens_q[0] == 0, cu_seqlens_q[batch_size] == total_tokens_q
// cu_seqlens_k[0] == 0, cu_seqlens_k[batch_size] == total_tokens_k
// max_seqlen_q >= any individual sequence q length
// 0 < head_dim <= 256, head_dim % 8 == 0
//
// Grid/block assignment:
// blockDim = (128, 1, 1) — 4 warps
// gridDim = (ceil(max_seqlen_q / BLOCK_Q), num_heads, 1)
#include <cuda_runtime.h>
#include <cuda_bf16.h>
#include <stdint.h>
#include <float.h>
#define BLOCK_Q 64
#define BLOCK_K 64
// Per-token element offset in packed layout [total_tokens, num_heads, head_dim]
// token_idx is the absolute position in the packed buffer.
__device__ __forceinline__ int elem_offset(int token_idx, int head_idx, int d, int num_heads, int head_dim) {
return (token_idx * num_heads + head_idx) * head_dim + d;
}
// Find sequence index s such that cu_seqlens[s] <= pos < cu_seqlens[s+1].
// Linear scan is correct for all batch sizes; constant-time for batch <= 64.
__device__ __forceinline__ int find_sequence(const int* cu_seqlens, int batch_size, int pos) {
int s = 0;
while (s < batch_size - 1 && cu_seqlens[s + 1] <= pos) {
s++;
}
return s;
}
extern "C" __global__ void flash_attention_varlen_forward(
const __nv_bfloat16* __restrict__ q,
const __nv_bfloat16* __restrict__ k,
const __nv_bfloat16* __restrict__ v,
__nv_bfloat16* __restrict__ out,
float* __restrict__ lse,
const int* __restrict__ cu_seqlens_q,
const int* __restrict__ cu_seqlens_k,
int max_seqlen_q,
int batch_size,
int num_heads,
int head_dim,
float softmax_scale,
int causal
) {
// -------------------------------------------------------------------------
// Identify which q-tile and head this block handles
// -------------------------------------------------------------------------
const int q_tile_idx = blockIdx.x; // which BLOCK_Q tile within max_seqlen_q
const int head_idx = blockIdx.y;
const int tid = threadIdx.x; // 0..127
const int q_tile_start_global = q_tile_idx * BLOCK_Q;
// -------------------------------------------------------------------------
// Identify which sequence owns this q-tile (linear scan over cu_seqlens_q)
// -------------------------------------------------------------------------
int s = find_sequence(cu_seqlens_q, batch_size, q_tile_start_global);
const int seq_start_q = cu_seqlens_q[s];
const int seq_end_q = cu_seqlens_q[s + 1];
const int seq_len_q = seq_end_q - seq_start_q;
const int seq_start_k = cu_seqlens_k[s];
const int seq_end_k = cu_seqlens_k[s + 1];
const int seq_len_k = seq_end_k - seq_start_k;
// Local q-tile start within this sequence
const int local_q_start = q_tile_start_global - seq_start_q;
// Early exit: this block is completely past the end of sequence s
if (local_q_start >= seq_len_q) return;
const int q_rows_this_tile = min(BLOCK_Q, seq_len_q - local_q_start);
// -------------------------------------------------------------------------
// Shared memory layout:
// [0 .. BLOCK_Q * head_dim) : Q tile (bf16)
// [BLOCK_Q*D*2 .. BLOCK_Q*D*2 + BLOCK_K*D*2) : K tile (bf16)
// [next .. next + BLOCK_K*D*2) : V tile (bf16)
// We access shared mem via float-aligned pointers for vectorised loads.
// -------------------------------------------------------------------------
extern __shared__ char smem_raw[];
__nv_bfloat16* smem_q = (__nv_bfloat16*)smem_raw;
__nv_bfloat16* smem_k = smem_q + BLOCK_Q * head_dim;
__nv_bfloat16* smem_v = smem_k + BLOCK_K * head_dim;
// -------------------------------------------------------------------------
// Load Q tile into shared memory
// Each of 128 threads loads elements strided across [q_row, d].
// -------------------------------------------------------------------------
for (int i = tid; i < q_rows_this_tile * head_dim; i += blockDim.x) {
int row = i / head_dim;
int col = i % head_dim;
int global_token = seq_start_q + local_q_start + row;
smem_q[row * head_dim + col] = q[elem_offset(global_token, head_idx, col, num_heads, head_dim)];
}
// Zero-pad rows that don't exist in this tile
for (int i = q_rows_this_tile * head_dim + tid; i < BLOCK_Q * head_dim; i += blockDim.x) {
smem_q[i] = __float2bfloat16(0.0f);
}
__syncthreads();
// -------------------------------------------------------------------------
// Per-row accumulators in registers (one row per thread — split 64 rows
// across 128 threads, so each thread "owns" rows tid/2 with half the threads
// handling even/odd columns via the inner d-loop).
// For simplicity we assign each thread one q-row cyclically.
// -------------------------------------------------------------------------
// We unroll over q rows: thread tid handles q-row (tid) if tid < BLOCK_Q.
// For 128 threads and BLOCK_Q=64, threads 0..63 process one row each.
// Threads 64..127 do nothing except participate in K/V loads.
// Accumulators: acc[d] for output, running_m and running_l for online softmax
float acc[256]; // max supported head_dim
float running_m = -FLT_MAX;
float running_l = 0.0f;
for (int d = 0; d < head_dim; d++) acc[d] = 0.0f;
const int my_q_row = tid; // my row within the q tile (tid 0..63 active, 64..127 idle)
const int my_active = (my_q_row < q_rows_this_tile) ? 1 : 0;
// -------------------------------------------------------------------------
// Outer loop over K/V tiles
// -------------------------------------------------------------------------
const int num_k_tiles = (seq_len_k + BLOCK_K - 1) / BLOCK_K;
for (int k_tile = 0; k_tile < num_k_tiles; k_tile++) {
const int k_tile_start = k_tile * BLOCK_K;
const int k_rows_this_tile = min(BLOCK_K, seq_len_k - k_tile_start);
// ------------------------------------------------------------------
// Load K tile into shared memory
// ------------------------------------------------------------------
for (int i = tid; i < k_rows_this_tile * head_dim; i += blockDim.x) {
int row = i / head_dim;
int col = i % head_dim;
int global_token = seq_start_k + k_tile_start + row;
smem_k[row * head_dim + col] = k[elem_offset(global_token, head_idx, col, num_heads, head_dim)];
}
for (int i = k_rows_this_tile * head_dim + tid; i < BLOCK_K * head_dim; i += blockDim.x) {
smem_k[i] = __float2bfloat16(0.0f);
}
// ------------------------------------------------------------------
// Load V tile into shared memory
// ------------------------------------------------------------------
for (int i = tid; i < k_rows_this_tile * head_dim; i += blockDim.x) {
int row = i / head_dim;
int col = i % head_dim;
int global_token = seq_start_k + k_tile_start + row;
smem_v[row * head_dim + col] = v[elem_offset(global_token, head_idx, col, num_heads, head_dim)];
}
for (int i = k_rows_this_tile * head_dim + tid; i < BLOCK_K * head_dim; i += blockDim.x) {
smem_v[i] = __float2bfloat16(0.0f);
}
__syncthreads();
if (!my_active) { __syncthreads(); continue; }
// ------------------------------------------------------------------
// Compute attention scores for my_q_row x all k_rows: S[j] = Q[my_q_row] . K[j]
// ------------------------------------------------------------------
float S[BLOCK_K];
for (int j = 0; j < k_rows_this_tile; j++) {
float dot = 0.0f;
for (int d = 0; d < head_dim; d++) {
dot += __bfloat162float(smem_q[my_q_row * head_dim + d]) *
__bfloat162float(smem_k[j * head_dim + d]);
}
S[j] = dot * softmax_scale;
// Causal mask: token (seq_start_q + local_q_start + my_q_row) cannot
// attend to token (seq_start_k + k_tile_start + j) if j > my_q_row+local_q_start.
if (causal) {
int q_pos = local_q_start + my_q_row;
int k_pos = k_tile_start + j;
if (k_pos > q_pos) S[j] = -FLT_MAX;
}
// Pad out-of-sequence k positions
if (j >= k_rows_this_tile) S[j] = -FLT_MAX;
}
for (int j = k_rows_this_tile; j < BLOCK_K; j++) S[j] = -FLT_MAX;
// ------------------------------------------------------------------
// Online softmax update (FA2 algorithm 1)
// ------------------------------------------------------------------
float tile_m = -FLT_MAX;
for (int j = 0; j < k_rows_this_tile; j++) tile_m = fmaxf(tile_m, S[j]);
float m_new = fmaxf(running_m, tile_m);
// Rescale existing accumulator
float rescale = expf(running_m - m_new);
for (int d = 0; d < head_dim; d++) acc[d] *= rescale;
float l_rescale = running_l * rescale;
// Accumulate weighted V
float tile_l = 0.0f;
for (int j = 0; j < k_rows_this_tile; j++) {
float p = expf(S[j] - m_new);
tile_l += p;
for (int d = 0; d < head_dim; d++) {
acc[d] += p * __bfloat162float(smem_v[j * head_dim + d]);
}
}
running_m = m_new;
running_l = l_rescale + tile_l;
__syncthreads();
}
// -------------------------------------------------------------------------
// Write output and LSE
// -------------------------------------------------------------------------
if (!my_active) return;
float inv_l = (running_l > 0.0f) ? (1.0f / running_l) : 0.0f;
int global_token_out = seq_start_q + local_q_start + my_q_row;
for (int d = 0; d < head_dim; d++) {
out[elem_offset(global_token_out, head_idx, d, num_heads, head_dim)] =
__float2bfloat16(acc[d] * inv_l);
}
// lse[token, head] = log(l) + m (log-sum-exp in standard form)
float lse_val = (running_l > 0.0f) ? (logf(running_l) + running_m) : -FLT_MAX;
lse[global_token_out * num_heads + head_idx] = lse_val;
}
```
- [ ] **Step 2: Verify file exists**
```bash
ls -la /slab/projects/rustyverse/rustytorch/crates/training/rtx-flash-attention/src/kernels/cuda/flash_attention_varlen.cu
```
Expected: file listed with non-zero size.
---
## Task 2: Rust Wrapper + CPU Simulation — `flash_varlen_forward.rs`
**Files:**
- Create: `crates/training/rtx-flash-attention/src/kernels/flash_varlen_forward.rs`
**Interfaces:**
- Consumes: `crate::error::{FlashError, FlashResult}` (always); `cudarc::driver::{CudaContext, CudaSlice, CudaStream, LaunchConfig, PushKernelArg}`, `cudarc::nvrtc::compile_ptx`, `half::bf16` (cuda feature only).
- Produces:
- `pub fn varlen_attention_cpu(q: &[f32], k: &[f32], v: &[f32], cu_seqlens: &[usize], num_heads: usize, head_dim: usize, softmax_scale: f32, causal: bool) -> Vec<f32>` — always available.
- `#[cfg(feature = "cuda")] pub struct FlashVarlenKernel` with:
- `pub fn new(ctx: &Arc<CudaContext>) -> FlashResult<Self>`
- `pub fn forward(&self, q: &CudaSlice<bf16>, k: &CudaSlice<bf16>, v: &CudaSlice<bf16>, out: &mut CudaSlice<bf16>, lse: &mut CudaSlice<f32>, cu_seqlens_q: &CudaSlice<i32>, cu_seqlens_k: &CudaSlice<i32>, max_seqlen_q: usize, batch_size: usize, num_heads: usize, head_dim: usize, softmax_scale: f32, causal: bool, stream: &Arc<CudaStream>) -> FlashResult<()>`
- [ ] **Step 1: Write the Rust file**
Create `/slab/projects/rustyverse/rustytorch/crates/training/rtx-flash-attention/src/kernels/flash_varlen_forward.rs` with the following content:
```rust
//! Variable-length (varlen) Flash Attention forward pass.
//!
//! This module provides:
//! - [`varlen_attention_cpu`] — pure-Rust O(n²) reference implementation for
//! testing and CPU fallback. No GPU required.
//! - [`FlashVarlenKernel`] — CUDA kernel wrapper (cuda feature only).
//!
//! # Packed layout
//!
//! All tensors use **packed** (also called "varlen" or "jagged") layout:
//! instead of `[batch, heads, seq_len, head_dim]` with padding, sequences are
//! concatenated along the token axis: `[total_tokens, heads, head_dim]`.
//!
//! A `cu_seqlens` prefix-sum array of length `batch_size + 1` describes the
//! boundaries:
//! - `cu_seqlens[0] = 0`
//! - `cu_seqlens[b + 1] = cu_seqlens[b] + seqlen_b`
//! - `cu_seqlens[batch_size] = total_tokens`
//!
//! Sequence `b` occupies token indices `cu_seqlens[b] .. cu_seqlens[b+1]`.
use crate::error::FlashResult;
#[cfg(feature = "cuda")]
use crate::error::FlashError;
#[cfg(feature = "cuda")]
use cudarc::driver::{CudaContext, CudaSlice, CudaStream, LaunchConfig, PushKernelArg};
#[cfg(feature = "cuda")]
use cudarc::nvrtc::compile_ptx;
#[cfg(feature = "cuda")]
use half::bf16;
#[cfg(feature = "cuda")]
use std::sync::Arc;
#[cfg(feature = "cuda")]
const VARLEN_CUDA_SOURCE: &str =
include_str!("cuda/flash_attention_varlen.cu");
// =============================================================================
// CPU reference implementation
// =============================================================================
/// Variable-length attention on the CPU using naive O(n²) per-sequence dot-product.
///
/// # Layout
///
/// `q`, `k`, `v` are packed flat slices in row-major order:
/// `[total_tokens * num_heads * head_dim]` where the logical index for
/// token `t`, head `h`, dimension `d` is `(t * num_heads + h) * head_dim + d`.
///
/// `cu_seqlens` has length `batch_size + 1`; `cu_seqlens[b..b+1]` is a half-open
/// range of token indices for sequence `b`. Sequences with zero length
/// (i.e. `cu_seqlens[b] == cu_seqlens[b+1]`) are silently skipped.
///
/// # Returns
///
/// A `Vec<f32>` of length `total_tokens * num_heads * head_dim` — the attention
/// output in the same packed layout.
///
/// # Panics
///
/// Panics if `cu_seqlens` is empty (must have at least one entry: `[0]`).
///
/// # Example
///
/// ```
/// use rtx_flash_attention::kernels::varlen_attention_cpu;
///
/// let head_dim = 4;
/// let num_heads = 1;
/// // Two sequences: seq0 = 2 tokens, seq1 = 3 tokens → 5 total tokens
/// let cu_seqlens = vec![0usize, 2, 5];
/// let total = cu_seqlens[cu_seqlens.len() - 1];
/// let scale = 1.0 / (head_dim as f32).sqrt();
///
/// let q: Vec<f32> = (0..total * num_heads * head_dim).map(|i| i as f32 * 0.01).collect();
/// let k = q.clone();
/// let v = q.clone();
///
/// let out = varlen_attention_cpu(&q, &k, &v, &cu_seqlens, num_heads, head_dim, scale, false);
/// assert_eq!(out.len(), total * num_heads * head_dim);
/// ```
pub fn varlen_attention_cpu(
q: &[f32],
k: &[f32],
v: &[f32],
cu_seqlens: &[usize],
num_heads: usize,
head_dim: usize,
softmax_scale: f32,
causal: bool,
) -> Vec<f32> {
assert!(!cu_seqlens.is_empty(), "cu_seqlens must have at least one entry");
let batch_size = cu_seqlens.len() - 1;
let total_tokens = cu_seqlens[batch_size];
let mut output = vec![0.0f32; total_tokens * num_heads * head_dim];
// Helper: flat index for packed layout [total_tokens, num_heads, head_dim]
let idx = |token: usize, head: usize, d: usize| -> usize {
(token * num_heads + head) * head_dim + d
};
for b in 0..batch_size {
let seq_start = cu_seqlens[b];
let seq_end = cu_seqlens[b + 1];
let seq_len = seq_end - seq_start;
// Empty sequence — skip without touching output
if seq_len == 0 {
continue;
}
for h in 0..num_heads {
// Compute attention weights for every (query, key) pair in this sequence.
// scores[qi * seq_len + ki] = Q[qi] . K[ki] * scale
let mut scores = vec![0.0f32; seq_len * seq_len];
for qi in 0..seq_len {
for ki in 0..seq_len {
// Causal: query at position qi cannot attend to key at ki > qi
if causal && ki > qi {
scores[qi * seq_len + ki] = f32::NEG_INFINITY;
continue;
}
let mut dot = 0.0f32;
for d in 0..head_dim {
dot += q[idx(seq_start + qi, h, d)] * k[idx(seq_start + ki, h, d)];
}
scores[qi * seq_len + ki] = dot * softmax_scale;
}
}
// Softmax row-by-row (online — subtract row max for numerical stability)
let mut weights = vec![0.0f32; seq_len * seq_len];
for qi in 0..seq_len {
let row = &scores[qi * seq_len..(qi + 1) * seq_len];
let row_max = row
.iter()
.copied()
.fold(f32::NEG_INFINITY, f32::max);
let exps: Vec<f32> = row.iter().map(|&s| (s - row_max).exp()).collect();
let sum: f32 = exps.iter().sum();
let inv_sum = if sum > 0.0 { 1.0 / sum } else { 0.0 };
for ki in 0..seq_len {
weights[qi * seq_len + ki] = exps[ki] * inv_sum;
}
}
// Output: O[qi, d] = Σ_ki weights[qi, ki] * V[ki, d]
for qi in 0..seq_len {
for d in 0..head_dim {
let mut acc = 0.0f32;
for ki in 0..seq_len {
acc += weights[qi * seq_len + ki]
* v[idx(seq_start + ki, h, d)];
}
output[idx(seq_start + qi, h, d)] = acc;
}
}
}
}
output
}
// =============================================================================
// CUDA kernel wrapper
// =============================================================================
/// Compiled FlashAttention varlen forward kernel for a single CUDA device.
///
/// Create once per device; reuse across calls. The compiled PTX is cached
/// inside the struct.
///
/// # Feature gate
///
/// Only available when the `cuda` feature is enabled.
#[cfg(feature = "cuda")]
pub struct FlashVarlenKernel {
module: Arc<cudarc::driver::CudaModule>,
_ctx: Arc<CudaContext>,
}
#[cfg(feature = "cuda")]
impl FlashVarlenKernel {
/// Compile and load the varlen CUDA kernel into `ctx`.
///
/// # Errors
///
/// Returns [`FlashError::Cuda`] if NVRTC compilation fails or the PTX
/// cannot be loaded into the device context.
pub fn new(ctx: &Arc<CudaContext>) -> FlashResult<Self> {
let ptx = compile_ptx(VARLEN_CUDA_SOURCE)
.map_err(|e| FlashError::cuda(format!("varlen NVRTC compilation failed: {e:?}")))?;
let module = ctx
.load_module(ptx)
.map_err(|e| FlashError::cuda(format!("varlen PTX load failed: {e:?}")))?;
Ok(Self {
module,
_ctx: Arc::clone(ctx),
})
}
/// Launch the varlen forward kernel on `stream`.
///
/// # Arguments
///
/// - `q`, `k`, `v` — packed `[total_tokens, num_heads, head_dim]` in BF16.
/// - `out` — pre-allocated output buffer, same shape as `q`.
/// - `lse` — pre-allocated log-sum-exp buffer `[total_tokens, num_heads]` in f32.
/// - `cu_seqlens_q` / `cu_seqlens_k` — device buffers of length `batch_size + 1`.
/// - `max_seqlen_q` — maximum query sequence length across the batch.
/// - `batch_size` — number of sequences in the batch.
///
/// # Safety invariants (enforced by caller)
///
/// - All device buffers must be allocated on the same device as `ctx`.
/// - Buffer lengths: `q.len() >= total_tokens_q * num_heads * head_dim`.
/// - `cu_seqlens_q[0] == 0`, `cu_seqlens_q[batch_size] == total_tokens_q`.
/// - `head_dim <= 256` and `head_dim % 8 == 0`.
///
/// # Errors
///
/// Returns [`FlashError::Cuda`] if the kernel symbol is not found or the
/// driver rejects the launch configuration.
#[allow(clippy::too_many_arguments)] // kernel interface requires all parameters
pub fn forward(
&self,
q: &CudaSlice<bf16>,
k: &CudaSlice<bf16>,
v: &CudaSlice<bf16>,
out: &mut CudaSlice<bf16>,
lse: &mut CudaSlice<f32>,
cu_seqlens_q: &CudaSlice<i32>,
cu_seqlens_k: &CudaSlice<i32>,
max_seqlen_q: usize,
batch_size: usize,
num_heads: usize,
head_dim: usize,
softmax_scale: f32,
causal: bool,
stream: &Arc<CudaStream>,
) -> FlashResult<()> {
let kernel = self
.module
.load_function("flash_attention_varlen_forward")
.map_err(|e| {
FlashError::cuda(format!("varlen kernel symbol not found: {e:?}"))
})?;
// Grid: (ceil(max_seqlen_q / BLOCK_Q), num_heads, 1)
const BLOCK_Q: usize = 64;
let grid_x = max_seqlen_q.div_ceil(BLOCK_Q) as u32;
let grid_y = num_heads as u32;
// Shared memory: 3 tiles × BLOCK_K × head_dim × sizeof(bf16)
const BLOCK_K: usize = 64;
let smem_bytes = (3 * BLOCK_K * head_dim * 2) as u32; // bf16 = 2 bytes
let launch_cfg = LaunchConfig {
block_dim: (128, 1, 1),
grid_dim: (grid_x, grid_y, 1),
shared_mem_bytes: smem_bytes,
};
let causal_int = causal as i32;
let batch_size_i32 = batch_size as i32;
let num_heads_i32 = num_heads as i32;
let head_dim_i32 = head_dim as i32;
let max_seqlen_q_i32 = max_seqlen_q as i32;
// Safety:
// - All CudaSlice buffers are device-resident on the same device as `_ctx`.
// - `out` and `lse` are mutable unique references, so no aliasing.
// - The kernel reads cu_seqlens_q/k as read-only int arrays; they have
// length `batch_size + 1` ≥ 2 (enforced by caller).
// - head_dim ≤ 256 ensures the `acc[256]` register array is not overflowed.
// - The CUDA kernel performs a bounds check per q-row (early return when
// the tile is past the end of its sequence).
unsafe {
let mut builder = stream.launch_builder(&kernel);
builder.arg(q);
builder.arg(k);
builder.arg(v);
builder.arg(out);
builder.arg(lse);
builder.arg(cu_seqlens_q);
builder.arg(cu_seqlens_k);
builder.arg(&max_seqlen_q_i32);
builder.arg(&batch_size_i32);
builder.arg(&num_heads_i32);
builder.arg(&head_dim_i32);
builder.arg(&softmax_scale);
builder.arg(&causal_int);
builder.launch(launch_cfg)
.map_err(|e| FlashError::cuda(format!("varlen kernel launch failed: {e:?}")))?;
}
Ok(())
}
}
// =============================================================================
// Tests — all CPU, no GPU required
// =============================================================================
#[cfg(test)]
mod tests {
use super::varlen_attention_cpu;
// Helper: flat index for packed [total_tokens, num_heads, head_dim]
fn idx(token: usize, head: usize, d: usize, num_heads: usize, head_dim: usize) -> usize {
(token * num_heads + head) * head_dim + d
}
// Build a simple packed buffer filled with a constant value per token.
fn constant_qkv(
cu_seqlens: &[usize],
num_heads: usize,
head_dim: usize,
val: f32,
) -> Vec<f32> {
let total = *cu_seqlens.last().unwrap();
vec![val; total * num_heads * head_dim]
}
// -------------------------------------------------------------------------
// Test 1: Single sequence — must match standard (non-varlen) attention
// -------------------------------------------------------------------------
#[test]
fn test_varlen_cpu_single_sequence_matches_regular() {
let num_heads = 2;
let head_dim = 8;
let seq_len = 4;
let scale = 1.0 / (head_dim as f32).sqrt();
// cu_seqlens for a single sequence of length 4
let cu_seqlens = vec![0usize, seq_len];
let total = seq_len;
// Random-ish deterministic data
let q: Vec<f32> = (0..total * num_heads * head_dim)
.map(|i| (i as f32) * 0.1)
.collect();
let k = q.clone();
let v = q.clone();
let out = varlen_attention_cpu(&q, &k, &v, &cu_seqlens, num_heads, head_dim, scale, false);
// Output has the right length
assert_eq!(out.len(), total * num_heads * head_dim);
// For a single sequence, varlen output must be identical to the same
// computation run as a single entry — we verify by calling again with
// the same data (idempotency check + non-zero output sanity).
let out2 = varlen_attention_cpu(&q, &k, &v, &cu_seqlens, num_heads, head_dim, scale, false);
for i in 0..out.len() {
assert!(
(out[i] - out2[i]).abs() < 1e-6,
"output not deterministic at index {i}: {} vs {}",
out[i],
out2[i]
);
}
// All values must be finite (no NaN/Inf for well-conditioned inputs)
for (i, &v) in out.iter().enumerate() {
assert!(v.is_finite(), "non-finite output at index {i}: {v}");
}
}
// -------------------------------------------------------------------------
// Test 2: Two sequences are independent — tokens from seq0 do not appear
// in seq1's attention output
// -------------------------------------------------------------------------
#[test]
fn test_varlen_cpu_two_sequences_independent() {
let num_heads = 1;
let head_dim = 4;
let scale = 1.0 / (head_dim as f32).sqrt();
// Sequence 0: 3 tokens, all-ones
// Sequence 1: 4 tokens, all-twos
let cu_seqlens = vec![0usize, 3, 7];
let total = 7;
let mut q = vec![0.0f32; total * num_heads * head_dim];
let mut k = vec![0.0f32; total * num_heads * head_dim];
let mut v = vec![0.0f32; total * num_heads * head_dim];
// seq0 tokens (0..3) → value 1.0
for t in 0..3 {
for d in 0..head_dim {
q[idx(t, 0, d, num_heads, head_dim)] = 1.0;
k[idx(t, 0, d, num_heads, head_dim)] = 1.0;
v[idx(t, 0, d, num_heads, head_dim)] = 1.0;
}
}
// seq1 tokens (3..7) → value 100.0 (very different magnitude)
for t in 3..7 {
for d in 0..head_dim {
q[idx(t, 0, d, num_heads, head_dim)] = 100.0;
k[idx(t, 0, d, num_heads, head_dim)] = 100.0;
v[idx(t, 0, d, num_heads, head_dim)] = 100.0;
}
}
let out = varlen_attention_cpu(&q, &k, &v, &cu_seqlens, num_heads, head_dim, scale, false);
// seq0 output must be ~1.0 (attending only to seq0 tokens, all equal)
for t in 0..3 {
for d in 0..head_dim {
let val = out[idx(t, 0, d, num_heads, head_dim)];
assert!(
(val - 1.0).abs() < 1e-5,
"seq0 token {t} dim {d}: expected ~1.0, got {val}"
);
}
}
// seq1 output must be ~100.0
for t in 3..7 {
for d in 0..head_dim {
let val = out[idx(t, 0, d, num_heads, head_dim)];
assert!(
(val - 100.0).abs() < 1e-5,
"seq1 token {t} dim {d}: expected ~100.0, got {val}"
);
}
}
}
// -------------------------------------------------------------------------
// Test 3: Output shape is [total_tokens, heads, head_dim]
// -------------------------------------------------------------------------
#[test]
fn test_varlen_cpu_output_shape() {
let num_heads = 4;
let head_dim = 16;
let cu_seqlens = vec![0usize, 5, 12, 15]; // 3 seqs, lengths 5, 7, 3
let total = 15;
let scale = 1.0 / (head_dim as f32).sqrt();
let q = vec![1.0f32; total * num_heads * head_dim];
let k = q.clone();
let v = q.clone();
let out = varlen_attention_cpu(&q, &k, &v, &cu_seqlens, num_heads, head_dim, scale, false);
assert_eq!(
out.len(),
total * num_heads * head_dim,
"output length must equal total_tokens * num_heads * head_dim"
);
}
// -------------------------------------------------------------------------
// Test 4: cu_seqlens semantics — [0, 3, 7] means seq0=3 tokens, seq1=4 tokens
// -------------------------------------------------------------------------
#[test]
fn test_varlen_cu_seqlens_correct() {
let num_heads = 1;
let head_dim = 4;
let scale = 1.0 / (head_dim as f32).sqrt();
let cu_seqlens = vec![0usize, 3, 7];
// Verify: batch_size = cu_seqlens.len() - 1 = 2
let batch_size = cu_seqlens.len() - 1;
assert_eq!(batch_size, 2);
// seq0 length = cu_seqlens[1] - cu_seqlens[0] = 3
assert_eq!(cu_seqlens[1] - cu_seqlens[0], 3);
// seq1 length = cu_seqlens[2] - cu_seqlens[1] = 4
assert_eq!(cu_seqlens[2] - cu_seqlens[1], 4);
let total = cu_seqlens[batch_size];
let q = vec![0.5f32; total * num_heads * head_dim];
let k = q.clone();
let v = q.clone();
let out = varlen_attention_cpu(&q, &k, &v, &cu_seqlens, num_heads, head_dim, scale, false);
assert_eq!(out.len(), total * num_heads * head_dim);
}
// -------------------------------------------------------------------------
// Test 5: Causal mask — token i cannot attend to token j > i (within sequence)
// -------------------------------------------------------------------------
#[test]
fn test_varlen_causal_mask() {
let num_heads = 1;
let head_dim = 4;
let seq_len = 4;
let scale = 1.0 / (head_dim as f32).sqrt();
let cu_seqlens = vec![0usize, seq_len];
// Use distinct V values per token so any cross-contamination is detectable.
// V[token t, head 0, all dims] = (t+1) as f32
let total = seq_len;
let mut q = vec![1.0f32; total * num_heads * head_dim];
let mut k = vec![1.0f32; total * num_heads * head_dim];
let mut v = vec![0.0f32; total * num_heads * head_dim];
for t in 0..seq_len {
for d in 0..head_dim {
v[idx(t, 0, d, num_heads, head_dim)] = (t + 1) as f32;
}
}
let out_causal =
varlen_attention_cpu(&q, &k, &v, &cu_seqlens, num_heads, head_dim, scale, true);
let out_non_causal =
varlen_attention_cpu(&q, &k, &v, &cu_seqlens, num_heads, head_dim, scale, false);
// Token 0 in causal mode: can only attend to itself (k=0).
// V[0,*] = 1.0, so output for token 0 must be 1.0 in causal mode.
for d in 0..head_dim {
let val = out_causal[idx(0, 0, d, num_heads, head_dim)];
assert!(
(val - 1.0).abs() < 1e-5,
"causal: token 0 dim {d} expected 1.0, got {val}"
);
}
// In non-causal mode, token 0 attends to all tokens uniformly (Q=K=1.0),
// so output = mean(V) = (1+2+3+4)/4 = 2.5
for d in 0..head_dim {
let val = out_non_causal[idx(0, 0, d, num_heads, head_dim)];
assert!(
(val - 2.5).abs() < 1e-5,
"non-causal: token 0 dim {d} expected 2.5, got {val}"
);
}
// Last token in causal mode attends to all tokens (same as non-causal for last token)
for d in 0..head_dim {
let causal_last = out_causal[idx(seq_len - 1, 0, d, num_heads, head_dim)];
let noncausal_last = out_non_causal[idx(seq_len - 1, 0, d, num_heads, head_dim)];
assert!(
(causal_last - noncausal_last).abs() < 1e-5,
"causal last token should match non-causal: got {causal_last} vs {noncausal_last}"
);
}
}
// -------------------------------------------------------------------------
// Test 6: Attention weights sum to 1.0 per row within 1e-5
// -------------------------------------------------------------------------
#[test]
fn test_varlen_softmax_sums_to_one() {
let num_heads = 2;
let head_dim = 8;
let scale = 1.0 / (head_dim as f32).sqrt();
// Two sequences: 3 and 5 tokens
let cu_seqlens = vec![0usize, 3, 8];
let total = 8;
// V is the identity: each token t, dim d = (t == d) ? 1.0 : 0.0
// This lets us recover the attention weight for each (q, k) pair
// as the output value at dimension k. Works only when head_dim >= seq_len,
// but here head_dim=8 >= max(3,5)=5, so the trick works.
let q = vec![1.0f32; total * num_heads * head_dim];
let k = vec![1.0f32; total * num_heads * head_dim];
// V = identity: V[t, h, d] = if d == t % head_dim { 1.0 } else { 0.0 }
// (use global token index t, so different tokens map to different dims)
let mut v = vec![0.0f32; total * num_heads * head_dim];
for t in 0..total {
for h in 0..num_heads {
let d = t % head_dim; // unique dim per token (modulo to stay in range)
v[idx(t, h, d, num_heads, head_dim)] = 1.0;
}
}
let out = varlen_attention_cpu(&q, &k, &v, &cu_seqlens, num_heads, head_dim, scale, false);
// For each sequence, for each q-row, sum weights extracted via V=identity.
// Since q=k=const, the attention is uniform → weights = 1/seq_len.
// The output at dim d = weight for token whose unique dim is d.
// Sum of all output dims within a sequence = sum of all weights = 1.0.
for seq_idx in 0..2 {
let seq_start = cu_seqlens[seq_idx];
let seq_end = cu_seqlens[seq_idx + 1];
let seq_len = seq_end - seq_start;
for qi in 0..seq_len {
let global_token = seq_start + qi;
for h in 0..num_heads {
// Sum over dims that correspond to tokens in *this* sequence
let mut weight_sum = 0.0f32;
for ki in 0..seq_len {
let kt = seq_start + ki;
let d = kt % head_dim;
weight_sum += out[idx(global_token, h, d, num_heads, head_dim)];
}
assert!(
(weight_sum - 1.0).abs() < 1e-5,
"seq{seq_idx} q={qi} h={h}: weights sum to {weight_sum}, expected 1.0"
);
}
}
}
}
// -------------------------------------------------------------------------
// Test 7: Empty sequence — cu_seqlens = [0, 0, 4] → skip seq0, process seq1
// -------------------------------------------------------------------------
#[test]
fn test_varlen_empty_sequence_handled() {
let num_heads = 1;
let head_dim = 4;
let scale = 1.0 / (head_dim as f32).sqrt();
// seq0 is empty (0 tokens), seq1 has 4 tokens
let cu_seqlens = vec![0usize, 0, 4];
let total = 4; // only seq1 tokens exist
let q = vec![1.0f32; total * num_heads * head_dim];
let k = q.clone();
let v = q.clone();
// Must not panic — empty seq0 is silently skipped
let out = varlen_attention_cpu(&q, &k, &v, &cu_seqlens, num_heads, head_dim, scale, false);
assert_eq!(out.len(), total * num_heads * head_dim);
// seq1 tokens (indices 0..4 in the packed buffer) should have finite values
for i in 0..out.len() {
assert!(out[i].is_finite(), "non-finite at index {i}: {}", out[i]);
}
// For uniform q/k/v=1.0, output should be 1.0
for i in 0..out.len() {
assert!(
(out[i] - 1.0).abs() < 1e-5,
"expected 1.0 at index {i}, got {}",
out[i]
);
}
}
// -------------------------------------------------------------------------
// Test 8: Standard head_dim=64 (production default)
// -------------------------------------------------------------------------
#[test]
fn test_varlen_head_dim_64() {
let num_heads = 8;
let head_dim = 64;
let scale = 1.0 / (head_dim as f32).sqrt();
let cu_seqlens = vec![0usize, 7, 15]; // seq0=7, seq1=8
let total = 15;
let q: Vec<f32> = (0..total * num_heads * head_dim)
.map(|i| (i as f32) * 0.001)
.collect();
let k = q.clone();
let v: Vec<f32> = (0..total * num_heads * head_dim)
.map(|i| (i as f32) * 0.002)
.collect();
let out = varlen_attention_cpu(&q, &k, &v, &cu_seqlens, num_heads, head_dim, scale, false);
assert_eq!(out.len(), total * num_heads * head_dim);
for (i, &val) in out.iter().enumerate() {
assert!(val.is_finite(), "non-finite at index {i}: {val}");
}
}
}
```
- [ ] **Step 2: Verify the file exists**
```bash
ls -la /slab/projects/rustyverse/rustytorch/crates/training/rtx-flash-attention/src/kernels/flash_varlen_forward.rs
```
Expected: file with non-zero size.
---
## Task 3: Wire Up Exports and Add `SdpaBackend::VarLen`
**Files:**
- Modify: `crates/training/rtx-flash-attention/src/kernels/mod.rs:1-57`
- Modify: `crates/training/rtx-flash-attention/src/backend_selector.rs` (lines 3868 for the enum, lines 205216 for `supports_backend`, lines 5568 for `Display`)
**Interfaces:**
- Consumes: `flash_varlen_forward::varlen_attention_cpu` and `flash_varlen_forward::FlashVarlenKernel` from Task 2.
- Produces: `pub use kernels::varlen_attention_cpu` available at crate root (via `lib.rs` re-export chain); `SdpaBackend::VarLen` in the public enum.
- [ ] **Step 1: Add module declaration and re-exports to `mod.rs`**
Edit `src/kernels/mod.rs`. Add after the existing `#[cfg(feature = "cuda")] pub mod ptx;` block (line 14) and before the `#[cfg(feature = "metal")] pub mod metal;` block (line 18):
```rust
// varlen module is always compiled (contains CPU simulation unconditionally)
pub mod flash_varlen_forward;
pub use flash_varlen_forward::varlen_attention_cpu;
#[cfg(feature = "cuda")]
pub use flash_varlen_forward::FlashVarlenKernel;
```
The resulting block after the edit looks like:
```rust
//! GPU kernels for Flash Attention implementation
#[cfg(feature = "cuda")]
pub mod flash_forward;
#[cfg(feature = "cuda")]
pub mod flash_backward;
#[cfg(feature = "cuda")]
pub mod flash_v3_forward;
#[cfg(feature = "cuda")]
pub mod utils;
#[cfg(feature = "cuda")]
pub mod simple;
#[cfg(feature = "cuda")]
pub mod ptx;
#[cfg(feature = "cuda")]
pub mod manager;
// varlen module is always compiled (contains CPU simulation unconditionally)
pub mod flash_varlen_forward;
pub use flash_varlen_forward::varlen_attention_cpu;
#[cfg(feature = "cuda")]
pub use flash_varlen_forward::FlashVarlenKernel;
#[cfg(feature = "metal")]
pub mod metal;
#[cfg(all(test, feature = "cuda"))]
mod manager_test;
#[cfg(all(test, feature = "metal"))]
mod metal_test;
#[cfg(all(test, feature = "cuda"))]
mod test_kernel_launch;
#[cfg(all(test, feature = "cuda"))]
mod cudarc_api_test;
// Re-export the implementation based on feature
#[cfg(feature = "cuda")]
pub use simple::{FlashCudaKernels, KernelResult};
#[cfg(feature = "cuda")]
pub use flash_forward::{
FlashForwardKernel,
KernelExecutionResult,
KernelPerformanceInfo,
KernelPerformanceCache,
FP8Format,
FP8KernelConfig,
};
#[cfg(feature = "cuda")]
pub use flash_v3_forward::FlashV3ForwardKernel;
#[cfg(feature = "metal")]
pub use metal::{FlashMetalKernels, MetalKernelResult};
// The complex CUDA kernel implementation has been replaced with a simplified version
// for compilation purposes. This allows the advanced crates to compile while maintaining
// the required API surface. In production, this would be replaced with optimized CUDA kernels.
```
- [ ] **Step 2: Add `SdpaBackend::VarLen` to `backend_selector.rs`**
In `src/backend_selector.rs`, locate the `SdpaBackend` enum (lines 3854) and add `VarLen` as a new variant. Add it after `Cpu`:
```rust
/// Available SDPA backends
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum SdpaBackend {
/// FlashAttention v2 - optimal for long sequences
FlashAttention,
/// FlashAttention v3 - WGMMA + TMA + warp specialization (Hopper/Blackwell)
FlashAttentionV3,
/// Standard mathematical attention - simple, debuggable
Math,
/// Memory-efficient chunked attention
MemoryEfficient,
/// NVIDIA cuDNN attention (Ampere+)
CuDnn,
/// Custom Metal implementation for Apple Silicon
Metal,
/// Fallback CPU implementation
Cpu,
/// Variable-length packed-sequence attention (no padding)
VarLen,
}
```
- [ ] **Step 3: Add `Display` arm for `VarLen`**
In the `impl std::fmt::Display for SdpaBackend` block (lines 5668), add:
```rust
SdpaBackend::VarLen => write!(f, "VarLen"),
```
after the `SdpaBackend::Cpu` arm.
- [ ] **Step 4: Add `supports_backend` arm for `VarLen`**
In `HardwareCapabilities::supports_backend` (lines 206217), add:
```rust
SdpaBackend::VarLen => true, // CPU simulation always available; CUDA variant when feature is on
```
after the `SdpaBackend::Cpu => true,` arm.
- [ ] **Step 5: Run `cargo check` and confirm clean**
```bash
~/.cargo/bin/cargo check -p rtx-flash-attention 2>&1
```
Expected: zero errors, zero warnings (or only pre-existing warnings from other files).
- [ ] **Step 6: Run `cargo test --lib` and confirm 50 tests pass**
```bash
~/.cargo/bin/cargo test -p rtx-flash-attention --lib 2>&1 | tail -30
```
Expected output contains:
```
test result: ok. 50 passed; 0 failed; ...
```
(42 existing + 8 new varlen tests = 50)
- [ ] **Step 7: Commit**
```bash
git -C /slab/projects/rustyverse/rustytorch add \
crates/training/rtx-flash-attention/src/kernels/cuda/flash_attention_varlen.cu \
crates/training/rtx-flash-attention/src/kernels/flash_varlen_forward.rs \
crates/training/rtx-flash-attention/src/kernels/mod.rs \
crates/training/rtx-flash-attention/src/backend_selector.rs
git -C /slab/projects/rustyverse/rustytorch commit -m "$(cat <<'EOF'
feat(flash-attention): add varlen packed-sequence support
Implements variable-length (varlen) FlashAttention that processes
mixed-length batches without padding waste:
- New CUDA kernel flash_attention_varlen_forward with BLOCK_Q=64 /
BLOCK_K=64 tiling; grid=(ceil(max_seqlen_q/64), num_heads, 1).
Each block uses a linear scan over cu_seqlens_q to identify its
owning sequence and exits early when past sequence end.
- New Rust module flash_varlen_forward: always-compiled CPU simulation
(varlen_attention_cpu) for testing + #[cfg(cuda)] FlashVarlenKernel.
- SdpaBackend::VarLen variant added to backend_selector.
- 8 new CPU-only tests; total test count: 50.
Co-Authored-By: Claude Sonnet 4.6 <[email protected]>
EOF
)"
```
---
## Self-Review Checklist
### Spec Coverage
| Spec requirement | Task covering it |
|---|---|
| New CUDA kernel `flash_attention_varlen.cu` | Task 1 |
| Grid mapping: `(ceil(max_seqlen_q/BLOCK_Q), num_heads, 1)` | Task 1 Step 1 |
| Sequence identification by linear scan | Task 1 Step 1 |
| Early exit when tile is past sequence end | Task 1 Step 1 |
| BLOCK_Q=64, BLOCK_K=64 | Task 1 Step 1 |
| `FlashVarlenKernel` Rust struct with `new` + `forward` | Task 2 Step 1 |
| `varlen_attention_cpu` with correct signature | Task 2 Step 1 |
| Export from `src/kernels/mod.rs` | Task 3 Step 1 |
| `SdpaBackend::VarLen` | Task 3 Steps 24 |
| 8 tests — all CPU | Task 2 Step 1 (tests block) |
| `cargo check` clean | Task 3 Step 5 |
| `cargo test --lib` passes (50 total) | Task 3 Step 6 |
| No new dependencies | Verified — only existing crate features used |
### Placeholder Scan
No "TBD", "TODO", "implement later", or "similar to Task N" placeholders — all code blocks are complete and self-contained.
### Type Consistency
- `varlen_attention_cpu(q: &[f32], k: &[f32], v: &[f32], cu_seqlens: &[usize], num_heads: usize, head_dim: usize, softmax_scale: f32, causal: bool) -> Vec<f32>` — used consistently in both the implementation (Task 2) and all 8 test calls.
- `FlashVarlenKernel::new(ctx: &Arc<CudaContext>) -> FlashResult<Self>` — used in Task 2.
- `FlashVarlenKernel::forward(...)` — full 15-argument signature defined once in Task 2, re-exported in Task 3.
- `SdpaBackend::VarLen` — added in Task 3 Step 2, Display arm in Step 3, `supports_backend` arm in Step 4 — three places, all consistent spelling.