Files
rustytorch/crates/training/rtx-flash-attention/cuda/flash_attention_backward.cu
T
2026-03-04 00:08:42 +00:00

452 lines
16 KiB
Plaintext

/*
* Flash Attention Backward CUDA Kernel
*
* Implements the backward pass of Flash Attention with O(n) memory complexity
* using SRAM tiling and efficient gradient computation.
*
* Reference: "FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness"
* https://arxiv.org/abs/2205.14135
*/
#include <cuda_runtime.h>
#include <cuda_fp16.h>
#include <mma.h>
#include <cooperative_groups.h>
// Use the configuration defines from Rust
#ifndef BLOCK_SIZE_Q
#define BLOCK_SIZE_Q 64
#endif
#ifndef BLOCK_SIZE_KV
#define BLOCK_SIZE_KV 64
#endif
#ifndef HEAD_DIM
#define HEAD_DIM 128
#endif
#ifndef NUM_HEADS
#define NUM_HEADS 32
#endif
#ifndef MAX_SEQ_LEN
#define MAX_SEQ_LEN 32768
#endif
// Constants
#define WARP_SIZE 32
#define MAX_THREADS_PER_BLOCK 1024
#define SHARED_MEM_ALIGNMENT 16
using namespace nvcuda;
namespace cg = cooperative_groups;
// Utility functions for half precision arithmetic
__device__ __forceinline__ float half_to_float(half x) {
return __half2float(x);
}
__device__ __forceinline__ half float_to_half(float x) {
return __float2half(x);
}
// Warp-level reduction for summing floats
__device__ __forceinline__ float warp_reduce_sum(float val) {
for (int offset = WARP_SIZE / 2; offset > 0; offset /= 2) {
val += __shfl_down_sync(0xffffffff, val, offset);
}
return val;
}
// Block-level reduction for summing floats
__device__ __forceinline__ float block_reduce_sum(float val) {
__shared__ float shared[32]; // Max warps per block
int warp_id = threadIdx.x / WARP_SIZE;
int lane_id = threadIdx.x % WARP_SIZE;
// Warp-level reduce
val = warp_reduce_sum(val);
// Store warp result
if (lane_id == 0) {
shared[warp_id] = val;
}
__syncthreads();
// Block-level reduce
if (warp_id == 0) {
val = (lane_id < (blockDim.x + WARP_SIZE - 1) / WARP_SIZE) ? shared[lane_id] : 0.0f;
val = warp_reduce_sum(val);
// Broadcast result
shared[0] = val;
}
__syncthreads();
return shared[0];
}
// Main Flash Attention backward kernel
extern "C" __global__ void flash_attention_backward_kernel(
const half* __restrict__ dout, // Gradient w.r.t output [batch * heads, seq_len, head_dim]
const half* __restrict__ q, // Query [batch * heads, seq_len, head_dim]
const half* __restrict__ k, // Key [batch * heads, seq_len, head_dim]
const half* __restrict__ v, // Value [batch * heads, seq_len, head_dim]
const half* __restrict__ o, // Forward output [batch * heads, seq_len, head_dim]
const float* __restrict__ lse, // Log-sum-exp [batch * heads, seq_len]
half* __restrict__ dq, // Query gradients (output) [batch * heads, seq_len, head_dim]
half* __restrict__ dk, // Key gradients (output) [batch * heads, seq_len, head_dim]
half* __restrict__ dv, // Value gradients (output) [batch * heads, seq_len, head_dim]
int seq_len,
int head_dim,
float softmax_scale,
unsigned int causal,
int block_size_q,
int block_size_kv
) {
// Shared memory for tiling
extern __shared__ char shared_mem[];
// Partition shared memory for backward pass
half* q_shared = reinterpret_cast<half*>(shared_mem);
half* k_shared = q_shared + BLOCK_SIZE_Q * HEAD_DIM;
half* v_shared = k_shared + BLOCK_SIZE_KV * HEAD_DIM;
half* dout_shared = v_shared + BLOCK_SIZE_KV * HEAD_DIM;
half* o_shared = dout_shared + BLOCK_SIZE_Q * HEAD_DIM;
float* scores_shared = reinterpret_cast<float*>(o_shared + BLOCK_SIZE_Q * HEAD_DIM);
float* ds_shared = scores_shared + BLOCK_SIZE_Q * BLOCK_SIZE_KV;
// Thread and block indices
int batch_head_idx = blockIdx.x;
int q_block_idx = blockIdx.y;
int tid = threadIdx.x;
int warp_id = tid / WARP_SIZE;
int lane_id = tid % WARP_SIZE;
// Calculate offsets
int tensor_offset = batch_head_idx * seq_len * head_dim;
int lse_offset = batch_head_idx * seq_len;
// Q block range
int q_start = q_block_idx * BLOCK_SIZE_Q;
int q_end = min(q_start + BLOCK_SIZE_Q, seq_len);
int q_size = q_end - q_start;
// Load Q, dout, and O blocks into shared memory
for (int i = tid; i < q_size * head_dim; i += blockDim.x) {
int q_row = i / head_dim;
int q_col = i % head_dim;
if (q_start + q_row < seq_len) {
int global_idx = tensor_offset + (q_start + q_row) * head_dim + q_col;
q_shared[q_row * head_dim + q_col] = q[global_idx];
dout_shared[q_row * head_dim + q_col] = dout[global_idx];
o_shared[q_row * head_dim + q_col] = o[global_idx];
}
}
__syncthreads();
// Initialize gradient accumulators
float dq_local[HEAD_DIM] = {0.0f};
// Compute row-wise dot product: dO * O for each query
float di_local[BLOCK_SIZE_Q] = {0.0f};
for (int q_local_idx = 0; q_local_idx < q_size; q_local_idx++) {
float sum = 0.0f;
for (int d = tid; d < head_dim; d += blockDim.x) {
float dout_val = half_to_float(dout_shared[q_local_idx * head_dim + d]);
float o_val = half_to_float(o_shared[q_local_idx * head_dim + d]);
sum += dout_val * o_val;
}
di_local[q_local_idx] = block_reduce_sum(sum);
}
__syncthreads();
// Iterate over KV blocks
for (int kv_block = 0; kv_block * BLOCK_SIZE_KV < seq_len; kv_block++) {
int kv_start = kv_block * BLOCK_SIZE_KV;
int kv_end = min(kv_start + BLOCK_SIZE_KV, seq_len);
int kv_size = kv_end - kv_start;
// Load K and V blocks into shared memory
for (int i = tid; i < kv_size * head_dim; i += blockDim.x) {
int kv_row = i / head_dim;
int kv_col = i % head_dim;
if (kv_start + kv_row < seq_len) {
int global_idx = tensor_offset + (kv_start + kv_row) * head_dim + kv_col;
k_shared[kv_row * head_dim + kv_col] = k[global_idx];
v_shared[kv_row * head_dim + kv_col] = v[global_idx];
}
}
__syncthreads();
// Recompute attention scores: Q @ K^T
for (int q_local_idx = 0; q_local_idx < q_size; q_local_idx++) {
for (int kv_idx = tid; kv_idx < kv_size; kv_idx += blockDim.x) {
float score = 0.0f;
// Dot product
for (int d = 0; d < head_dim; d++) {
float q_val = half_to_float(q_shared[q_local_idx * head_dim + d]);
float k_val = half_to_float(k_shared[kv_idx * head_dim + d]);
score += q_val * k_val;
}
score *= softmax_scale;
// Apply causal mask
int q_pos = q_start + q_local_idx;
int k_pos = kv_start + kv_idx;
if (causal && k_pos > q_pos) {
score = -INFINITY;
}
scores_shared[q_local_idx * BLOCK_SIZE_KV + kv_idx] = score;
}
}
__syncthreads();
// Compute softmax probabilities from scores and LSE
for (int q_local_idx = 0; q_local_idx < q_size; q_local_idx++) {
int q_global_idx = q_start + q_local_idx;
float lse_val = lse[lse_offset + q_global_idx];
for (int kv_idx = tid; kv_idx < kv_size; kv_idx += blockDim.x) {
float score = scores_shared[q_local_idx * BLOCK_SIZE_KV + kv_idx];
float prob = expf(score - lse_val);
scores_shared[q_local_idx * BLOCK_SIZE_KV + kv_idx] = prob;
}
}
__syncthreads();
// Compute dS = P * (dO @ V^T - di)
for (int q_local_idx = 0; q_local_idx < q_size; q_local_idx++) {
for (int kv_idx = tid; kv_idx < kv_size; kv_idx += blockDim.x) {
// Compute dO @ V^T for this (q, kv) pair
float dout_v_sum = 0.0f;
for (int d = 0; d < head_dim; d++) {
float dout_val = half_to_float(dout_shared[q_local_idx * head_dim + d]);
float v_val = half_to_float(v_shared[kv_idx * head_dim + d]);
dout_v_sum += dout_val * v_val;
}
float prob = scores_shared[q_local_idx * BLOCK_SIZE_KV + kv_idx];
float ds_val = prob * (dout_v_sum - di_local[q_local_idx]);
ds_shared[q_local_idx * BLOCK_SIZE_KV + kv_idx] = ds_val;
}
}
__syncthreads();
// Compute dQ: dQ += dS @ K
for (int q_local_idx = 0; q_local_idx < q_size; q_local_idx++) {
for (int d = tid; d < head_dim; d += blockDim.x) {
float dq_val = 0.0f;
for (int kv_idx = 0; kv_idx < kv_size; kv_idx++) {
float ds_val = ds_shared[q_local_idx * BLOCK_SIZE_KV + kv_idx];
float k_val = half_to_float(k_shared[kv_idx * head_dim + d]);
dq_val += ds_val * k_val;
}
dq_local[d] += dq_val * softmax_scale;
}
}
// Compute dK: dK += dS^T @ Q (accumulate across Q blocks)
for (int kv_idx = 0; kv_idx < kv_size; kv_idx++) {
float dk_local[HEAD_DIM] = {0.0f};
for (int d = tid; d < head_dim; d += blockDim.x) {
float dk_val = 0.0f;
for (int q_local_idx = 0; q_local_idx < q_size; q_local_idx++) {
float ds_val = ds_shared[q_local_idx * BLOCK_SIZE_KV + kv_idx];
float q_val = half_to_float(q_shared[q_local_idx * head_dim + d]);
dk_val += ds_val * q_val;
}
dk_local[d] = dk_val * softmax_scale;
}
// Atomic accumulation to global dK (multiple Q blocks contribute)
int kv_global_idx = kv_start + kv_idx;
if (kv_global_idx < seq_len) {
for (int d = tid; d < head_dim; d += blockDim.x) {
half* dk_ptr = &dk[tensor_offset + kv_global_idx * head_dim + d];
atomicAdd(dk_ptr, float_to_half(dk_local[d]));
}
}
}
// Compute dV: dV += P^T @ dO (accumulate across Q blocks)
for (int kv_idx = 0; kv_idx < kv_size; kv_idx++) {
float dv_local[HEAD_DIM] = {0.0f};
for (int d = tid; d < head_dim; d += blockDim.x) {
float dv_val = 0.0f;
for (int q_local_idx = 0; q_local_idx < q_size; q_local_idx++) {
float prob = scores_shared[q_local_idx * BLOCK_SIZE_KV + kv_idx];
float dout_val = half_to_float(dout_shared[q_local_idx * head_dim + d]);
dv_val += prob * dout_val;
}
dv_local[d] = dv_val;
}
// Atomic accumulation to global dV (multiple Q blocks contribute)
int kv_global_idx = kv_start + kv_idx;
if (kv_global_idx < seq_len) {
for (int d = tid; d < head_dim; d += blockDim.x) {
half* dv_ptr = &dv[tensor_offset + kv_global_idx * head_dim + d];
atomicAdd(dv_ptr, float_to_half(dv_local[d]));
}
}
}
__syncthreads();
}
// Store dQ gradients
for (int q_local_idx = 0; q_local_idx < q_size; q_local_idx++) {
int q_global_idx = q_start + q_local_idx;
if (q_global_idx < seq_len) {
for (int d = tid; d < head_dim; d += blockDim.x) {
dq[tensor_offset + q_global_idx * head_dim + d] = float_to_half(dq_local[d]);
}
}
}
}
// Specialized kernel for small sequences
extern "C" __global__ void flash_attention_backward_small_kernel(
const half* __restrict__ dout,
const half* __restrict__ q,
const half* __restrict__ k,
const half* __restrict__ v,
const half* __restrict__ o,
const float* __restrict__ lse,
half* __restrict__ dq,
half* __restrict__ dk,
half* __restrict__ dv,
int seq_len,
int head_dim,
float softmax_scale,
unsigned int causal
) {
// For small sequences, we can fit everything in shared memory
extern __shared__ char shared_mem[];
// Partition shared memory
half* q_all = reinterpret_cast<half*>(shared_mem);
half* k_all = q_all + seq_len * head_dim;
half* v_all = k_all + seq_len * head_dim;
half* dout_all = v_all + seq_len * head_dim;
half* o_all = dout_all + seq_len * head_dim;
float* scores_all = reinterpret_cast<float*>(o_all + seq_len * head_dim);
float* ds_all = scores_all + seq_len * seq_len;
int batch_head_idx = blockIdx.x;
int tid = threadIdx.x;
// Load all tensors into shared memory
int offset = batch_head_idx * seq_len * head_dim;
for (int i = tid; i < seq_len * head_dim; i += blockDim.x) {
q_all[i] = q[offset + i];
k_all[i] = k[offset + i];
v_all[i] = v[offset + i];
dout_all[i] = dout[offset + i];
o_all[i] = o[offset + i];
}
__syncthreads();
// Recompute attention scores and probabilities
for (int i = tid; i < seq_len * seq_len; i += blockDim.x) {
int q_idx = i / seq_len;
int k_idx = i % seq_len;
float score = 0.0f;
for (int d = 0; d < head_dim; d++) {
score += half_to_float(q_all[q_idx * head_dim + d]) * half_to_float(k_all[k_idx * head_dim + d]);
}
score *= softmax_scale;
// Apply causal mask
if (causal && k_idx > q_idx) {
score = -INFINITY;
}
float lse_val = lse[batch_head_idx * seq_len + q_idx];
float prob = expf(score - lse_val);
scores_all[i] = prob;
}
__syncthreads();
// Compute di = dO * O for each query
__shared__ float di[MAX_SEQ_LEN];
for (int q_idx = tid; q_idx < seq_len; q_idx += blockDim.x) {
float sum = 0.0f;
for (int d = 0; d < head_dim; d++) {
sum += half_to_float(dout_all[q_idx * head_dim + d]) * half_to_float(o_all[q_idx * head_dim + d]);
}
di[q_idx] = sum;
}
__syncthreads();
// Compute dS = P * (dO @ V^T - di)
for (int i = tid; i < seq_len * seq_len; i += blockDim.x) {
int q_idx = i / seq_len;
int k_idx = i % seq_len;
float dout_v_sum = 0.0f;
for (int d = 0; d < head_dim; d++) {
dout_v_sum += half_to_float(dout_all[q_idx * head_dim + d]) * half_to_float(v_all[k_idx * head_dim + d]);
}
float prob = scores_all[i];
float ds_val = prob * (dout_v_sum - di[q_idx]);
ds_all[i] = ds_val;
}
__syncthreads();
// Compute gradients
// dQ = dS @ K
for (int i = tid; i < seq_len * head_dim; i += blockDim.x) {
int q_idx = i / head_dim;
int d = i % head_dim;
float dq_val = 0.0f;
for (int k_idx = 0; k_idx < seq_len; k_idx++) {
float ds_val = ds_all[q_idx * seq_len + k_idx];
float k_val = half_to_float(k_all[k_idx * head_dim + d]);
dq_val += ds_val * k_val;
}
dq[offset + i] = float_to_half(dq_val * softmax_scale);
}
// dK = dS^T @ Q
for (int i = tid; i < seq_len * head_dim; i += blockDim.x) {
int k_idx = i / head_dim;
int d = i % head_dim;
float dk_val = 0.0f;
for (int q_idx = 0; q_idx < seq_len; q_idx++) {
float ds_val = ds_all[q_idx * seq_len + k_idx];
float q_val = half_to_float(q_all[q_idx * head_dim + d]);
dk_val += ds_val * q_val;
}
dk[offset + i] = float_to_half(dk_val * softmax_scale);
}
// dV = P^T @ dO
for (int i = tid; i < seq_len * head_dim; i += blockDim.x) {
int v_idx = i / head_dim;
int d = i % head_dim;
float dv_val = 0.0f;
for (int q_idx = 0; q_idx < seq_len; q_idx++) {
float prob = scores_all[q_idx * seq_len + v_idx];
float dout_val = half_to_float(dout_all[q_idx * head_dim + d]);
dv_val += prob * dout_val;
}
dv[offset + i] = float_to_half(dv_val);
}
}