Files
rustytorch/crates/training/rtx-flash-metal-attention/src/forward.rs
T
osobhandClaude Opus 4.6 02d382d5f6 style: apply rustfmt across all crates and demos
Consistent formatting pass: line wrapping, import sorting, trailing
whitespace removal, let-chain indentation, merged derive attributes,
and unsafe block reformatting.

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>
2026-04-12 07:01:58 -07:00

204 lines
6.0 KiB
Rust

//! Flash Attention forward pass implementation
use crate::FlashAttention;
use crate::error::{FlashError, FlashResult};
use objc2_metal::{
MTLCommandBuffer, MTLCommandEncoder, MTLCommandQueue, MTLComputeCommandEncoder, MTLSize,
};
use rtx_tensor::Tensor;
use std::ptr::NonNull;
use tracing::debug;
/// Parameters passed to the Metal forward kernel
#[repr(C)]
struct ForwardParams {
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 forward pass
///
/// Computes: O = softmax(Q @ K^T * scale) @ V
///
/// # Arguments
/// * `attn` - Flash Attention instance with compiled pipelines
/// * `q` - Query tensor [batch, heads, seq_q, head_dim]
/// * `k` - Key tensor [batch, heads, seq_kv, head_dim]
/// * `v` - Value tensor [batch, heads, seq_kv, head_dim]
///
/// # Returns
/// * Output tensor [batch, heads, seq_q, head_dim]
/// * Log-sum-exp tensor [batch, heads, seq_q] for backward pass
pub fn flash_attention_forward(
attn: &FlashAttention,
q: &Tensor,
k: &Tensor,
v: &Tensor,
) -> FlashResult<(Tensor, Tensor)> {
// Validate input shapes
let q_shape = q.shape();
let k_shape = k.shape();
let v_shape = v.shape();
if q_shape.len() != 4 {
return Err(FlashError::shape(format!(
"Q must be 4D [batch, heads, seq, head_dim], 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 K shape
if k_shape != [batch, heads, seq_kv, head_dim] {
return Err(FlashError::dim_mismatch(format!(
"K shape {:?} incompatible with Q shape {:?}",
k_shape, q_shape
)));
}
// Validate V shape
if v_shape != [batch, heads, seq_kv, head_dim] {
return Err(FlashError::dim_mismatch(format!(
"V shape {:?} incompatible with K shape {:?}",
v_shape, k_shape
)));
}
// Validate head_dim
attn.config
.validate(head_dim)
.map_err(FlashError::Configuration)?;
debug!(
"Flash attention forward: batch={}, heads={}, seq_q={}, seq_kv={}, head_dim={}",
batch, heads, seq_q, seq_kv, head_dim
);
// Create output tensors (same dtype as input for output, F32 for LSE)
let output = Tensor::zeros_typed([batch, heads, seq_q, head_dim], q.dtype(), q.device())
.map_err(|e| FlashError::execution(format!("Failed to allocate output: {}", e)))?;
let lse = Tensor::zeros([batch, heads, seq_q], q.device())
.map_err(|e| FlashError::execution(format!("Failed to allocate LSE: {}", 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 (lse_buffer, _) = lse
.storage_ref()
.get_metal_data()
.ok_or_else(|| FlashError::not_metal("LSE tensor is not on Metal device"))?;
// Create command buffer
let cmd_buffer = attn
.command_queue
.commandBuffer()
.ok_or_else(|| FlashError::device("Failed to create command buffer"))?;
// Create compute encoder
let encoder = cmd_buffer
.computeCommandEncoder()
.ok_or_else(|| FlashError::device("Failed to create compute encoder"))?;
// Set pipeline
encoder.setComputePipelineState(&attn.forward_pipeline);
// Bind buffers
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(&lse_buffer), 0, 4);
}
// Set parameters
let params = ForwardParams {
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,
};
unsafe {
let params_ptr = NonNull::new_unchecked(&raw const params as *mut std::ffi::c_void);
encoder.setBytes_length_atIndex(params_ptr, std::mem::size_of::<ForwardParams>(), 5);
}
// Calculate grid dimensions
let block_q = attn.config.block_q;
let num_q_blocks = (seq_q + block_q - 1) / block_q;
let grid = MTLSize {
width: num_q_blocks,
height: heads,
depth: batch,
};
let threadgroup = MTLSize {
width: block_q,
height: 1,
depth: 1,
};
debug!(
"Dispatching forward kernel: grid=({}, {}, {}), threadgroup=({}, {}, {})",
grid.width,
grid.height,
grid.depth,
threadgroup.width,
threadgroup.height,
threadgroup.depth
);
encoder.dispatchThreadgroups_threadsPerThreadgroup(grid, threadgroup);
encoder.endEncoding();
// Execute and wait
cmd_buffer.commit();
cmd_buffer.waitUntilCompleted();
// Check for errors
if let Some(error) = cmd_buffer.error() {
return Err(FlashError::execution(format!(
"Command buffer execution failed: {:?}",
error
)));
}
Ok((output, lse))
}
#[cfg(test)]
mod tests {
// Tests would go here, but require Metal device
// In practice, test against CPU reference implementation
}