//! Flash Attention backward pass implementation use crate::FlashAttention; use crate::error::{FlashError, FlashResult}; use objc2::rc::Retained; use objc2::runtime::ProtocolObject; use objc2_metal::{ MTLBuffer, MTLCommandBuffer, MTLCommandEncoder, MTLCommandQueue, MTLComputeCommandEncoder, MTLSize, }; use rtx_tensor::Tensor; use std::ptr::NonNull; use tracing::debug; /// Parameters passed to the Metal backward kernels #[repr(C)] struct BackwardParams { batch_size: u32, num_heads: u32, seq_len_q: u32, seq_len_kv: u32, head_dim: u32, softmax_scale: f32, causal: u32, } /// Execute Flash Attention backward pass /// /// Computes gradients dQ, dK, dV given the upstream gradient dO. /// /// # Arguments /// * `attn` - Flash Attention instance with compiled pipelines /// * `grad_output` - Gradient of loss w.r.t. output [batch, heads, seq_q, head_dim] /// * `q` - Query tensor from forward pass [batch, heads, seq_q, head_dim] /// * `k` - Key tensor from forward pass [batch, heads, seq_kv, head_dim] /// * `v` - Value tensor from forward pass [batch, heads, seq_kv, head_dim] /// * `output` - Output tensor from forward pass [batch, heads, seq_q, head_dim] /// * `lse` - Log-sum-exp from forward pass [batch, heads, seq_q] /// /// # Returns /// * `(dQ, dK, dV)` - Gradients with same shapes as Q, K, V /// /// # Errors /// Returns an error if: /// - Tensor shapes are invalid or mismatched /// - Tensors are not on a Metal device /// - Kernel execution fails pub fn flash_attention_backward( attn: &FlashAttention, grad_output: &Tensor, q: &Tensor, k: &Tensor, v: &Tensor, output: &Tensor, lse: &Tensor, ) -> FlashResult<(Tensor, Tensor, Tensor)> { // Validate shapes let q_shape = q.shape(); let k_shape = k.shape(); let do_shape = grad_output.shape(); if q_shape.len() != 4 { return Err(FlashError::shape(format!( "Q must be 4D, got {:?}", q_shape ))); } let batch = q_shape[0]; let heads = q_shape[1]; let seq_q = q_shape[2]; let head_dim = q_shape[3]; let seq_kv = k_shape[2]; // Validate grad_output shape if do_shape != q_shape { return Err(FlashError::dim_mismatch(format!( "grad_output shape {:?} must match Q shape {:?}", do_shape, q_shape ))); } // Validate output shape if output.shape() != q_shape { return Err(FlashError::dim_mismatch(format!( "output shape {:?} must match Q shape {:?}", output.shape(), q_shape ))); } // Validate LSE shape let expected_lse_shape = vec![batch, heads, seq_q]; if lse.shape() != expected_lse_shape { return Err(FlashError::dim_mismatch(format!( "LSE shape {:?} must be {:?}", lse.shape(), expected_lse_shape ))); } debug!( "Flash attention backward: batch={}, heads={}, seq_q={}, seq_kv={}, head_dim={}", batch, heads, seq_q, seq_kv, head_dim ); // Allocate gradient tensors (same dtype as inputs) let dq = Tensor::zeros_typed([batch, heads, seq_q, head_dim], q.dtype(), q.device()) .map_err(|e| FlashError::execution(format!("Failed to allocate dQ: {}", e)))?; let dk = Tensor::zeros_typed([batch, heads, seq_kv, head_dim], k.dtype(), k.device()) .map_err(|e| FlashError::execution(format!("Failed to allocate dK: {}", e)))?; let dv = Tensor::zeros_typed([batch, heads, seq_kv, head_dim], v.dtype(), v.device()) .map_err(|e| FlashError::execution(format!("Failed to allocate dV: {}", e)))?; // Get Metal buffers via storage API let (q_buffer, _) = q .storage_ref() .get_metal_data() .ok_or_else(|| FlashError::not_metal("Q tensor is not on Metal device"))?; let (k_buffer, _) = k .storage_ref() .get_metal_data() .ok_or_else(|| FlashError::not_metal("K tensor is not on Metal device"))?; let (v_buffer, _) = v .storage_ref() .get_metal_data() .ok_or_else(|| FlashError::not_metal("V tensor is not on Metal device"))?; let (o_buffer, _) = output .storage_ref() .get_metal_data() .ok_or_else(|| FlashError::not_metal("Output tensor is not on Metal device"))?; let (do_buffer, _) = grad_output .storage_ref() .get_metal_data() .ok_or_else(|| FlashError::not_metal("grad_output tensor is not on Metal device"))?; let (lse_buffer, _) = lse .storage_ref() .get_metal_data() .ok_or_else(|| FlashError::not_metal("LSE tensor is not on Metal device"))?; let (dq_buffer, _) = dq .storage_ref() .get_metal_data() .ok_or_else(|| FlashError::not_metal("dQ tensor is not on Metal device"))?; let (dk_buffer, _) = dk .storage_ref() .get_metal_data() .ok_or_else(|| FlashError::not_metal("dK tensor is not on Metal device"))?; let (dv_buffer, _) = dv .storage_ref() .get_metal_data() .ok_or_else(|| FlashError::not_metal("dV tensor is not on Metal device"))?; let params = BackwardParams { batch_size: batch as u32, num_heads: heads as u32, seq_len_q: seq_q as u32, seq_len_kv: seq_kv as u32, head_dim: head_dim as u32, softmax_scale: attn.config.get_softmax_scale(head_dim), causal: attn.config.causal as u32, }; // Execute dQ kernel execute_dq_kernel( attn, ¶ms, &q_buffer, &k_buffer, &v_buffer, &o_buffer, &do_buffer, &lse_buffer, &dq_buffer, )?; // Execute dK/dV kernel execute_dkv_kernel( attn, ¶ms, &q_buffer, &k_buffer, &v_buffer, &o_buffer, &do_buffer, &lse_buffer, &dk_buffer, &dv_buffer, )?; Ok((dq, dk, dv)) } /// Execute the dQ backward kernel fn execute_dq_kernel( attn: &FlashAttention, params: &BackwardParams, q_buffer: &Retained>, k_buffer: &Retained>, v_buffer: &Retained>, o_buffer: &Retained>, do_buffer: &Retained>, lse_buffer: &Retained>, dq_buffer: &Retained>, ) -> FlashResult<()> { let cmd_buffer = attn .command_queue .commandBuffer() .ok_or_else(|| FlashError::device("Failed to create command buffer for dQ"))?; let encoder = cmd_buffer .computeCommandEncoder() .ok_or_else(|| FlashError::device("Failed to create encoder for dQ"))?; encoder.setComputePipelineState(&attn.backward_dq_pipeline); unsafe { encoder.setBuffer_offset_atIndex(Some(q_buffer), 0, 0); encoder.setBuffer_offset_atIndex(Some(k_buffer), 0, 1); encoder.setBuffer_offset_atIndex(Some(v_buffer), 0, 2); encoder.setBuffer_offset_atIndex(Some(o_buffer), 0, 3); encoder.setBuffer_offset_atIndex(Some(do_buffer), 0, 4); encoder.setBuffer_offset_atIndex(Some(lse_buffer), 0, 5); encoder.setBuffer_offset_atIndex(Some(dq_buffer), 0, 6); let params_ptr = NonNull::new_unchecked(std::ptr::from_ref::(params) as *mut std::ffi::c_void); encoder.setBytes_length_atIndex(params_ptr, std::mem::size_of::(), 7); } let block_q = attn.config.block_q; let num_q_blocks = (params.seq_len_q as usize + block_q - 1) / block_q; let grid = MTLSize { width: num_q_blocks, height: params.num_heads as usize, depth: params.batch_size as usize, }; let threadgroup = MTLSize { width: block_q, height: 1, depth: 1, }; debug!( "Dispatching dQ kernel: grid=({}, {}, {})", grid.width, grid.height, grid.depth ); encoder.dispatchThreadgroups_threadsPerThreadgroup(grid, threadgroup); encoder.endEncoding(); cmd_buffer.commit(); cmd_buffer.waitUntilCompleted(); if let Some(error) = cmd_buffer.error() { return Err(FlashError::execution(format!( "dQ kernel failed: {:?}", error ))); } Ok(()) } /// Execute the dK/dV backward kernel fn execute_dkv_kernel( attn: &FlashAttention, params: &BackwardParams, q_buffer: &Retained>, k_buffer: &Retained>, v_buffer: &Retained>, o_buffer: &Retained>, do_buffer: &Retained>, lse_buffer: &Retained>, dk_buffer: &Retained>, dv_buffer: &Retained>, ) -> FlashResult<()> { let cmd_buffer = attn .command_queue .commandBuffer() .ok_or_else(|| FlashError::device("Failed to create command buffer for dKV"))?; let encoder = cmd_buffer .computeCommandEncoder() .ok_or_else(|| FlashError::device("Failed to create encoder for dKV"))?; encoder.setComputePipelineState(&attn.backward_dkv_pipeline); unsafe { encoder.setBuffer_offset_atIndex(Some(q_buffer), 0, 0); encoder.setBuffer_offset_atIndex(Some(k_buffer), 0, 1); encoder.setBuffer_offset_atIndex(Some(v_buffer), 0, 2); encoder.setBuffer_offset_atIndex(Some(o_buffer), 0, 3); encoder.setBuffer_offset_atIndex(Some(do_buffer), 0, 4); encoder.setBuffer_offset_atIndex(Some(lse_buffer), 0, 5); encoder.setBuffer_offset_atIndex(Some(dk_buffer), 0, 6); encoder.setBuffer_offset_atIndex(Some(dv_buffer), 0, 7); let params_ptr = NonNull::new_unchecked(std::ptr::from_ref::(params) as *mut std::ffi::c_void); encoder.setBytes_length_atIndex(params_ptr, std::mem::size_of::(), 8); } let block_kv = attn.config.block_kv; let num_kv_blocks = (params.seq_len_kv as usize + block_kv - 1) / block_kv; let grid = MTLSize { width: num_kv_blocks, height: params.num_heads as usize, depth: params.batch_size as usize, }; let threadgroup = MTLSize { width: block_kv, height: 1, depth: 1, }; debug!( "Dispatching dKV kernel: grid=({}, {}, {})", grid.width, grid.height, grid.depth ); encoder.dispatchThreadgroups_threadsPerThreadgroup(grid, threadgroup); encoder.endEncoding(); cmd_buffer.commit(); cmd_buffer.waitUntilCompleted(); if let Some(error) = cmd_buffer.error() { return Err(FlashError::execution(format!( "dKV kernel failed: {:?}", error ))); } Ok(()) }