452 lines
16 KiB
Plaintext
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);
|
|
}
|
|
} |