Initial commit
This commit is contained in:
@@ -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,
|
||||
¶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(())
|
||||
}
|
||||
Reference in New Issue
Block a user