/* * 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 #include #include #include // 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(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(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(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(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); } }