#[cfg(all(target_os = "macos", feature = "metal-gpu"))]
mod gpu {
use crate::forward::cpu::{validate_gemm_bt, validate_gemm_nn};
use metal::*;
use std::sync::OnceLock;
const GPU_DISPATCH_THRESHOLD: u64 = 64 * 64 * 64;
struct MetalState {
device: Device,
queue: CommandQueue,
pipeline_bt: ComputePipelineState,
pipeline_nn: ComputePipelineState,
}
static METAL: OnceLock<Option<MetalState>> = OnceLock::new();
const SHADER_SOURCE: &str = include_str!("shaders/gemm.metal");
fn init_metal() -> Option<MetalState> {
let device = Device::system_default()?;
tracing::info!(
name = device.name(),
"Metal GPU initialized for GEMM acceleration"
);
let queue = device.new_command_queue();
let opts = CompileOptions::new();
let library = device.new_library_with_source(SHADER_SOURCE, &opts).ok()?;
let fn_bt = library.get_function("gemm_bt", None).ok()?;
let fn_nn = library.get_function("gemm_nn", None).ok()?;
let pipeline_bt = device
.new_compute_pipeline_state_with_function(&fn_bt)
.ok()?;
let pipeline_nn = device
.new_compute_pipeline_state_with_function(&fn_nn)
.ok()?;
Some(MetalState {
device,
queue,
pipeline_bt,
pipeline_nn,
})
}
fn get_metal() -> Option<&'static MetalState> {
METAL.get_or_init(init_metal).as_ref()
}
fn run_gemm(
state: &MetalState,
pipeline: &ComputePipelineState,
a: &[f32],
b: &[f32],
c: &mut [f32],
m: u32,
n: u32,
k: u32,
) {
let a_bytes = std::mem::size_of_val(a) as u64;
let b_bytes = std::mem::size_of_val(b) as u64;
let c_bytes = (m as u64) * (n as u64) * 4;
let buf_a = state.device.new_buffer_with_bytes_no_copy(
a.as_ptr() as *const _,
a_bytes,
MTLResourceOptions::StorageModeShared,
None,
);
let buf_b = state.device.new_buffer_with_bytes_no_copy(
b.as_ptr() as *const _,
b_bytes,
MTLResourceOptions::StorageModeShared,
None,
);
let buf_c = state.device.new_buffer_with_bytes_no_copy(
c.as_mut_ptr() as *mut _ as *const _,
c_bytes,
MTLResourceOptions::StorageModeShared,
None,
);
let cmd = state.queue.new_command_buffer();
let enc = cmd.new_compute_command_encoder();
enc.set_compute_pipeline_state(pipeline);
enc.set_buffer(0, Some(&buf_a), 0);
enc.set_buffer(1, Some(&buf_b), 0);
enc.set_buffer(2, Some(&buf_c), 0);
enc.set_bytes(3, 4, &m as *const u32 as *const _);
enc.set_bytes(4, 4, &n as *const u32 as *const _);
enc.set_bytes(5, 4, &k as *const u32 as *const _);
let tile = 16u64;
let grid = MTLSize::new(
(n as u64).div_ceil(tile) * tile,
(m as u64).div_ceil(tile) * tile,
1,
);
let tg = MTLSize::new(tile, tile, 1);
enc.dispatch_threads(grid, tg);
enc.end_encoding();
cmd.commit();
cmd.wait_until_completed();
}
pub fn metal_matmul_bt(
a: &[f32],
b: &[f32],
c: &mut [f32],
m: usize,
k: usize,
n: usize,
) -> bool {
validate_gemm_bt(a.len(), b.len(), c.len(), m, k, n, "metal_matmul_bt");
let work = (m as u64) * (n as u64) * (k as u64);
if work < GPU_DISPATCH_THRESHOLD {
return false;
}
let Some(state) = get_metal() else {
return false;
};
run_gemm(
state,
&state.pipeline_bt,
a,
b,
c,
m as u32,
n as u32,
k as u32,
);
true
}
#[allow(dead_code)] pub fn is_available() -> bool {
get_metal().is_some()
}
pub fn metal_matmul(a: &[f32], b: &[f32], c: &mut [f32], m: usize, k: usize, n: usize) -> bool {
validate_gemm_nn(a.len(), b.len(), c.len(), m, k, n, "metal_matmul");
let work = (m as u64) * (n as u64) * (k as u64);
if work < GPU_DISPATCH_THRESHOLD {
return false;
}
let Some(state) = get_metal() else {
return false;
};
run_gemm(
state,
&state.pipeline_nn,
a,
b,
c,
m as u32,
n as u32,
k as u32,
);
true
}
}
#[cfg(all(target_os = "macos", feature = "metal-gpu"))]
pub use gpu::{is_available, metal_matmul, metal_matmul_bt};
#[cfg(not(all(target_os = "macos", feature = "metal-gpu")))]
pub fn metal_matmul_bt(
_a: &[f32],
_b: &[f32],
_c: &mut [f32],
_m: usize,
_k: usize,
_n: usize,
) -> bool {
false
}
#[cfg(not(all(target_os = "macos", feature = "metal-gpu")))]
pub fn metal_matmul(
_a: &[f32],
_b: &[f32],
_c: &mut [f32],
_m: usize,
_k: usize,
_n: usize,
) -> bool {
false
}
#[cfg(all(test, target_os = "macos", feature = "metal-gpu"))]
mod tests {
use super::{metal_matmul, metal_matmul_bt};
#[test]
#[should_panic(expected = "b too short for n*k")]
fn metal_matmul_bt_rejects_short_b() {
let a = [0.0f32; 2];
let b = [0.0f32; 1]; let mut c = [0.0f32; 1];
metal_matmul_bt(&a, &b, &mut c, 1, 1, 2);
}
#[test]
#[should_panic(expected = "shape overflow: n*k")]
fn metal_matmul_bt_rejects_overflow() {
let a = [0.0f32; 2];
let b = [0.0f32; 2];
let mut c = [0.0f32; 2];
metal_matmul_bt(&a, &b, &mut c, 2, 2, usize::MAX);
}
#[test]
#[should_panic(expected = "b too short for k*n")]
fn metal_matmul_rejects_short_b() {
let a = [0.0f32; 2];
let b = [0.0f32; 1]; let mut c = [0.0f32; 1];
metal_matmul(&a, &b, &mut c, 1, 1, 2);
}
}