Initial commit

This commit is contained in:
redclawsystems
2026-03-04 00:08:42 +00:00
commit 4d88dc0584
4449 changed files with 1556714 additions and 0 deletions
@@ -0,0 +1,339 @@
//! 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,
&params,
&q_buffer,
&k_buffer,
&v_buffer,
&o_buffer,
&do_buffer,
&lse_buffer,
&dq_buffer,
)?;
// Execute dK/dV kernel
execute_dkv_kernel(
attn,
&params,
&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<ProtocolObject<dyn MTLBuffer>>,
k_buffer: &Retained<ProtocolObject<dyn MTLBuffer>>,
v_buffer: &Retained<ProtocolObject<dyn MTLBuffer>>,
o_buffer: &Retained<ProtocolObject<dyn MTLBuffer>>,
do_buffer: &Retained<ProtocolObject<dyn MTLBuffer>>,
lse_buffer: &Retained<ProtocolObject<dyn MTLBuffer>>,
dq_buffer: &Retained<ProtocolObject<dyn MTLBuffer>>,
) -> 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::<BackwardParams>(params) as *mut std::ffi::c_void);
encoder.setBytes_length_atIndex(params_ptr, std::mem::size_of::<BackwardParams>(), 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<ProtocolObject<dyn MTLBuffer>>,
k_buffer: &Retained<ProtocolObject<dyn MTLBuffer>>,
v_buffer: &Retained<ProtocolObject<dyn MTLBuffer>>,
o_buffer: &Retained<ProtocolObject<dyn MTLBuffer>>,
do_buffer: &Retained<ProtocolObject<dyn MTLBuffer>>,
lse_buffer: &Retained<ProtocolObject<dyn MTLBuffer>>,
dk_buffer: &Retained<ProtocolObject<dyn MTLBuffer>>,
dv_buffer: &Retained<ProtocolObject<dyn MTLBuffer>>,
) -> 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::<BackwardParams>(params) as *mut std::ffi::c_void);
encoder.setBytes_length_atIndex(params_ptr, std::mem::size_of::<BackwardParams>(), 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(())
}