#[cfg(gemmology_simd)]
mod imp {
use std::os::raw::{c_char, c_void};
extern "C" {
fn gemmology_prepare_b(b_transposed: *const i8, n: usize, k: usize) -> *mut c_void;
fn gemmology_free_b(handle: *mut c_void);
fn gemmology_multiply(
handle: *mut c_void,
a: *const u8,
m: usize,
unquant: f32,
bias: *const f32,
out: *mut f32,
);
fn gemmology_prepared_bytes() -> usize;
fn gemmology_read_row(handle: *const c_void, id: usize, out: *mut i8);
fn gemmology_backend_name() -> *const c_char;
fn gemmology_gemm_threads() -> usize;
}
pub fn gemm_threads() -> usize {
unsafe { gemmology_gemm_threads() }
}
pub fn prepared_bytes() -> usize {
unsafe { gemmology_prepared_bytes() }
}
pub fn backend() -> &'static str {
let name = unsafe { std::ffi::CStr::from_ptr(gemmology_backend_name()) };
name.to_str().unwrap_or("unknown")
}
pub struct PreparedB {
handle: *mut c_void,
n: usize,
k: usize,
}
impl PreparedB {
pub fn new(b_transposed: &[i8], n: usize, k: usize) -> Option<PreparedB> {
assert_eq!(b_transposed.len(), n * k, "B length must be n * k");
let handle = unsafe { gemmology_prepare_b(b_transposed.as_ptr(), n, k) };
if handle.is_null() {
None
} else {
Some(PreparedB { handle, n, k })
}
}
pub fn matmul(&self, a: &[u8], m: usize, unquant: f32, bias: &[f32]) -> Vec<f32> {
let mut out = Vec::new();
self.matmul_into(a, m, unquant, bias, &mut out);
out
}
pub fn matmul_into(
&self,
a: &[u8],
m: usize,
unquant: f32,
bias: &[f32],
out: &mut Vec<f32>,
) {
assert_eq!(a.len(), m * self.k, "A length must be m * k");
assert_eq!(bias.len(), self.n, "bias length must be n");
out.clear();
out.resize(m * self.n, 0.0);
unsafe {
gemmology_multiply(
self.handle,
a.as_ptr(),
m,
unquant,
bias.as_ptr(),
out.as_mut_ptr(),
);
}
}
pub fn read_row(&self, id: usize, out: &mut [i8]) {
assert_eq!(out.len(), self.k, "out length must be k");
assert!(id < self.n, "row id {id} out of range (n = {})", self.n);
unsafe { gemmology_read_row(self.handle, id, out.as_mut_ptr()) };
}
}
impl Drop for PreparedB {
fn drop(&mut self) {
unsafe { gemmology_free_b(self.handle) };
}
}
unsafe impl Sync for PreparedB {}
}
#[cfg(all(
not(gemmology_simd),
target_arch = "wasm32",
target_feature = "simd128"
))]
mod imp {
use core::arch::wasm32::*;
pub fn prepared_bytes() -> usize {
0
}
pub fn backend() -> &'static str {
"wasm-simd128"
}
pub fn gemm_threads() -> usize {
1
}
const LANES: usize = 16;
pub struct PreparedB {
packed: Vec<i8>,
n: usize,
k: usize,
k_padded: usize,
}
impl PreparedB {
pub fn new(b_transposed: &[i8], n: usize, k: usize) -> Option<PreparedB> {
assert_eq!(b_transposed.len(), n * k, "B length must be n * k");
if k % LANES != 0 {
return None;
}
let k_padded = k;
Some(PreparedB {
packed: b_transposed.to_vec(),
n,
k,
k_padded,
})
}
pub fn matmul(&self, a: &[u8], m: usize, unquant: f32, bias: &[f32]) -> Vec<f32> {
let mut out = Vec::new();
self.matmul_into(a, m, unquant, bias, &mut out);
out
}
pub fn matmul_into(
&self,
a: &[u8],
m: usize,
unquant: f32,
bias: &[f32],
out: &mut Vec<f32>,
) {
assert_eq!(a.len(), m * self.k, "A length must be m * k");
assert_eq!(bias.len(), self.n, "bias length must be n");
out.clear();
out.resize(m * self.n, 0.0);
for row in 0..m {
let a_row = &a[row * self.k..(row + 1) * self.k];
for col in 0..self.n {
let b_col = &self.packed[col * self.k_padded..col * self.k_padded + self.k];
let acc = unsafe { dot(a_row, b_col) };
out[row * self.n + col] = unquant * acc as f32 + bias[col];
}
}
}
pub fn read_row(&self, id: usize, out: &mut [i8]) {
assert_eq!(out.len(), self.k, "out length must be k");
assert!(id < self.n, "row id {id} out of range (n = {})", self.n);
let base = id * self.k_padded;
out.copy_from_slice(&self.packed[base..base + self.k]);
}
}
#[target_feature(enable = "simd128")]
unsafe fn dot(a: &[u8], b: &[i8]) -> i32 {
debug_assert_eq!(a.len(), b.len());
debug_assert_eq!(a.len() % LANES, 0);
let mut acc = i32x4_splat(0);
let mut off = 0;
while off < a.len() {
let av = v128_load(a.as_ptr().add(off) as *const v128);
let bv = v128_load(b.as_ptr().add(off) as *const v128);
let a_lo = i16x8_extend_low_u8x16(av);
let a_hi = i16x8_extend_high_u8x16(av);
let b_lo = i16x8_extend_low_i8x16(bv);
let b_hi = i16x8_extend_high_i8x16(bv);
acc = i32x4_add(acc, i32x4_dot_i16x8(a_lo, b_lo));
acc = i32x4_add(acc, i32x4_dot_i16x8(a_hi, b_hi));
off += LANES;
}
i32x4_extract_lane::<0>(acc)
+ i32x4_extract_lane::<1>(acc)
+ i32x4_extract_lane::<2>(acc)
+ i32x4_extract_lane::<3>(acc)
}
}
#[cfg(all(
not(gemmology_simd),
not(all(target_arch = "wasm32", target_feature = "simd128"))
))]
mod imp {
pub fn prepared_bytes() -> usize {
0
}
pub fn backend() -> &'static str {
"scalar"
}
pub fn gemm_threads() -> usize {
0
}
pub struct PreparedB {
_never: (),
}
impl PreparedB {
pub fn new(b_transposed: &[i8], n: usize, k: usize) -> Option<PreparedB> {
debug_assert_eq!(b_transposed.len(), n * k, "B length must be n * k");
None
}
pub fn matmul(&self, _a: &[u8], _m: usize, _unquant: f32, _bias: &[f32]) -> Vec<f32> {
unreachable!("scalar-fallback PreparedB is never constructed")
}
pub fn matmul_into(
&self,
_a: &[u8],
_m: usize,
_unquant: f32,
_bias: &[f32],
_out: &mut Vec<f32>,
) {
unreachable!("scalar-fallback PreparedB is never constructed")
}
pub fn read_row(&self, _id: usize, _out: &mut [i8]) {
unreachable!("scalar-fallback PreparedB is never constructed")
}
}
}
pub use imp::{backend, gemm_threads, prepared_bytes, PreparedB};