380 lines
12 KiB
Markdown
380 lines
12 KiB
Markdown
# Metal GPU Support for rtx-flash-attention
|
|
|
|
## Overview
|
|
|
|
This document describes the Metal GPU support implementation for the rtx-flash-attention crate, enabling Flash Attention on Apple Silicon and macOS devices.
|
|
|
|
## Implementation Status
|
|
|
|
### ✅ Completed
|
|
|
|
1. **Feature Configuration** (`Cargo.toml`)
|
|
- Added `metal` feature flag
|
|
- Added Metal dependencies (objc2, objc2-metal, objc2-foundation, block2)
|
|
- Feature-gated CUDA dependencies with `dep:` syntax
|
|
- Properly configured for macOS target
|
|
|
|
2. **CUDA Code Feature-Gating**
|
|
- All CUDA-specific code wrapped in `#[cfg(feature = "cuda")]`
|
|
- Clean separation between CUDA and Metal backends
|
|
- Files affected:
|
|
- `src/lib.rs` - Feature-gated CUDA imports
|
|
- `src/core.rs` - Separate impl blocks for CUDA and Metal
|
|
- `src/kernels/mod.rs` - Feature-gated kernel modules
|
|
- `src/kernels/flash_forward.rs` - CUDA-only
|
|
- `src/kernels/flash_backward.rs` - CUDA-only
|
|
- `src/kernels/simple.rs` - CUDA-only
|
|
|
|
3. **Metal Kernel Implementation** (`src/kernels/metal.rs`)
|
|
- `FlashMetalKernels` struct with Metal device and command queue
|
|
- `MetalKernelResult` for performance metrics
|
|
- Forward pass stub implementation
|
|
- Backward pass stub (returns not-implemented error)
|
|
- Block size optimization for Metal hardware
|
|
- Device information querying
|
|
- Tensor validation
|
|
|
|
4. **Metal Shader Source** (`metal/flash_attention.metal`)
|
|
- MSL (Metal Shading Language) shader structure
|
|
- Stub implementations showing:
|
|
- Forward kernel signature with threadgroup parameters
|
|
- Backward kernels (dQ, dK/dV) signatures
|
|
- Comprehensive implementation notes
|
|
- Ready for full tiled attention implementation
|
|
|
|
5. **Metal Backend in Core** (`src/core.rs`)
|
|
- Separate `FlashAttention` struct for Metal feature
|
|
- Full `FlashAttentionBackend` trait implementation
|
|
- Performance statistics tracking
|
|
- Validation and error handling
|
|
- Factory methods for Metal backend creation
|
|
|
|
6. **Build System** (`build.rs`)
|
|
- Feature-gated CUDA compilation
|
|
- Metal shader detection and path setup
|
|
- Proper warnings for missing backends
|
|
- Environment variable setup for shader paths
|
|
|
|
7. **Test Suite** (`src/kernels/metal_test.rs`)
|
|
- Device availability tests
|
|
- Kernel creation tests
|
|
- Block size optimization validation
|
|
- Thread safety tests
|
|
- Forward/backward pass behavior verification
|
|
- Multiple kernel instance tests
|
|
|
|
### Key Design Decisions
|
|
|
|
#### Feature Flag Strategy
|
|
```toml
|
|
[features]
|
|
default = ["cuda"]
|
|
cuda = ["rtx-kernel/cuda", "rtx-tensor/cuda", "dep:cudarc"]
|
|
metal = ["rtx-kernel/metal", "rtx-tensor/metal", "dep:objc2", "dep:objc2-metal", "dep:objc2-foundation", "dep:block2"]
|
|
```
|
|
|
|
- **Default**: CUDA for backward compatibility
|
|
- **Explicit dependencies**: Using `dep:` prefix for optional dependencies
|
|
- **Cascading features**: Enabling rtx-kernel and rtx-tensor GPU features
|
|
|
|
#### Backend Abstraction
|
|
```rust
|
|
pub trait FlashAttentionBackend {
|
|
async fn forward(...) -> FlashResult<FlashOutput>;
|
|
async fn backward(...) -> FlashResult<FlashGradOutput>;
|
|
fn name(&self) -> &str;
|
|
fn supports_config(&self, config: &FlashAttentionConfig) -> bool;
|
|
fn optimize_config(&self, config: FlashAttentionConfig) -> FlashResult<FlashAttentionConfig>;
|
|
}
|
|
```
|
|
|
|
- **Trait-based**: Clean abstraction over CUDA and Metal
|
|
- **Async**: Non-blocking GPU operations
|
|
- **Configuration**: Backend-specific optimization
|
|
|
|
#### Metal Hardware Considerations
|
|
- **Threadgroup memory**: 32KB typical (vs CUDA's 160KB on RTX 5090)
|
|
- **Max threads**: 1024 per threadgroup
|
|
- **Block sizes**: Optimized for 32-256 threads
|
|
- **No explicit tensor cores**: Metal abstracts this
|
|
|
|
## Usage
|
|
|
|
### Building with Metal Support
|
|
|
|
```bash
|
|
# Metal only (macOS)
|
|
cargo build -p rtx-flash-attention --no-default-features --features metal
|
|
|
|
# Both CUDA and Metal
|
|
cargo build -p rtx-flash-attention --features "cuda,metal"
|
|
|
|
# Default (CUDA only)
|
|
cargo build -p rtx-flash-attention
|
|
```
|
|
|
|
### Runtime Detection
|
|
|
|
```rust
|
|
use rtx_flash_attention::core::FlashAttentionFactory;
|
|
|
|
// Check backend availability
|
|
let cuda_available = FlashAttentionFactory::cuda_available();
|
|
let metal_available = FlashAttentionFactory::metal_available();
|
|
|
|
// Get device capabilities
|
|
let caps = FlashAttentionFactory::device_capabilities()?;
|
|
println!("GPU: {:?}, Memory: {} GB", caps.compute_capability, caps.total_memory / 1024 / 1024 / 1024);
|
|
```
|
|
|
|
### Creating Flash Attention with Metal
|
|
|
|
```rust
|
|
use rtx_flash_attention::{FlashAttention, FlashAttentionConfig};
|
|
|
|
// Create configuration
|
|
let config = FlashAttentionConfig::new(8, 64);
|
|
|
|
// Create Flash Attention instance (automatically selects Metal on macOS with metal feature)
|
|
let flash = FlashAttention::new(config)?;
|
|
|
|
// Run forward pass
|
|
let output = flash.forward(&q, &k, &v, false, 0.125).await?;
|
|
```
|
|
|
|
## Current Limitations
|
|
|
|
### Known Issues
|
|
|
|
1. **rtx-runtime Dependency**: The crate currently cannot compile with `--no-default-features --features metal` because rtx-runtime has hardcoded CUDA dependencies that are not properly feature-gated. This needs to be fixed in rtx-runtime first.
|
|
|
|
2. **Stub Implementation**: The Metal kernels are stubs that:
|
|
- Return zero tensors for forward pass
|
|
- Return not-implemented error for backward pass
|
|
- Don't perform actual attention computation
|
|
|
|
3. **No Tensor Integration**: Tests use rtx-tensor but Metal device support in rtx-tensor needs verification
|
|
|
|
### Next Steps for Full Implementation
|
|
|
|
#### 1. Complete Metal Shader Implementation
|
|
|
|
The shader stubs in `metal/flash_attention.metal` need full implementation:
|
|
|
|
```metal
|
|
kernel void flash_attention_forward(...) {
|
|
// 1. Load Q tile into threadgroup memory
|
|
threadgroup half Q_tile[BLOCK_SIZE_Q][HEAD_DIM];
|
|
|
|
// 2. Loop over K/V tiles
|
|
for (uint kv_block = 0; kv_block < num_kv_blocks; ++kv_block) {
|
|
// Load K, V tiles
|
|
|
|
// Compute attention scores: S = Q @ K^T
|
|
|
|
// Online softmax with running max/sum
|
|
float m_new = simd_max(m_old, row_max(scores));
|
|
float l_new = l_old * exp(m_old - m_new) + row_sum(exp(scores - m_new));
|
|
|
|
// Accumulate output: O = softmax(S) @ V
|
|
// Apply causal masking if needed
|
|
}
|
|
|
|
// Write final output and LSE
|
|
}
|
|
```
|
|
|
|
Key Metal APIs to use:
|
|
- `threadgroup` memory for tiles
|
|
- `simd_max()` / `simd_sum()` for reductions
|
|
- `simd_shuffle_down()` for warp-level communication
|
|
- `threadgroup_barrier()` for synchronization
|
|
|
|
#### 2. Kernel Loading and Compilation
|
|
|
|
Update `FlashMetalKernels` to actually compile and load shaders:
|
|
|
|
```rust
|
|
impl FlashMetalKernels {
|
|
pub fn new(config: &FlashAttentionConfig) -> FlashResult<Self> {
|
|
let device = unsafe { MTLCreateSystemDefaultDevice() }?;
|
|
let queue = device.newCommandQueue()?;
|
|
|
|
// Compile Metal shader
|
|
let shader_source = include_str!(env!("METAL_SHADER_PATH"));
|
|
let library = device.newLibraryWithSource_options_error(shader_source, None)?;
|
|
|
|
// Create pipeline states
|
|
let forward_function = library.newFunctionWithName(ns_string!("flash_attention_forward"))?;
|
|
let forward_pipeline = device.newComputePipelineStateWithFunction_error(forward_function)?;
|
|
|
|
Ok(Self { device, queue, forward_pipeline, ... })
|
|
}
|
|
}
|
|
```
|
|
|
|
#### 3. Buffer Management
|
|
|
|
Implement Metal buffer creation and data transfer:
|
|
|
|
```rust
|
|
pub async fn flash_attention_forward(
|
|
&self,
|
|
q: &Tensor,
|
|
k: &Tensor,
|
|
v: &Tensor,
|
|
output: &mut Tensor,
|
|
lse: &mut Tensor,
|
|
softmax_scale: f32,
|
|
causal: bool,
|
|
) -> FlashResult<MetalKernelResult> {
|
|
// Get Metal buffers from tensors
|
|
let q_buffer = q.metal_buffer()?;
|
|
let k_buffer = k.metal_buffer()?;
|
|
let v_buffer = v.metal_buffer()?;
|
|
let output_buffer = output.metal_buffer_mut()?;
|
|
let lse_buffer = lse.metal_buffer_mut()?;
|
|
|
|
// Create command buffer
|
|
let command_buffer = self.command_queue.commandBuffer()?;
|
|
let encoder = command_buffer.computeCommandEncoder()?;
|
|
|
|
// Set pipeline and buffers
|
|
encoder.setComputePipelineState(&self.forward_pipeline);
|
|
encoder.setBuffer_offset_atIndex(&q_buffer, 0, 0);
|
|
encoder.setBuffer_offset_atIndex(&k_buffer, 0, 1);
|
|
encoder.setBuffer_offset_atIndex(&v_buffer, 0, 2);
|
|
encoder.setBuffer_offset_atIndex(&output_buffer, 0, 3);
|
|
encoder.setBuffer_offset_atIndex(&lse_buffer, 0, 4);
|
|
|
|
// Dispatch
|
|
let grid_size = MTLSize::new(num_blocks, num_heads, batch_size);
|
|
let threadgroup_size = MTLSize::new(threads_per_group, 1, 1);
|
|
encoder.dispatchThreadgroups_threadsPerThreadgroup(grid_size, threadgroup_size);
|
|
|
|
encoder.endEncoding();
|
|
command_buffer.commit();
|
|
command_buffer.waitUntilCompleted();
|
|
|
|
Ok(MetalKernelResult { ... })
|
|
}
|
|
```
|
|
|
|
#### 4. Performance Optimization
|
|
|
|
Metal-specific optimizations:
|
|
- Use `simdgroup` functions for warp-level operations
|
|
- Optimize threadgroup memory layout for bank-conflict-free access
|
|
- Use Metal Performance Shaders (MPS) for matrix operations if beneficial
|
|
- Profile with Xcode Instruments Metal profiler
|
|
- Consider using `MTLGPUFamily` capabilities detection
|
|
|
|
#### 5. Fix rtx-runtime Dependencies
|
|
|
|
The rtx-runtime crate needs:
|
|
- Feature-gate all CUDA backend code with `#[cfg(feature = "cuda")]`
|
|
- Implement Metal stream and kernel launch abstractions
|
|
- Update `Stream::raw_stream()` to support Metal command queues
|
|
- Feature-gate kernel launch system
|
|
|
|
#### 6. Integration Tests
|
|
|
|
Create end-to-end tests:
|
|
```rust
|
|
#[tokio::test]
|
|
#[cfg(all(feature = "metal", target_os = "macos"))]
|
|
async fn test_metal_flash_attention_correctness() {
|
|
let config = FlashAttentionConfig::new(8, 64);
|
|
let flash = FlashAttention::new(config)?;
|
|
|
|
// Create random inputs
|
|
let device = Device::Metal(0);
|
|
let q = Tensor::randn(&[2, 8, 512, 64], &device)?;
|
|
let k = Tensor::randn(&[2, 8, 512, 64], &device)?;
|
|
let v = Tensor::randn(&[2, 8, 512, 64], &device)?;
|
|
|
|
// Run flash attention
|
|
let output = flash.forward(&q, &k, &v, false, 0.125).await?;
|
|
|
|
// Compare with naive attention
|
|
let expected = naive_attention(&q, &k, &v, 0.125)?;
|
|
|
|
assert_tensors_close(&output.output, &expected, 1e-3)?;
|
|
}
|
|
```
|
|
|
|
## Architecture Diagram
|
|
|
|
```
|
|
rtx-flash-attention
|
|
├── Cargo.toml (features: cuda, metal)
|
|
├── build.rs (feature-gated compilation)
|
|
├── metal/
|
|
│ └── flash_attention.metal (MSL shaders)
|
|
├── cuda/ (feature = "cuda")
|
|
│ ├── flash_attention_forward.cu
|
|
│ └── flash_attention_backward.cu
|
|
└── src/
|
|
├── lib.rs (public API)
|
|
├── core.rs (backend trait + impls)
|
|
├── config.rs (configuration)
|
|
├── error.rs (error types)
|
|
└── kernels/
|
|
├── mod.rs (feature-gated exports)
|
|
├── metal.rs (feature = "metal")
|
|
├── flash_forward.rs (feature = "cuda")
|
|
└── flash_backward.rs (feature = "cuda")
|
|
```
|
|
|
|
## Testing
|
|
|
|
### Run Metal Tests (macOS only)
|
|
|
|
```bash
|
|
# All Metal tests
|
|
cargo test -p rtx-flash-attention --features metal -- --test-threads=1
|
|
|
|
# Specific test
|
|
cargo test -p rtx-flash-attention --features metal test_metal_device_available
|
|
|
|
# With output
|
|
cargo test -p rtx-flash-attention --features metal -- --nocapture
|
|
```
|
|
|
|
### Expected Test Results (Current Stub Implementation)
|
|
|
|
- ✅ Device availability tests pass
|
|
- ✅ Kernel creation tests pass
|
|
- ✅ Block size optimization tests pass
|
|
- ⚠️ Forward pass returns zeros (stub)
|
|
- ❌ Backward pass returns not-implemented error
|
|
|
|
## References
|
|
|
|
### Metal Documentation
|
|
- [Metal Shading Language Specification](https://developer.apple.com/metal/Metal-Shading-Language-Specification.pdf)
|
|
- [Metal Best Practices Guide](https://developer.apple.com/documentation/metal/best_practices)
|
|
- [objc2-metal crate](https://docs.rs/objc2-metal/latest/objc2_metal/)
|
|
|
|
### Flash Attention
|
|
- [Flash Attention Paper](https://arxiv.org/abs/2205.14135) - Dao et al., 2022
|
|
- [Flash Attention 2](https://arxiv.org/abs/2307.08691) - Dao, 2023
|
|
|
|
### Similar Implementations
|
|
- [MLX Flash Attention](https://github.com/ml-explore/mlx) - Apple's ML framework with Metal Flash Attention
|
|
- [PyTorch Metal Backend](https://github.com/pytorch/pytorch/tree/main/aten/src/ATen/mps)
|
|
|
|
## Contributing
|
|
|
|
When implementing the full Metal backend:
|
|
|
|
1. **Follow TDD**: Write tests first, then implementation
|
|
2. **No Stubs**: All functions must have real implementations
|
|
3. **Performance First**: Optimize for Metal hardware characteristics
|
|
4. **Correctness**: Ensure numerical accuracy matches CUDA implementation
|
|
5. **Documentation**: Document all Metal-specific optimizations
|
|
|
|
## License
|
|
|
|
Same as rtx-flash-attention: MIT OR Apache-2.0
|