340 lines
11 KiB
Rust
340 lines
11 KiB
Rust
//! 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<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(())
|
|
}
|