Files
rustytorch/crates/core/rtx-backend-metal/src/lib.rs
T
quantumandClaude Fable 5 9297976929 rtx-backend-metal: GPU index_select / index_add via one-hot CSR SpMM
Override the Backend trait's host-round-trip defaults: gather is S @ X with
S the [E x N] one-hot selection CSR; scatter-add is the adjoint S^T @ X,
whose CSR is built directly by counting sort so duplicate indices land in
one row and the spmm kernel (one thread per output element) accumulates
them without atomics. CSR matrices are cached per thread keyed by the
exact index list + dims, so a static graph topology (GNN message passing)
builds each matrix once. Host fallback on degenerate shapes or any
sparse-pipeline failure.

13 new parity tests vs CPU reference: duplicates, unreferenced rows,
D=1/2/3, 15k x 5k x 64 gather/scatter, cache reuse, adjoint roundtrip.
Verified on-device that the SpMM path (not the fallback) serves all 13.

Co-Authored-By: Claude Fable 5 <[email protected]>
2026-08-21 05:56:23 -07:00

439 lines
13 KiB
Rust

//! # 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<const D: usize> = MetalTensorPrimitive<D>;
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<const D: usize>(shape: [usize; D], device: &Self::Device) -> Self::TensorPrimitive<D> {
ops::creation::zeros(shape, device)
}
fn ones<const D: usize>(shape: [usize; D], device: &Self::Device) -> Self::TensorPrimitive<D> {
ops::creation::ones(shape, device)
}
fn full<const D: usize>(
shape: [usize; D],
fill_value: Self::FloatElem,
device: &Self::Device,
) -> Self::TensorPrimitive<D> {
ops::creation::full(shape, fill_value, device)
}
fn rand<const D: usize>(shape: [usize; D], device: &Self::Device) -> Self::TensorPrimitive<D> {
ops::creation::rand(shape, device)
}
fn randn<const D: usize>(shape: [usize; D], device: &Self::Device) -> Self::TensorPrimitive<D> {
ops::creation::randn(shape, device)
}
fn from_data<const D: usize>(
data: &[Self::FloatElem],
shape: [usize; D],
device: &Self::Device,
) -> Self::TensorPrimitive<D> {
ops::creation::from_data(data, shape, device)
}
// ==================== Basic Operations ====================
fn add<const D: usize>(
lhs: Self::TensorPrimitive<D>,
rhs: Self::TensorPrimitive<D>,
) -> Self::TensorPrimitive<D> {
ops::basic::add(&lhs, &rhs)
}
fn sub<const D: usize>(
lhs: Self::TensorPrimitive<D>,
rhs: Self::TensorPrimitive<D>,
) -> Self::TensorPrimitive<D> {
ops::basic::sub(&lhs, &rhs)
}
fn mul<const D: usize>(
lhs: Self::TensorPrimitive<D>,
rhs: Self::TensorPrimitive<D>,
) -> Self::TensorPrimitive<D> {
ops::basic::mul(&lhs, &rhs)
}
fn div<const D: usize>(
lhs: Self::TensorPrimitive<D>,
rhs: Self::TensorPrimitive<D>,
) -> Self::TensorPrimitive<D> {
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<const D: usize>(tensor: Self::TensorPrimitive<D>) -> Self::TensorPrimitive<D> {
ops::unary::neg(&tensor)
}
fn exp<const D: usize>(tensor: Self::TensorPrimitive<D>) -> Self::TensorPrimitive<D> {
ops::unary::exp(&tensor)
}
fn log<const D: usize>(tensor: Self::TensorPrimitive<D>) -> Self::TensorPrimitive<D> {
ops::unary::log(&tensor)
}
fn sqrt<const D: usize>(tensor: Self::TensorPrimitive<D>) -> Self::TensorPrimitive<D> {
ops::unary::sqrt(&tensor)
}
fn abs<const D: usize>(tensor: Self::TensorPrimitive<D>) -> Self::TensorPrimitive<D> {
ops::unary::abs(&tensor)
}
fn sin<const D: usize>(tensor: Self::TensorPrimitive<D>) -> Self::TensorPrimitive<D> {
ops::unary::sin(&tensor)
}
fn cos<const D: usize>(tensor: Self::TensorPrimitive<D>) -> Self::TensorPrimitive<D> {
ops::unary::cos(&tensor)
}
fn pow<const D: usize>(
tensor: Self::TensorPrimitive<D>,
exp: Self::FloatElem,
) -> Self::TensorPrimitive<D> {
ops::unary::pow(&tensor, exp)
}
fn clamp<const D: usize>(
tensor: Self::TensorPrimitive<D>,
min: Self::FloatElem,
max: Self::FloatElem,
) -> Self::TensorPrimitive<D> {
ops::unary::clamp(&tensor, min, max)
}
// ==================== Activation Functions ====================
fn relu<const D: usize>(tensor: Self::TensorPrimitive<D>) -> Self::TensorPrimitive<D> {
ops::activation::relu(&tensor)
}
fn sigmoid<const D: usize>(tensor: Self::TensorPrimitive<D>) -> Self::TensorPrimitive<D> {
ops::activation::sigmoid(&tensor)
}
fn tanh<const D: usize>(tensor: Self::TensorPrimitive<D>) -> Self::TensorPrimitive<D> {
ops::activation::tanh(&tensor)
}
// ==================== Reduction Operations ====================
fn sum<const D: usize>(tensor: Self::TensorPrimitive<D>) -> Self::TensorPrimitive<1> {
ops::reduction::sum(&tensor)
}
fn sum_dim<const D: usize>(
tensor: Self::TensorPrimitive<D>,
dim: usize,
) -> Self::TensorPrimitive<D> {
ops::reduction::sum_dim(&tensor, dim)
}
fn mean<const D: usize>(tensor: Self::TensorPrimitive<D>) -> Self::TensorPrimitive<1> {
ops::reduction::mean(&tensor)
}
fn mean_dim<const D: usize>(
tensor: Self::TensorPrimitive<D>,
dim: usize,
) -> Self::TensorPrimitive<D> {
ops::reduction::mean_dim(&tensor, dim)
}
fn max<const D: usize>(tensor: Self::TensorPrimitive<D>) -> Self::TensorPrimitive<1> {
ops::reduction::max(&tensor)
}
fn min<const D: usize>(tensor: Self::TensorPrimitive<D>) -> Self::TensorPrimitive<1> {
ops::reduction::min(&tensor)
}
fn var<const D: usize>(tensor: Self::TensorPrimitive<D>) -> Self::TensorPrimitive<1> {
ops::reduction::var(&tensor)
}
fn var_dim<const D: usize>(
tensor: Self::TensorPrimitive<D>,
dim: usize,
) -> Self::TensorPrimitive<D> {
ops::reduction::var_dim(&tensor, dim)
}
// ==================== Shape Operations ====================
fn shape<const D: usize>(tensor: &Self::TensorPrimitive<D>) -> [usize; D] {
tensor.shape
}
fn reshape<const D1: usize, const D2: usize>(
tensor: Self::TensorPrimitive<D1>,
shape: [usize; D2],
) -> Self::TensorPrimitive<D2> {
ops::shape::reshape(tensor, shape)
}
fn transpose<const D: usize>(tensor: Self::TensorPrimitive<D>) -> Self::TensorPrimitive<D> {
ops::shape::transpose(&tensor)
}
fn swap_dims<const D: usize>(
tensor: Self::TensorPrimitive<D>,
dim1: usize,
dim2: usize,
) -> Self::TensorPrimitive<D> {
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<const D: usize>(
tensor: Self::TensorPrimitive<D>,
dim: usize,
) -> Self::TensorPrimitive<D> {
ops::activation::softmax(&tensor, dim)
}
fn layer_norm<const D: usize>(
tensor: Self::TensorPrimitive<D>,
weight: &Self::TensorPrimitive<1>,
bias: Option<&Self::TensorPrimitive<1>>,
eps: Self::FloatElem,
) -> Self::TensorPrimitive<D> {
ops::normalization::layer_norm(&tensor, weight, bias, eps)
}
fn rms_norm<const D: usize>(
tensor: Self::TensorPrimitive<D>,
weight: &Self::TensorPrimitive<1>,
eps: Self::FloatElem,
) -> Self::TensorPrimitive<D> {
ops::normalization::rms_norm(&tensor, weight, eps)
}
fn rope<const D: usize>(
tensor: Self::TensorPrimitive<D>,
cos: &Self::TensorPrimitive<2>,
sin: &Self::TensorPrimitive<2>,
) -> Self::TensorPrimitive<D> {
ops::attention::rope(&tensor, cos, sin)
}
fn gelu<const D: usize>(tensor: Self::TensorPrimitive<D>) -> Self::TensorPrimitive<D> {
ops::activation::gelu(&tensor)
}
fn silu<const D: usize>(tensor: Self::TensorPrimitive<D>) -> Self::TensorPrimitive<D> {
ops::activation::silu(&tensor)
}
fn leaky_relu<const D: usize>(
tensor: Self::TensorPrimitive<D>,
negative_slope: Self::FloatElem,
) -> Self::TensorPrimitive<D> {
ops::activation::leaky_relu(&tensor, negative_slope)
}
fn elu<const D: usize>(
tensor: Self::TensorPrimitive<D>,
alpha: Self::FloatElem,
) -> Self::TensorPrimitive<D> {
ops::activation::elu(&tensor, alpha)
}
// ==================== Comparison Operations ====================
fn gt_scalar<const D: usize>(
tensor: Self::TensorPrimitive<D>,
value: Self::FloatElem,
) -> Self::TensorPrimitive<D> {
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<const D: usize>(
tensor: Self::TensorPrimitive<D>,
indices: &[usize],
) -> Self::TensorPrimitive<D> {
ops::index::index_select(&tensor, indices)
}
fn index_add<const D: usize>(
tensor: Self::TensorPrimitive<D>,
indices: &[usize],
num_rows: usize,
) -> Self::TensorPrimitive<D> {
ops::index::index_add(&tensor, indices, num_rows)
}
// ==================== Device Management ====================
fn device<const D: usize>(tensor: &Self::TensorPrimitive<D>) -> Self::Device {
tensor.device.clone()
}
fn to_device<const D: usize>(
tensor: Self::TensorPrimitive<D>,
device: &Self::Device,
) -> Self::TensorPrimitive<D> {
if tensor.device == *device {
tensor
} else {
ops::device::copy_to_device(tensor, device)
}
}
fn to_data<const D: usize>(tensor: &Self::TensorPrimitive<D>) -> Vec<Self::FloatElem> {
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;