#[cfg(all(feature = "accelerate", target_os = "macos"))]
#[link(name = "Accelerate", kind = "framework")]
unsafe extern "C" {
fn cblas_sgemm(
order: i32,
trans_a: i32,
trans_b: i32,
m: i32,
n: i32,
k: i32,
alpha: f32,
a: *const f32,
lda: i32,
b: *const f32,
ldb: i32,
beta: f32,
c: *mut f32,
ldc: i32,
);
}
pub fn sgemm_forward(
y: &mut [f32],
x: &[f32],
w: &[f32],
bias: Option<&[f32]>,
batch: usize,
n_in: usize,
n_out: usize,
) {
if let Some(b) = bias {
for row in 0..batch {
let off = row * n_out;
y[off..off + n_out].copy_from_slice(&b[..n_out]);
}
} else {
y[..batch * n_out].fill(0.0);
}
#[cfg(all(feature = "accelerate", target_os = "macos"))]
unsafe {
cblas_sgemm(
101, 111, 111, batch as i32, n_out as i32, n_in as i32, 1.0, x.as_ptr(), n_in as i32, w.as_ptr(), n_out as i32, 1.0, y.as_mut_ptr(), n_out as i32, );
}
#[cfg(all(
feature = "gemm-blas",
not(all(feature = "accelerate", target_os = "macos"))
))]
unsafe {
gemm::gemm(
batch,
n_out,
n_in,
y.as_mut_ptr(),
1, n_out as isize, true, x.as_ptr(),
1, n_in as isize, w.as_ptr(),
1, n_out as isize, 1.0, 1.0, false,
false,
false,
gemm::Parallelism::None, );
}
#[cfg(not(any(
all(feature = "accelerate", target_os = "macos"),
feature = "gemm-blas"
)))]
{
for row in 0..batch {
let x_off = row * n_in;
let y_off = row * n_out;
for k in 0..n_in {
let xv = x[x_off + k];
let w_off = k * n_out;
for j in 0..n_out {
y[y_off + j] += xv * w[w_off + j];
}
}
}
}
}
pub fn matvec_forward(
y: &mut [f32],
x: &[f32],
w: &[f32],
bias: Option<&[f32]>,
n_in: usize,
n_out: usize,
) {
sgemm_forward(y, x, w, bias, 1, n_in, n_out);
}
pub fn sgemm_backward(
dx: &mut [f32],
dw: &mut [f32],
db: Option<&mut [f32]>,
dy: &[f32],
x_saved: &[f32],
w: &[f32],
dims: (usize, usize, usize), ) {
let (batch, n_in, n_out) = dims;
dx[..batch * n_in].fill(0.0);
#[cfg(all(feature = "accelerate", target_os = "macos"))]
unsafe {
cblas_sgemm(
101, 111, 112, batch as i32, n_in as i32, n_out as i32, 1.0,
dy.as_ptr(),
n_out as i32,
w.as_ptr(),
n_out as i32, 0.0,
dx.as_mut_ptr(),
n_in as i32,
);
}
#[cfg(all(
feature = "gemm-blas",
not(all(feature = "accelerate", target_os = "macos"))
))]
unsafe {
gemm::gemm(
batch,
n_in,
n_out,
dx.as_mut_ptr(),
1, n_in as isize, false,
dy.as_ptr(),
1, n_out as isize, w.as_ptr(),
n_out as isize, 1, 0.0, 1.0, false,
false,
false,
gemm::Parallelism::None,
);
}
#[cfg(not(any(
all(feature = "accelerate", target_os = "macos"),
feature = "gemm-blas"
)))]
{
for row in 0..batch {
let dy_off = row * n_out;
let dx_off = row * n_in;
for j in 0..n_out {
let dv = dy[dy_off + j];
for k in 0..n_in {
dx[dx_off + k] += dv * w[k * n_out + j];
}
}
}
}
#[cfg(all(feature = "accelerate", target_os = "macos"))]
unsafe {
cblas_sgemm(
101,
112, 111, n_in as i32, n_out as i32, batch as i32, 1.0,
x_saved.as_ptr(),
n_in as i32, dy.as_ptr(),
n_out as i32,
1.0, dw.as_mut_ptr(),
n_out as i32,
);
}
#[cfg(all(
feature = "gemm-blas",
not(all(feature = "accelerate", target_os = "macos"))
))]
unsafe {
gemm::gemm(
n_in,
n_out,
batch,
dw.as_mut_ptr(),
1, n_out as isize, true, x_saved.as_ptr(),
n_in as isize, 1, dy.as_ptr(),
1, n_out as isize, 1.0, 1.0, false,
false,
false,
gemm::Parallelism::None,
);
}
#[cfg(not(any(
all(feature = "accelerate", target_os = "macos"),
feature = "gemm-blas"
)))]
{
for row in 0..batch {
let x_off = row * n_in;
let dy_off = row * n_out;
for k in 0..n_in {
let xv = x_saved[x_off + k];
let w_off = k * n_out;
for j in 0..n_out {
dw[w_off + j] += xv * dy[dy_off + j];
}
}
}
}
if let Some(db) = db {
for row in 0..batch {
let dy_off = row * n_out;
for j in 0..n_out {
db[j] += dy[dy_off + j];
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_sgemm_forward_identity() {
let n = 4;
let mut w = vec![0.0; n * n];
for i in 0..n {
w[i * n + i] = 1.0;
}
let x: Vec<f32> = (0..n).map(|i| (i + 1) as f32).collect();
let mut y = vec![0.0; n];
sgemm_forward(&mut y, &x, &w, None, 1, n, n);
for (i, yi) in y.iter().enumerate().take(n) {
assert!((*yi - (i + 1) as f32).abs() < 1e-6);
}
}
#[test]
fn test_sgemm_forward_with_bias() {
let w = vec![1.0, 0.0, 0.0, 1.0];
let x = vec![3.0, 4.0];
let bias = vec![10.0, 20.0];
let mut y = vec![0.0; 2];
sgemm_forward(&mut y, &x, &w, Some(&bias), 1, 2, 2);
assert!((y[0] - 13.0).abs() < 1e-6);
assert!((y[1] - 24.0).abs() < 1e-6);
}
#[test]
fn test_sgemm_backward_gradient() {
let batch = 2;
let n_in = 3;
let n_out = 2;
let w = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
let x = vec![1.0, 0.0, 0.0, 0.0, 1.0, 0.0];
let dy = vec![1.0, 1.0, 1.0, 1.0];
let mut dx = vec![0.0; batch * n_in];
let mut dw = vec![0.0; n_in * n_out];
sgemm_backward(&mut dx, &mut dw, None, &dy, &x, &w, (batch, n_in, n_out));
assert!((dx[0] - 3.0).abs() < 1e-5);
assert!((dx[1] - 7.0).abs() < 1e-5);
assert!((dx[2] - 11.0).abs() < 1e-5);
}
}