//! # RustyTorch++ Metal Backend //! //! Apple Silicon GPU backend implementation using Metal and MPS. //! //! ## Features //! //! - **Unified Memory**: Zero-copy CPU-GPU data sharing on Apple Silicon //! - **MPS Integration**: Metal Performance Shaders for optimized BLAS //! - **Flash Attention**: Hand-optimized attention kernels for M-series //! - **Low Power**: Efficient for laptop/mobile deployment //! //! ## Architecture //! //! ```text //! MetalBackend //! ├── MetalTensorPrimitive - Unified memory buffer //! ├── MetalDevice - Device context and command queues //! └── Ops //! ├── Basic - Add, mul, etc. //! ├── GEMM - Matrix multiply (MPS) //! └── Attention - Flash Attention (MSL kernels) //! ``` //! //! ## Example //! //! ```rust,ignore //! use rtx_backend_metal::{MetalBackend, MetalDevice}; //! use rtx_backend::Backend; //! //! let device = MetalDevice::default()?; //! let a = MetalBackend::zeros([1024, 1024], &device); //! let b = MetalBackend::randn([1024, 1024], &device); //! let c = MetalBackend::matmul(&a, &b); //! ``` #![warn(missing_docs)] mod device; mod error; /// Operations module - re-exported for tests and direct access. pub mod ops; mod tensor; pub use device::MetalDeviceWrapper; pub use error::{MetalBackendError, MetalBackendResult}; pub use tensor::MetalTensorPrimitive; use rtx_backend::{Backend, BoolU8}; /// Metal backend for RustyTorch++. /// /// This backend uses Apple Silicon GPUs via Metal and provides: /// - MPS for matrix operations /// - Hand-optimized MSL shaders /// - Unified memory for zero-copy transfers /// - Efficient power consumption #[derive(Clone, Debug, Default)] pub struct MetalBackend; impl Backend for MetalBackend { type TensorPrimitive = MetalTensorPrimitive; type Device = MetalDeviceWrapper; type FloatElem = f32; type IntElem = i32; type BoolElem = BoolU8; fn name() -> &'static str { "metal" } fn seed(seed: u64) { ops::seed_rng(seed); } // ==================== Tensor Creation ==================== fn zeros(shape: [usize; D], device: &Self::Device) -> Self::TensorPrimitive { ops::creation::zeros(shape, device) } fn ones(shape: [usize; D], device: &Self::Device) -> Self::TensorPrimitive { ops::creation::ones(shape, device) } fn full( shape: [usize; D], fill_value: Self::FloatElem, device: &Self::Device, ) -> Self::TensorPrimitive { ops::creation::full(shape, fill_value, device) } fn rand(shape: [usize; D], device: &Self::Device) -> Self::TensorPrimitive { ops::creation::rand(shape, device) } fn randn(shape: [usize; D], device: &Self::Device) -> Self::TensorPrimitive { ops::creation::randn(shape, device) } fn from_data( data: &[Self::FloatElem], shape: [usize; D], device: &Self::Device, ) -> Self::TensorPrimitive { ops::creation::from_data(data, shape, device) } // ==================== Basic Operations ==================== fn add( lhs: Self::TensorPrimitive, rhs: Self::TensorPrimitive, ) -> Self::TensorPrimitive { ops::basic::add(&lhs, &rhs) } fn sub( lhs: Self::TensorPrimitive, rhs: Self::TensorPrimitive, ) -> Self::TensorPrimitive { ops::basic::sub(&lhs, &rhs) } fn mul( lhs: Self::TensorPrimitive, rhs: Self::TensorPrimitive, ) -> Self::TensorPrimitive { ops::basic::mul(&lhs, &rhs) } fn div( lhs: Self::TensorPrimitive, rhs: Self::TensorPrimitive, ) -> Self::TensorPrimitive { ops::basic::div(&lhs, &rhs) } fn matmul( lhs: Self::TensorPrimitive<2>, rhs: Self::TensorPrimitive<2>, ) -> Self::TensorPrimitive<2> { ops::gemm::matmul(&lhs, &rhs) } fn bmm( lhs: Self::TensorPrimitive<3>, rhs: Self::TensorPrimitive<3>, ) -> Self::TensorPrimitive<3> { ops::gemm::bmm(&lhs, &rhs) } // ==================== Unary Operations ==================== fn neg(tensor: Self::TensorPrimitive) -> Self::TensorPrimitive { ops::unary::neg(&tensor) } fn exp(tensor: Self::TensorPrimitive) -> Self::TensorPrimitive { ops::unary::exp(&tensor) } fn log(tensor: Self::TensorPrimitive) -> Self::TensorPrimitive { ops::unary::log(&tensor) } fn sqrt(tensor: Self::TensorPrimitive) -> Self::TensorPrimitive { ops::unary::sqrt(&tensor) } fn abs(tensor: Self::TensorPrimitive) -> Self::TensorPrimitive { ops::unary::abs(&tensor) } fn sin(tensor: Self::TensorPrimitive) -> Self::TensorPrimitive { ops::unary::sin(&tensor) } fn cos(tensor: Self::TensorPrimitive) -> Self::TensorPrimitive { ops::unary::cos(&tensor) } fn pow( tensor: Self::TensorPrimitive, exp: Self::FloatElem, ) -> Self::TensorPrimitive { ops::unary::pow(&tensor, exp) } fn clamp( tensor: Self::TensorPrimitive, min: Self::FloatElem, max: Self::FloatElem, ) -> Self::TensorPrimitive { ops::unary::clamp(&tensor, min, max) } // ==================== Activation Functions ==================== fn relu(tensor: Self::TensorPrimitive) -> Self::TensorPrimitive { ops::activation::relu(&tensor) } fn sigmoid(tensor: Self::TensorPrimitive) -> Self::TensorPrimitive { ops::activation::sigmoid(&tensor) } fn tanh(tensor: Self::TensorPrimitive) -> Self::TensorPrimitive { ops::activation::tanh(&tensor) } // ==================== Reduction Operations ==================== fn sum(tensor: Self::TensorPrimitive) -> Self::TensorPrimitive<1> { ops::reduction::sum(&tensor) } fn sum_dim( tensor: Self::TensorPrimitive, dim: usize, ) -> Self::TensorPrimitive { ops::reduction::sum_dim(&tensor, dim) } fn mean(tensor: Self::TensorPrimitive) -> Self::TensorPrimitive<1> { ops::reduction::mean(&tensor) } fn mean_dim( tensor: Self::TensorPrimitive, dim: usize, ) -> Self::TensorPrimitive { ops::reduction::mean_dim(&tensor, dim) } fn max(tensor: Self::TensorPrimitive) -> Self::TensorPrimitive<1> { ops::reduction::max(&tensor) } fn min(tensor: Self::TensorPrimitive) -> Self::TensorPrimitive<1> { ops::reduction::min(&tensor) } fn var(tensor: Self::TensorPrimitive) -> Self::TensorPrimitive<1> { ops::reduction::var(&tensor) } fn var_dim( tensor: Self::TensorPrimitive, dim: usize, ) -> Self::TensorPrimitive { ops::reduction::var_dim(&tensor, dim) } // ==================== Shape Operations ==================== fn shape(tensor: &Self::TensorPrimitive) -> [usize; D] { tensor.shape } fn reshape( tensor: Self::TensorPrimitive, shape: [usize; D2], ) -> Self::TensorPrimitive { ops::shape::reshape(tensor, shape) } fn transpose(tensor: Self::TensorPrimitive) -> Self::TensorPrimitive { ops::shape::transpose(&tensor) } fn swap_dims( tensor: Self::TensorPrimitive, dim1: usize, dim2: usize, ) -> Self::TensorPrimitive { ops::shape::swap_dims(&tensor, dim1, dim2) } // ==================== LLM-Specific Operations ==================== fn flash_attention( query: Self::TensorPrimitive<4>, key: Self::TensorPrimitive<4>, value: Self::TensorPrimitive<4>, mask: Option<&Self::TensorPrimitive<4>>, scale: Self::FloatElem, causal: bool, ) -> Self::TensorPrimitive<4> { ops::attention::flash_attention(&query, &key, &value, mask, scale, causal) } fn softmax( tensor: Self::TensorPrimitive, dim: usize, ) -> Self::TensorPrimitive { ops::activation::softmax(&tensor, dim) } fn layer_norm( tensor: Self::TensorPrimitive, weight: &Self::TensorPrimitive<1>, bias: Option<&Self::TensorPrimitive<1>>, eps: Self::FloatElem, ) -> Self::TensorPrimitive { ops::normalization::layer_norm(&tensor, weight, bias, eps) } fn rms_norm( tensor: Self::TensorPrimitive, weight: &Self::TensorPrimitive<1>, eps: Self::FloatElem, ) -> Self::TensorPrimitive { ops::normalization::rms_norm(&tensor, weight, eps) } fn rope( tensor: Self::TensorPrimitive, cos: &Self::TensorPrimitive<2>, sin: &Self::TensorPrimitive<2>, ) -> Self::TensorPrimitive { ops::attention::rope(&tensor, cos, sin) } fn gelu(tensor: Self::TensorPrimitive) -> Self::TensorPrimitive { ops::activation::gelu(&tensor) } fn silu(tensor: Self::TensorPrimitive) -> Self::TensorPrimitive { ops::activation::silu(&tensor) } fn leaky_relu( tensor: Self::TensorPrimitive, negative_slope: Self::FloatElem, ) -> Self::TensorPrimitive { ops::activation::leaky_relu(&tensor, negative_slope) } fn elu( tensor: Self::TensorPrimitive, alpha: Self::FloatElem, ) -> Self::TensorPrimitive { ops::activation::elu(&tensor, alpha) } // ==================== Comparison Operations ==================== fn gt_scalar( tensor: Self::TensorPrimitive, value: Self::FloatElem, ) -> Self::TensorPrimitive { ops::unary::gt_scalar(&tensor, value) } // ==================== Convolution Operations ==================== fn conv2d( input: Self::TensorPrimitive<4>, weight: &Self::TensorPrimitive<4>, bias: Option<&Self::TensorPrimitive<1>>, stride: [usize; 2], padding: [usize; 2], dilation: [usize; 2], groups: usize, ) -> Self::TensorPrimitive<4> { ops::conv::conv2d(&input, weight, bias, stride, padding, dilation, groups) } // ==================== Pooling Operations ==================== fn max_pool2d( input: Self::TensorPrimitive<4>, kernel_size: [usize; 2], stride: [usize; 2], padding: [usize; 2], ) -> Self::TensorPrimitive<4> { ops::conv::max_pool2d(&input, kernel_size, stride, padding) } fn avg_pool2d( input: Self::TensorPrimitive<4>, kernel_size: [usize; 2], stride: [usize; 2], padding: [usize; 2], count_include_pad: bool, ) -> Self::TensorPrimitive<4> { ops::conv::avg_pool2d(&input, kernel_size, stride, padding, count_include_pad) } // ==================== Index Operations ==================== fn index_select( tensor: Self::TensorPrimitive, indices: &[usize], ) -> Self::TensorPrimitive { ops::index::index_select(&tensor, indices) } fn index_add( tensor: Self::TensorPrimitive, indices: &[usize], num_rows: usize, ) -> Self::TensorPrimitive { ops::index::index_add(&tensor, indices, num_rows) } // ==================== Device Management ==================== fn device(tensor: &Self::TensorPrimitive) -> Self::Device { tensor.device.clone() } fn to_device( tensor: Self::TensorPrimitive, device: &Self::Device, ) -> Self::TensorPrimitive { if tensor.device == *device { tensor } else { ops::device::copy_to_device(tensor, device) } } fn to_data(tensor: &Self::TensorPrimitive) -> Vec { ops::device::copy_to_host(tensor) } fn sync(device: &Self::Device) { device.synchronize(); } } /// Type alias for training with Metal + autodiff. pub type MetalTraining = MetalBackend; /// Type alias for inference with Metal (no autodiff overhead). pub type MetalInference = MetalBackend;