//! Metal Compute Pipeline and Kernel Management //! //! This module provides compute pipeline creation, shader compilation, //! and kernel execution for Metal compute operations. use crate::device::MetalDevice; use crate::error::{MetalError, Result}; use crate::memory::MetalBuffer; use dashmap::DashMap; use std::sync::Arc; use tracing::{debug, info}; #[cfg(target_os = "macos")] use objc2::rc::Retained; #[cfg(target_os = "macos")] use objc2::runtime::ProtocolObject; #[cfg(target_os = "macos")] use objc2_foundation::NSString; #[cfg(target_os = "macos")] use objc2_metal::{ MTLCommandBuffer, MTLCommandQueue, MTLComputeCommandEncoder, MTLComputePipelineState, MTLDevice, MTLLibrary, MTLSize, MTLCommandEncoder, }; /// Metal kernel launch configuration #[derive(Debug, Clone, Copy)] pub struct MetalLaunchConfig { /// Total threads in grid (width, height, depth) pub grid_size: (usize, usize, usize), /// Threads per threadgroup (width, height, depth) pub threadgroup_size: (usize, usize, usize), } impl MetalLaunchConfig { /// Create a 1D launch configuration pub fn new_1d(total_threads: usize, threads_per_group: usize) -> Self { Self { grid_size: (total_threads, 1, 1), threadgroup_size: (threads_per_group.min(1024), 1, 1), } } /// Create a 2D launch configuration pub fn new_2d(width: usize, height: usize, tile_size: usize) -> Self { let tile = tile_size.min(32); Self { grid_size: (width, height, 1), threadgroup_size: (tile, tile, 1), } } /// Validate the launch configuration pub fn validate(&self) -> Result<()> { let (gx, gy, gz) = self.grid_size; let (tx, ty, tz) = self.threadgroup_size; if gx == 0 || gy == 0 || gz == 0 { return Err(MetalError::InvalidLaunchConfig( "Grid size dimensions must be positive".into() )); } if tx == 0 || ty == 0 || tz == 0 { return Err(MetalError::InvalidLaunchConfig( "Threadgroup size dimensions must be positive".into() )); } let total_threads = tx * ty * tz; if total_threads > 1024 { return Err(MetalError::InvalidLaunchConfig(format!( "Too many threads per threadgroup: {total_threads} (max 1024)" ))); } Ok(()) } } /// Metal command queue for submitting work to the GPU #[cfg(target_os = "macos")] pub struct MetalCommandQueue { queue: Retained>, } #[cfg(not(target_os = "macos"))] pub struct MetalCommandQueue { _private: (), } #[cfg(target_os = "macos")] impl MetalCommandQueue { /// Create a new command queue from a device pub fn new(device: &MetalDevice) -> Result { let queue = device.new_command_queue()?; Ok(Self { queue }) } /// Get the underlying MTLCommandQueue pub fn mtl_queue(&self) -> &ProtocolObject { &self.queue } /// Create a new command buffer pub fn command_buffer(&self) -> Result>> { self.queue.commandBuffer() .ok_or_else(|| MetalError::CommandBuffer("Failed to create command buffer".into())) } } #[cfg(not(target_os = "macos"))] impl MetalCommandQueue { /// Create a new command queue (stub) pub fn new(_device: &MetalDevice) -> Result { Err(MetalError::NotAvailable) } } /// Compiled Metal compute kernel #[cfg(target_os = "macos")] pub struct MetalKernel { /// Kernel function name name: String, /// Compiled pipeline state pipeline: Retained>, /// Maximum total threads per threadgroup max_threads_per_threadgroup: usize, /// Thread execution width (SIMD group size) thread_execution_width: usize, } #[cfg(not(target_os = "macos"))] pub struct MetalKernel { name: String, } #[cfg(target_os = "macos")] impl MetalKernel { /// Get kernel name pub fn name(&self) -> &str { &self.name } /// Get maximum threads per threadgroup for this kernel pub fn max_threads_per_threadgroup(&self) -> usize { self.max_threads_per_threadgroup } /// Get thread execution width (SIMD group size) pub fn thread_execution_width(&self) -> usize { self.thread_execution_width } /// Get the pipeline state pub fn pipeline(&self) -> &ProtocolObject { &self.pipeline } /// Calculate optimal threadgroup size for 1D dispatch pub fn optimal_threadgroup_size_1d(&self, total_elements: usize) -> usize { let max = self.max_threads_per_threadgroup; let width = self.thread_execution_width; // Prefer multiples of thread execution width if total_elements >= max { max } else if total_elements >= width { (total_elements / width) * width } else { total_elements } } } #[cfg(not(target_os = "macos"))] impl MetalKernel { /// Get kernel name pub fn name(&self) -> &str { &self.name } /// Get maximum threads per threadgroup pub fn max_threads_per_threadgroup(&self) -> usize { 0 } /// Get thread execution width pub fn thread_execution_width(&self) -> usize { 0 } } /// Metal compute pipeline manager /// /// Handles shader compilation, kernel caching, and execution #[cfg(target_os = "macos")] pub struct MetalComputePipeline { device: Arc, /// Cached compiled kernels kernel_cache: DashMap>, /// Compiled shader library library: Option>>, } #[cfg(not(target_os = "macos"))] pub struct MetalComputePipeline { _private: (), } #[cfg(target_os = "macos")] impl MetalComputePipeline { /// Create a new compute pipeline manager pub fn new(device: Arc) -> Self { Self { device, kernel_cache: DashMap::new(), library: None, } } /// Compile MSL shader source code into a library pub fn compile_library(&mut self, source: &str) -> Result<()> { debug!("Compiling Metal shader library"); let ns_source = NSString::from_str(source); let library = self.device.mtl_device() .newLibraryWithSource_options_error(&ns_source, None) .map_err(|e| MetalError::ShaderCompilation(format!("{e:?}")))?; self.library = Some(library); info!("Metal shader library compiled successfully"); Ok(()) } /// Get or compile a kernel by name pub fn get_kernel(&self, name: &str) -> Result> { // Check cache first if let Some(kernel) = self.kernel_cache.get(name) { return Ok(kernel.value().clone()); } // Need library to compile kernel let library = self.library.as_ref() .ok_or_else(|| MetalError::ShaderCompilation("No shader library compiled".into()))?; // Get function from library let ns_name = NSString::from_str(name); let function = library.newFunctionWithName(&ns_name) .ok_or_else(|| MetalError::KernelNotFound(name.to_string()))?; // Create pipeline state let pipeline = self.device.mtl_device() .newComputePipelineStateWithFunction_error(&function) .map_err(|e| MetalError::PipelineCreation(format!("{e:?}")))?; let kernel = Arc::new(MetalKernel { name: name.to_string(), max_threads_per_threadgroup: pipeline.maxTotalThreadsPerThreadgroup(), thread_execution_width: pipeline.threadExecutionWidth(), pipeline, }); self.kernel_cache.insert(name.to_string(), kernel.clone()); debug!("Cached Metal kernel: {}", name); Ok(kernel) } /// Execute a kernel with the given buffers pub fn execute( &self, kernel: &MetalKernel, config: MetalLaunchConfig, buffers: &[&MetalBuffer], ) -> Result<()> { config.validate()?; // Create command buffer let cmd_buffer = self.device.command_queue_retained().commandBuffer() .ok_or_else(|| MetalError::CommandBuffer("Failed to create command buffer".into()))?; // Create compute encoder let encoder = cmd_buffer.computeCommandEncoder() .ok_or_else(|| MetalError::Execution("Failed to create compute encoder".into()))?; unsafe { // Set pipeline state encoder.setComputePipelineState(kernel.pipeline()); // Bind buffers for (index, buffer) in buffers.iter().enumerate() { encoder.setBuffer_offset_atIndex( Some(buffer.mtl_buffer()), 0, index ); } // Dispatch threads let grid_size = MTLSize { width: config.grid_size.0, height: config.grid_size.1, depth: config.grid_size.2, }; let threadgroup_size = MTLSize { width: config.threadgroup_size.0, height: config.threadgroup_size.1, depth: config.threadgroup_size.2, }; encoder.dispatchThreads_threadsPerThreadgroup(grid_size, threadgroup_size); } // End encoding and execute encoder.endEncoding(); cmd_buffer.commit(); cmd_buffer.waitUntilCompleted(); Ok(()) } /// Execute a kernel with data buffers and scalar constant buffers /// /// This is useful for kernels that need both data buffers (A, B, C matrices) /// and scalar parameters (M, N, K dimensions) passed as separate buffers. /// /// Buffers are bound at indices 0..buffers.len(), /// scalars are bound at indices buffers.len().. pub fn execute_with_scalars( &self, kernel: &MetalKernel, config: MetalLaunchConfig, buffers: &[&MetalBuffer], scalars: &[&dyn crate::memory::AsRawMetalBuffer], ) -> Result<()> { config.validate()?; // Create command buffer let cmd_buffer = self.device.command_queue_retained().commandBuffer() .ok_or_else(|| MetalError::CommandBuffer("Failed to create command buffer".into()))?; // Create compute encoder let encoder = cmd_buffer.computeCommandEncoder() .ok_or_else(|| MetalError::Execution("Failed to create compute encoder".into()))?; unsafe { // Set pipeline state encoder.setComputePipelineState(kernel.pipeline()); // Bind data buffers at indices 0..buffers.len() for (index, buffer) in buffers.iter().enumerate() { encoder.setBuffer_offset_atIndex( Some(buffer.mtl_buffer()), 0, index ); } // Bind scalar buffers at indices buffers.len().. let scalar_offset = buffers.len(); for (index, scalar) in scalars.iter().enumerate() { encoder.setBuffer_offset_atIndex( Some(scalar.as_raw_buffer()), 0, scalar_offset + index ); } // Dispatch threads let grid_size = MTLSize { width: config.grid_size.0, height: config.grid_size.1, depth: config.grid_size.2, }; let threadgroup_size = MTLSize { width: config.threadgroup_size.0, height: config.threadgroup_size.1, depth: config.threadgroup_size.2, }; encoder.dispatchThreads_threadsPerThreadgroup(grid_size, threadgroup_size); } // End encoding and execute encoder.endEncoding(); cmd_buffer.commit(); cmd_buffer.waitUntilCompleted(); Ok(()) } /// Clear kernel cache pub fn clear_cache(&self) { self.kernel_cache.clear(); debug!("Metal kernel cache cleared"); } /// Get number of cached kernels pub fn cache_size(&self) -> usize { self.kernel_cache.len() } /// Dispatch a kernel with a closure for setting up encoder pub fn dispatch( &self, kernel_name: &str, grid_size: [u32; 3], threadgroup_size: [u32; 3], setup: F, ) -> Result<()> where F: FnOnce(&ComputeEncoder), { let kernel = self.get_kernel(kernel_name)?; let cmd_buffer = self .device .command_queue() .commandBuffer() .ok_or_else(|| MetalError::CommandBuffer("Failed to create command buffer".into()))?; let encoder = cmd_buffer .computeCommandEncoder() .ok_or_else(|| MetalError::Execution("Failed to create compute encoder".into()))?; unsafe { encoder.setComputePipelineState(kernel.pipeline()); } // Wrap encoder and call setup let wrapper = ComputeEncoder { encoder: &encoder }; setup(&wrapper); unsafe { let grid = MTLSize { width: grid_size[0] as usize, height: grid_size[1] as usize, depth: grid_size[2] as usize, }; let threads = MTLSize { width: threadgroup_size[0] as usize, height: threadgroup_size[1] as usize, depth: threadgroup_size[2] as usize, }; encoder.dispatchThreads_threadsPerThreadgroup(grid, threads); } encoder.endEncoding(); cmd_buffer.commit(); cmd_buffer.waitUntilCompleted(); Ok(()) } /// Dispatch a 1D kernel with automatic threadgroup sizing pub fn dispatch_1d(&self, kernel_name: &str, total_threads: u32, setup: F) -> Result<()> where F: FnOnce(&ComputeEncoder), { let threads_per_group = 256.min(total_threads); self.dispatch( kernel_name, [total_threads, 1, 1], [threads_per_group, 1, 1], setup, ) } /// Dispatch a 3D kernel with automatic threadgroup sizing pub fn dispatch_3d(&self, kernel_name: &str, grid_size: [u32; 3], setup: F) -> Result<()> where F: FnOnce(&ComputeEncoder), { let tx = 8.min(grid_size[0]); let ty = 8.min(grid_size[1]); let tz = 1.min(grid_size[2]); self.dispatch(kernel_name, grid_size, [tx, ty, tz], setup) } } /// Wrapper for compute encoder that provides safe buffer/bytes setting #[cfg(target_os = "macos")] pub struct ComputeEncoder<'a> { encoder: &'a ProtocolObject, } #[cfg(target_os = "macos")] impl<'a> ComputeEncoder<'a> { /// Set a buffer at the given index pub fn set_buffer(&self, index: u64, buffer: &MetalBuffer) { unsafe { self.encoder .setBuffer_offset_atIndex(Some(buffer.mtl_buffer()), 0, index as usize); } } /// Set bytes (scalar value) at the given index pub fn set_bytes(&self, index: u64, value: &T) { unsafe { self.encoder.setBytes_length_atIndex( std::ptr::NonNull::new(value as *const T as *mut std::ffi::c_void).unwrap(), std::mem::size_of::(), index as usize, ); } } } #[cfg(not(target_os = "macos"))] impl MetalComputePipeline { /// Create a new compute pipeline manager (stub) pub fn new(_device: Arc) -> Self { Self { _private: () } } /// Compile shader source (stub) pub fn compile_library(&mut self, _source: &str) -> Result<()> { Err(MetalError::NotAvailable) } /// Get kernel (stub) pub fn get_kernel(&self, _name: &str) -> Result> { Err(MetalError::NotAvailable) } /// Execute kernel (stub) pub fn execute( &self, _kernel: &MetalKernel, _config: MetalLaunchConfig, _buffers: &[&MetalBuffer], ) -> Result<()> { Err(MetalError::NotAvailable) } /// Execute kernel with scalars (stub) pub fn execute_with_scalars( &self, _kernel: &MetalKernel, _config: MetalLaunchConfig, _buffers: &[&MetalBuffer], _scalars: &[&dyn crate::memory::AsRawMetalBuffer], ) -> Result<()> { Err(MetalError::NotAvailable) } /// Clear cache pub fn clear_cache(&self) {} /// Get cache size pub fn cache_size(&self) -> usize { 0 } /// Dispatch a kernel (stub) pub fn dispatch( &self, _kernel_name: &str, _grid_size: [u32; 3], _threadgroup_size: [u32; 3], _setup: F, ) -> Result<()> where F: FnOnce(&ComputeEncoder), { Err(MetalError::NotAvailable) } /// Dispatch a 1D kernel (stub) pub fn dispatch_1d(&self, _kernel_name: &str, _total_threads: u32, _setup: F) -> Result<()> where F: FnOnce(&ComputeEncoder), { Err(MetalError::NotAvailable) } /// Dispatch a 3D kernel (stub) pub fn dispatch_3d(&self, _kernel_name: &str, _grid_size: [u32; 3], _setup: F) -> Result<()> where F: FnOnce(&ComputeEncoder), { Err(MetalError::NotAvailable) } } /// Wrapper for compute encoder (stub) #[cfg(not(target_os = "macos"))] pub struct ComputeEncoder<'a> { _phantom: std::marker::PhantomData<&'a ()>, } #[cfg(not(target_os = "macos"))] impl<'a> ComputeEncoder<'a> { /// Set a buffer at the given index (stub) pub fn set_buffer(&self, _index: u64, _buffer: &MetalBuffer) {} /// Set bytes at the given index (stub) pub fn set_bytes(&self, _index: u64, _value: &T) {} } #[cfg(test)] mod tests { use super::*; #[test] fn test_launch_config_validation() { let valid = MetalLaunchConfig::new_1d(1000, 256); assert!(valid.validate().is_ok()); let invalid = MetalLaunchConfig { grid_size: (0, 1, 1), threadgroup_size: (256, 1, 1), }; assert!(invalid.validate().is_err()); let too_many_threads = MetalLaunchConfig { grid_size: (1000, 1, 1), threadgroup_size: (1025, 1, 1), }; assert!(too_many_threads.validate().is_err()); } }