use std::ffi::CStr;
use std::os::raw::c_char;
use std::sync::{Mutex, OnceLock};
use crate::{
kernels, Dispatch, KernelInfo,
DotKernel, DotF32Kernel,
DotI32Kernel, DotI64Kernel, DotI16Kernel, DotI8Kernel,
ReduceKernel, ReduceSumF32Kernel,
ReduceI32Kernel, ReduceI64Kernel, ReduceI16Kernel, ReduceI8Kernel,
SoftmaxKernel, UnaryOpF32Kernel,
ClampKernel, ClampF32Kernel,
BinaryOpKernel, AddF32Kernel,
BinaryI32Kernel, BinaryI64Kernel, BinaryI16Kernel, BinaryI8Kernel,
UnaryI32Kernel,
MatMulKernel, MatMulF32Kernel,
ArgmaxKernel, ArgmaxF32Kernel,
ArgmaxI32Kernel, ArgmaxI64Kernel, ArgmaxI16Kernel, ArgmaxI8Kernel,
MemchrKernel,
MatMulBiasReluKernel,
SortI32Kernel, SortF64Kernel, SortF32Kernel,
GemvF64Kernel, GemvF32Kernel,
LayerNormF64Kernel, LayerNormF32Kernel,
};
use himada_core::HardwareDNA;
pub struct HwdnaDispatch(Dispatch<DotKernel>);
#[no_mangle]
pub unsafe extern "C" fn himada_dispatch_create(name: *const c_char) -> *mut HwdnaDispatch {
let name = if name.is_null() {
"ffi-dispatch".into()
} else {
unsafe { CStr::from_ptr(name) }
.to_str()
.unwrap_or("ffi-dispatch")
.to_string()
};
let scalar_fn: DotKernel = kernels::dot_scalar;
let dispatch = Dispatch::new(&name, vec![
KernelInfo {
name: "scalar",
func: scalar_fn,
is_supported: |_: &HardwareDNA| true,
thermal_priority: 1,
},
#[cfg(target_arch = "x86_64")]
KernelInfo {
name: "SSE",
func: kernels::dot_sse as DotKernel,
is_supported: crate::kernels::dot_sse_supported,
thermal_priority: 2,
},
#[cfg(target_arch = "x86_64")]
KernelInfo {
name: "AVX2",
func: kernels::dot_avx2 as DotKernel,
is_supported: crate::kernels::dot_avx2_supported,
thermal_priority: 3,
},
#[cfg(target_arch = "aarch64")]
KernelInfo {
name: "NEON",
func: kernels::dot_neon as DotKernel,
is_supported: crate::kernels::dot_neon_supported,
thermal_priority: 2,
},
]);
Box::into_raw(Box::new(HwdnaDispatch(dispatch)))
}
#[no_mangle]
pub unsafe extern "C" fn himada_dispatch_destroy(handle: *mut HwdnaDispatch) {
if !handle.is_null() {
unsafe { drop(Box::from_raw(handle)) };
}
}
#[no_mangle]
pub unsafe extern "C" fn himada_dispatch_dot(
handle: *mut HwdnaDispatch,
a: *const f64,
b: *const f64,
n: usize,
) -> f64 {
if handle.is_null() || a.is_null() || b.is_null() || n == 0 {
return 0.0;
}
let dispatch = unsafe { &mut (*handle).0 };
let a_slice = unsafe { std::slice::from_raw_parts(a, n) };
let b_slice = unsafe { std::slice::from_raw_parts(b, n) };
dispatch.compute(a_slice, b_slice)
}
#[no_mangle]
pub unsafe extern "C" fn himada_dispatch_selected_name(
handle: *const HwdnaDispatch,
buf: *mut c_char,
max_len: usize,
) {
if handle.is_null() || buf.is_null() || max_len == 0 {
return;
}
let name = unsafe { (*handle).0.selected_name() };
let bytes = name.as_bytes();
let len = bytes.len().min(max_len.saturating_sub(1));
unsafe { std::ptr::copy_nonoverlapping(bytes.as_ptr(), buf as *mut u8, len) };
unsafe { *buf.add(len) = 0 };
}
#[no_mangle]
pub unsafe extern "C" fn himada_dot(a: *const f64, b: *const f64, n: usize) -> f64 {
if a.is_null() || b.is_null() || n == 0 {
return 0.0;
}
let handle = himada_dispatch_create(std::ptr::null());
let result = himada_dispatch_dot(handle, a, b, n);
himada_dispatch_destroy(handle);
result
}
#[allow(dead_code)]
fn add_sse_kernel<T: Copy>(name: &'static str, func: T, supported: fn(&HardwareDNA) -> bool, priority: u8) -> KernelInfo<T> {
KernelInfo { name, func, is_supported: supported, thermal_priority: priority }
}
#[allow(dead_code)]
fn add_avx2_kernel<T: Copy>(name: &'static str, func: T, supported: fn(&HardwareDNA) -> bool, priority: u8) -> KernelInfo<T> {
KernelInfo { name, func, is_supported: supported, thermal_priority: priority }
}
#[allow(dead_code)]
fn add_neon_kernel<T: Copy>(name: &'static str, func: T, supported: fn(&HardwareDNA) -> bool, priority: u8) -> KernelInfo<T> {
KernelInfo { name, func, is_supported: supported, thermal_priority: priority }
}
fn reduce_sum_dispatch() -> &'static Mutex<Dispatch<ReduceKernel>> {
static D: OnceLock<Mutex<Dispatch<ReduceKernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("reduce_sum_f64_ffi", vec![
KernelInfo { name: "scalar", func: kernels::reduce_sum_scalar, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::reduce_sum_sse as ReduceKernel, kernels::reduce_sum_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::reduce_sum_avx2 as ReduceKernel, kernels::reduce_sum_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::reduce_sum_neon as ReduceKernel, kernels::reduce_sum_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_reduce_sum(a: *const f64, n: usize) -> f64 {
if a.is_null() || n == 0 {
return 0.0;
}
let a = unsafe { std::slice::from_raw_parts(a, n) };
let mut guard = match reduce_sum_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return 0.0,
};
guard.compute(a)
}
fn matmul_dispatch() -> &'static Mutex<Dispatch<MatMulKernel>> {
static D: OnceLock<Mutex<Dispatch<MatMulKernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("matmul_f64_ffi", vec![
KernelInfo { name: "scalar", func: kernels::matmul_scalar, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
KernelInfo { name: "tiled", func: kernels::matmul_tiled, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
KernelInfo { name: "cache_tiled", func: kernels::matmul_cache_tiled, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::matmul_sse as MatMulKernel, kernels::matmul_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::matmul_avx2 as MatMulKernel, kernels::matmul_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::matmul_neon as MatMulKernel, kernels::matmul_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_matmul(
a: *const f64,
b: *const f64,
c: *mut f64,
n: usize,
) {
if a.is_null() || b.is_null() || c.is_null() || n == 0 {
return;
}
let a = unsafe { std::slice::from_raw_parts(a, n * n) };
let b = unsafe { std::slice::from_raw_parts(b, n * n) };
let c = unsafe { std::slice::from_raw_parts_mut(c, n * n) };
c.fill(0.0);
let mut guard = match matmul_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return,
};
guard.compute(a, b, c, n)
}
fn dot_f32_dispatch() -> &'static Mutex<Dispatch<DotF32Kernel>> {
static D: OnceLock<Mutex<Dispatch<DotF32Kernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("dot_f32_ffi", vec![
KernelInfo { name: "scalar", func: kernels::dot_f32_scalar, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::dot_f32_sse as DotF32Kernel, kernels::dot_f32_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::dot_f32_avx2 as DotF32Kernel, kernels::dot_f32_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::dot_f32_neon as DotF32Kernel, kernels::dot_f32_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_dot_f32(a: *const f32, b: *const f32, n: usize) -> f32 {
if a.is_null() || b.is_null() || n == 0 {
return 0.0;
}
let a = unsafe { std::slice::from_raw_parts(a, n) };
let b = unsafe { std::slice::from_raw_parts(b, n) };
let mut guard = match dot_f32_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return 0.0,
};
guard.compute(a, b)
}
fn euclidean_f64_dispatch() -> &'static Mutex<Dispatch<DotKernel>> {
static D: OnceLock<Mutex<Dispatch<DotKernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("euclidean_f64_ffi", vec![
KernelInfo { name: "scalar", func: kernels::euclidean_distance_f64_scalar as DotKernel, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::euclidean_distance_f64_sse as DotKernel, kernels::euclidean_distance_f64_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::euclidean_distance_f64_avx2 as DotKernel, kernels::euclidean_distance_f64_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::euclidean_distance_f64_neon as DotKernel, kernels::euclidean_distance_f64_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_euclidean_distance(a: *const f64, b: *const f64, n: usize) -> f64 {
if a.is_null() || b.is_null() || n == 0 {
return 0.0;
}
let a = unsafe { std::slice::from_raw_parts(a, n) };
let b = unsafe { std::slice::from_raw_parts(b, n) };
let mut guard = match euclidean_f64_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return 0.0,
};
guard.compute(a, b)
}
fn euclidean_f32_dispatch() -> &'static Mutex<Dispatch<DotF32Kernel>> {
static D: OnceLock<Mutex<Dispatch<DotF32Kernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("euclidean_f32_ffi", vec![
KernelInfo { name: "scalar", func: kernels::euclidean_distance_f32_scalar as DotF32Kernel, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::euclidean_distance_f32_sse as DotF32Kernel, kernels::euclidean_distance_f32_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::euclidean_distance_f32_avx2 as DotF32Kernel, kernels::euclidean_distance_f32_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::euclidean_distance_f32_neon as DotF32Kernel, kernels::euclidean_distance_f32_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_euclidean_distance_f32(a: *const f32, b: *const f32, n: usize) -> f32 {
if a.is_null() || b.is_null() || n == 0 {
return 0.0;
}
let a = unsafe { std::slice::from_raw_parts(a, n) };
let b = unsafe { std::slice::from_raw_parts(b, n) };
let mut guard = match euclidean_f32_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return 0.0,
};
guard.compute(a, b)
}
fn cosine_f64_dispatch() -> &'static Mutex<Dispatch<DotKernel>> {
static D: OnceLock<Mutex<Dispatch<DotKernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("cosine_f64_ffi", vec![
KernelInfo { name: "scalar", func: kernels::cosine_similarity_f64_scalar as DotKernel, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::cosine_similarity_f64_sse as DotKernel, kernels::cosine_similarity_f64_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::cosine_similarity_f64_avx2 as DotKernel, kernels::cosine_similarity_f64_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::cosine_similarity_f64_neon as DotKernel, kernels::cosine_similarity_f64_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_cosine_similarity(a: *const f64, b: *const f64, n: usize) -> f64 {
if a.is_null() || b.is_null() || n == 0 {
return 0.0;
}
let a = unsafe { std::slice::from_raw_parts(a, n) };
let b = unsafe { std::slice::from_raw_parts(b, n) };
let mut guard = match cosine_f64_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return 0.0,
};
guard.compute(a, b)
}
fn cosine_f32_dispatch() -> &'static Mutex<Dispatch<DotF32Kernel>> {
static D: OnceLock<Mutex<Dispatch<DotF32Kernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("cosine_f32_ffi", vec![
KernelInfo { name: "scalar", func: kernels::cosine_similarity_f32_scalar as DotF32Kernel, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::cosine_similarity_f32_sse as DotF32Kernel, kernels::cosine_similarity_f32_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::cosine_similarity_f32_avx2 as DotF32Kernel, kernels::cosine_similarity_f32_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::cosine_similarity_f32_neon as DotF32Kernel, kernels::cosine_similarity_f32_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_cosine_similarity_f32(a: *const f32, b: *const f32, n: usize) -> f32 {
if a.is_null() || b.is_null() || n == 0 {
return 0.0;
}
let a = unsafe { std::slice::from_raw_parts(a, n) };
let b = unsafe { std::slice::from_raw_parts(b, n) };
let mut guard = match cosine_f32_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return 0.0,
};
guard.compute(a, b)
}
fn reduce_sum_f32_dispatch() -> &'static Mutex<Dispatch<ReduceSumF32Kernel>> {
static D: OnceLock<Mutex<Dispatch<ReduceSumF32Kernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("reduce_sum_f32_ffi", vec![
KernelInfo { name: "scalar", func: kernels::reduce_sum_f32_scalar, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::reduce_sum_f32_sse as ReduceSumF32Kernel, kernels::reduce_sum_f32_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::reduce_sum_f32_avx2 as ReduceSumF32Kernel, kernels::reduce_sum_f32_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::reduce_sum_f32_neon as ReduceSumF32Kernel, kernels::reduce_sum_f32_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_reduce_sum_f32(a: *const f32, n: usize) -> f32 {
if a.is_null() || n == 0 {
return 0.0;
}
let a = unsafe { std::slice::from_raw_parts(a, n) };
let mut guard = match reduce_sum_f32_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return 0.0,
};
guard.compute(a)
}
fn reduce_max_f64_dispatch() -> &'static Mutex<Dispatch<ReduceKernel>> {
static D: OnceLock<Mutex<Dispatch<ReduceKernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("reduce_max_f64_ffi", vec![
KernelInfo { name: "scalar", func: kernels::reduce_max_scalar, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::reduce_max_sse as ReduceKernel, kernels::reduce_max_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::reduce_max_avx2 as ReduceKernel, kernels::reduce_max_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::reduce_max_neon as ReduceKernel, kernels::reduce_max_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_reduce_max(a: *const f64, n: usize) -> f64 {
if a.is_null() || n == 0 {
return 0.0;
}
let a = unsafe { std::slice::from_raw_parts(a, n) };
let mut guard = match reduce_max_f64_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return 0.0,
};
guard.compute(a)
}
fn reduce_max_f32_dispatch() -> &'static Mutex<Dispatch<ReduceSumF32Kernel>> {
static D: OnceLock<Mutex<Dispatch<ReduceSumF32Kernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("reduce_max_f32_ffi", vec![
KernelInfo { name: "scalar", func: kernels::reduce_max_f32_scalar as ReduceSumF32Kernel, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::reduce_max_f32_sse as ReduceSumF32Kernel, kernels::reduce_max_f32_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::reduce_max_f32_avx2 as ReduceSumF32Kernel, kernels::reduce_max_f32_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::reduce_max_f32_neon as ReduceSumF32Kernel, kernels::reduce_max_f32_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_reduce_max_f32(a: *const f32, n: usize) -> f32 {
if a.is_null() || n == 0 {
return 0.0;
}
let a = unsafe { std::slice::from_raw_parts(a, n) };
let mut guard = match reduce_max_f32_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return 0.0,
};
guard.compute(a)
}
fn abs_max_f64_dispatch() -> &'static Mutex<Dispatch<ReduceKernel>> {
static D: OnceLock<Mutex<Dispatch<ReduceKernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("abs_max_f64_ffi", vec![
KernelInfo { name: "scalar", func: kernels::abs_max_f64_scalar as ReduceKernel, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::abs_max_f64_sse as ReduceKernel, kernels::abs_max_f64_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::abs_max_f64_avx2 as ReduceKernel, kernels::abs_max_f64_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::abs_max_f64_neon as ReduceKernel, kernels::abs_max_f64_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_abs_max(a: *const f64, n: usize) -> f64 {
if a.is_null() || n == 0 {
return 0.0;
}
let a = unsafe { std::slice::from_raw_parts(a, n) };
let mut guard = match abs_max_f64_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return 0.0,
};
guard.compute(a)
}
fn abs_max_f32_dispatch() -> &'static Mutex<Dispatch<ReduceSumF32Kernel>> {
static D: OnceLock<Mutex<Dispatch<ReduceSumF32Kernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("abs_max_f32_ffi", vec![
KernelInfo { name: "scalar", func: kernels::abs_max_f32_scalar as ReduceSumF32Kernel, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::abs_max_f32_sse as ReduceSumF32Kernel, kernels::abs_max_f32_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::abs_max_f32_avx2 as ReduceSumF32Kernel, kernels::abs_max_f32_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::abs_max_f32_neon as ReduceSumF32Kernel, kernels::abs_max_f32_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_abs_max_f32(a: *const f32, n: usize) -> f32 {
if a.is_null() || n == 0 {
return 0.0;
}
let a = unsafe { std::slice::from_raw_parts(a, n) };
let mut guard = match abs_max_f32_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return 0.0,
};
guard.compute(a)
}
fn argmax_f64_dispatch() -> &'static Mutex<Dispatch<ArgmaxKernel>> {
static D: OnceLock<Mutex<Dispatch<ArgmaxKernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("argmax_f64_ffi", vec![
KernelInfo { name: "scalar", func: kernels::argmax_f64_scalar, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::argmax_f64_sse as ArgmaxKernel, kernels::argmax_f64_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::argmax_f64_avx2 as ArgmaxKernel, kernels::argmax_f64_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::argmax_f64_neon as ArgmaxKernel, kernels::argmax_f64_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_argmax(a: *const f64, n: usize) -> usize {
if a.is_null() || n == 0 {
return 0;
}
let a = unsafe { std::slice::from_raw_parts(a, n) };
let mut guard = match argmax_f64_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return 0,
};
guard.compute(a)
}
fn argmax_f32_dispatch() -> &'static Mutex<Dispatch<ArgmaxF32Kernel>> {
static D: OnceLock<Mutex<Dispatch<ArgmaxF32Kernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("argmax_f32_ffi", vec![
KernelInfo { name: "scalar", func: kernels::argmax_f32_scalar, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::argmax_f32_sse as ArgmaxF32Kernel, kernels::argmax_f32_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::argmax_f32_avx2 as ArgmaxF32Kernel, kernels::argmax_f32_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::argmax_f32_neon as ArgmaxF32Kernel, kernels::argmax_f32_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_argmax_f32(a: *const f32, n: usize) -> usize {
if a.is_null() || n == 0 {
return 0;
}
let a = unsafe { std::slice::from_raw_parts(a, n) };
let mut guard = match argmax_f32_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return 0,
};
guard.compute(a)
}
fn softmax_f64_dispatch() -> &'static Mutex<Dispatch<SoftmaxKernel>> {
static D: OnceLock<Mutex<Dispatch<SoftmaxKernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("softmax_f64_ffi", vec![
KernelInfo { name: "scalar", func: kernels::softmax_scalar, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::softmax_sse as SoftmaxKernel, kernels::softmax_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::softmax_avx2 as SoftmaxKernel, kernels::softmax_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::softmax_neon as SoftmaxKernel, kernels::softmax_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_softmax(input: *const f64, output: *mut f64, n: usize) {
if input.is_null() || output.is_null() || n == 0 {
return;
}
let input = unsafe { std::slice::from_raw_parts(input, n) };
let output = unsafe { std::slice::from_raw_parts_mut(output, n) };
let mut guard = match softmax_f64_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return,
};
guard.compute(input, output)
}
fn softmax_f32_dispatch() -> &'static Mutex<Dispatch<UnaryOpF32Kernel>> {
static D: OnceLock<Mutex<Dispatch<UnaryOpF32Kernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("softmax_f32_ffi", vec![
KernelInfo { name: "scalar", func: kernels::softmax_f32_scalar as UnaryOpF32Kernel, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::softmax_f32_sse as UnaryOpF32Kernel, kernels::softmax_f32_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::softmax_f32_avx2 as UnaryOpF32Kernel, kernels::softmax_f32_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::softmax_f32_neon as UnaryOpF32Kernel, kernels::softmax_f32_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_softmax_f32(input: *const f32, output: *mut f32, n: usize) {
if input.is_null() || output.is_null() || n == 0 {
return;
}
let input = unsafe { std::slice::from_raw_parts(input, n) };
let output = unsafe { std::slice::from_raw_parts_mut(output, n) };
let mut guard = match softmax_f32_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return,
};
guard.compute(input, output)
}
fn hadamard_f64_dispatch() -> &'static Mutex<Dispatch<BinaryOpKernel>> {
static D: OnceLock<Mutex<Dispatch<BinaryOpKernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("hadamard_f64_ffi", vec![
KernelInfo { name: "scalar", func: kernels::hadamard_product_f64_scalar, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::hadamard_product_f64_sse as BinaryOpKernel, kernels::hadamard_product_f64_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::hadamard_product_f64_avx2 as BinaryOpKernel, kernels::hadamard_product_f64_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::hadamard_product_f64_neon as BinaryOpKernel, kernels::hadamard_product_f64_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_hadamard_product(a: *const f64, b: *const f64, c: *mut f64, n: usize) {
if a.is_null() || b.is_null() || c.is_null() || n == 0 {
return;
}
let a = unsafe { std::slice::from_raw_parts(a, n) };
let b = unsafe { std::slice::from_raw_parts(b, n) };
let c = unsafe { std::slice::from_raw_parts_mut(c, n) };
let mut guard = match hadamard_f64_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return,
};
guard.compute(a, b, c)
}
fn hadamard_f32_dispatch() -> &'static Mutex<Dispatch<AddF32Kernel>> {
static D: OnceLock<Mutex<Dispatch<AddF32Kernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("hadamard_f32_ffi", vec![
KernelInfo { name: "scalar", func: kernels::hadamard_product_f32_scalar as AddF32Kernel, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::hadamard_product_f32_sse as AddF32Kernel, kernels::hadamard_product_f32_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::hadamard_product_f32_avx2 as AddF32Kernel, kernels::hadamard_product_f32_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::hadamard_product_f32_neon as AddF32Kernel, kernels::hadamard_product_f32_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_hadamard_product_f32(a: *const f32, b: *const f32, c: *mut f32, n: usize) {
if a.is_null() || b.is_null() || c.is_null() || n == 0 {
return;
}
let a = unsafe { std::slice::from_raw_parts(a, n) };
let b = unsafe { std::slice::from_raw_parts(b, n) };
let c = unsafe { std::slice::from_raw_parts_mut(c, n) };
let mut guard = match hadamard_f32_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return,
};
guard.compute(a, b, c)
}
fn add_f64_dispatch() -> &'static Mutex<Dispatch<BinaryOpKernel>> {
static D: OnceLock<Mutex<Dispatch<BinaryOpKernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("add_f64_ffi", vec![
KernelInfo { name: "scalar", func: kernels::add_f64_scalar, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_add(a: *const f64, b: *const f64, c: *mut f64, n: usize) {
if a.is_null() || b.is_null() || c.is_null() || n == 0 {
return;
}
let a = unsafe { std::slice::from_raw_parts(a, n) };
let b = unsafe { std::slice::from_raw_parts(b, n) };
let c = unsafe { std::slice::from_raw_parts_mut(c, n) };
let mut guard = match add_f64_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return,
};
guard.compute(a, b, c)
}
fn add_f32_dispatch() -> &'static Mutex<Dispatch<AddF32Kernel>> {
static D: OnceLock<Mutex<Dispatch<AddF32Kernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("add_f32_ffi", vec![
KernelInfo { name: "scalar", func: kernels::add_f32_scalar, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::add_f32_sse as AddF32Kernel, kernels::add_f32_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::add_f32_avx2 as AddF32Kernel, kernels::add_f32_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::add_f32_neon as AddF32Kernel, kernels::add_f32_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_add_f32(a: *const f32, b: *const f32, c: *mut f32, n: usize) {
if a.is_null() || b.is_null() || c.is_null() || n == 0 {
return;
}
let a = unsafe { std::slice::from_raw_parts(a, n) };
let b = unsafe { std::slice::from_raw_parts(b, n) };
let c = unsafe { std::slice::from_raw_parts_mut(c, n) };
let mut guard = match add_f32_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return,
};
guard.compute(a, b, c)
}
fn sub_f64_dispatch() -> &'static Mutex<Dispatch<BinaryOpKernel>> {
static D: OnceLock<Mutex<Dispatch<BinaryOpKernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("sub_f64_ffi", vec![
KernelInfo { name: "scalar", func: kernels::sub_f64_scalar, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_sub(a: *const f64, b: *const f64, c: *mut f64, n: usize) {
if a.is_null() || b.is_null() || c.is_null() || n == 0 {
return;
}
let a = unsafe { std::slice::from_raw_parts(a, n) };
let b = unsafe { std::slice::from_raw_parts(b, n) };
let c = unsafe { std::slice::from_raw_parts_mut(c, n) };
let mut guard = match sub_f64_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return,
};
guard.compute(a, b, c)
}
fn sub_f32_dispatch() -> &'static Mutex<Dispatch<AddF32Kernel>> {
static D: OnceLock<Mutex<Dispatch<AddF32Kernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("sub_f32_ffi", vec![
KernelInfo { name: "scalar", func: kernels::sub_f32_scalar, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::sub_f32_sse as AddF32Kernel, kernels::sub_f32_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::sub_f32_avx2 as AddF32Kernel, kernels::sub_f32_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::sub_f32_neon as AddF32Kernel, kernels::sub_f32_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_sub_f32(a: *const f32, b: *const f32, c: *mut f32, n: usize) {
if a.is_null() || b.is_null() || c.is_null() || n == 0 {
return;
}
let a = unsafe { std::slice::from_raw_parts(a, n) };
let b = unsafe { std::slice::from_raw_parts(b, n) };
let c = unsafe { std::slice::from_raw_parts_mut(c, n) };
let mut guard = match sub_f32_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return,
};
guard.compute(a, b, c)
}
fn mul_f64_dispatch() -> &'static Mutex<Dispatch<BinaryOpKernel>> {
static D: OnceLock<Mutex<Dispatch<BinaryOpKernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("mul_f64_ffi", vec![
KernelInfo { name: "scalar", func: kernels::mul_f64_scalar, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_mul(a: *const f64, b: *const f64, c: *mut f64, n: usize) {
if a.is_null() || b.is_null() || c.is_null() || n == 0 {
return;
}
let a = unsafe { std::slice::from_raw_parts(a, n) };
let b = unsafe { std::slice::from_raw_parts(b, n) };
let c = unsafe { std::slice::from_raw_parts_mut(c, n) };
let mut guard = match mul_f64_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return,
};
guard.compute(a, b, c)
}
fn mul_f32_dispatch() -> &'static Mutex<Dispatch<AddF32Kernel>> {
static D: OnceLock<Mutex<Dispatch<AddF32Kernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("mul_f32_ffi", vec![
KernelInfo { name: "scalar", func: kernels::mul_f32_scalar, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::mul_f32_sse as AddF32Kernel, kernels::mul_f32_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::mul_f32_avx2 as AddF32Kernel, kernels::mul_f32_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::mul_f32_neon as AddF32Kernel, kernels::mul_f32_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_mul_f32(a: *const f32, b: *const f32, c: *mut f32, n: usize) {
if a.is_null() || b.is_null() || c.is_null() || n == 0 {
return;
}
let a = unsafe { std::slice::from_raw_parts(a, n) };
let b = unsafe { std::slice::from_raw_parts(b, n) };
let c = unsafe { std::slice::from_raw_parts_mut(c, n) };
let mut guard = match mul_f32_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return,
};
guard.compute(a, b, c)
}
fn negate_f64_dispatch() -> &'static Mutex<Dispatch<SoftmaxKernel>> {
static D: OnceLock<Mutex<Dispatch<SoftmaxKernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("negate_f64_ffi", vec![
KernelInfo { name: "scalar", func: kernels::negate_f64_scalar as SoftmaxKernel, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::negate_f64_sse as SoftmaxKernel, kernels::negate_f64_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::negate_f64_avx2 as SoftmaxKernel, kernels::negate_f64_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::negate_f64_neon as SoftmaxKernel, kernels::negate_f64_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_negate(a: *const f64, c: *mut f64, n: usize) {
if a.is_null() || c.is_null() || n == 0 {
return;
}
let a = unsafe { std::slice::from_raw_parts(a, n) };
let c = unsafe { std::slice::from_raw_parts_mut(c, n) };
let mut guard = match negate_f64_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return,
};
guard.compute(a, c)
}
fn negate_f32_dispatch() -> &'static Mutex<Dispatch<UnaryOpF32Kernel>> {
static D: OnceLock<Mutex<Dispatch<UnaryOpF32Kernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("negate_f32_ffi", vec![
KernelInfo { name: "scalar", func: kernels::negate_f32_scalar, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::negate_f32_sse as UnaryOpF32Kernel, kernels::negate_f32_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::negate_f32_avx2 as UnaryOpF32Kernel, kernels::negate_f32_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::negate_f32_neon as UnaryOpF32Kernel, kernels::negate_f32_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_negate_f32(a: *const f32, c: *mut f32, n: usize) {
if a.is_null() || c.is_null() || n == 0 {
return;
}
let a = unsafe { std::slice::from_raw_parts(a, n) };
let c = unsafe { std::slice::from_raw_parts_mut(c, n) };
let mut guard = match negate_f32_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return,
};
guard.compute(a, c)
}
fn clamp_f64_dispatch() -> &'static Mutex<Dispatch<ClampKernel>> {
static D: OnceLock<Mutex<Dispatch<ClampKernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("clamp_f64_ffi", vec![
KernelInfo { name: "scalar", func: kernels::clamp_f64_scalar, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::clamp_f64_sse as ClampKernel, kernels::clamp_f64_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::clamp_f64_avx2 as ClampKernel, kernels::clamp_f64_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::clamp_f64_neon as ClampKernel, kernels::clamp_f64_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_clamp(a: *const f64, lo: f64, hi: f64, c: *mut f64, n: usize) {
if a.is_null() || c.is_null() || n == 0 {
return;
}
let a = unsafe { std::slice::from_raw_parts(a, n) };
let c = unsafe { std::slice::from_raw_parts_mut(c, n) };
let mut guard = match clamp_f64_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return,
};
guard.compute(a, lo, hi, c)
}
fn clamp_f32_dispatch() -> &'static Mutex<Dispatch<ClampF32Kernel>> {
static D: OnceLock<Mutex<Dispatch<ClampF32Kernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("clamp_f32_ffi", vec![
KernelInfo { name: "scalar", func: kernels::clamp_f32_scalar, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::clamp_f32_sse as ClampF32Kernel, kernels::clamp_f32_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::clamp_f32_avx2 as ClampF32Kernel, kernels::clamp_f32_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::clamp_f32_neon as ClampF32Kernel, kernels::clamp_f32_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_clamp_f32(a: *const f32, lo: f32, hi: f32, c: *mut f32, n: usize) {
if a.is_null() || c.is_null() || n == 0 {
return;
}
let a = unsafe { std::slice::from_raw_parts(a, n) };
let c = unsafe { std::slice::from_raw_parts_mut(c, n) };
let mut guard = match clamp_f32_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return,
};
guard.compute(a, lo, hi, c)
}
fn matmul_f32_dispatch() -> &'static Mutex<Dispatch<MatMulF32Kernel>> {
static D: OnceLock<Mutex<Dispatch<MatMulF32Kernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("matmul_f32_ffi", vec![
KernelInfo { name: "scalar", func: kernels::matmul_f32_scalar, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
KernelInfo { name: "tiled", func: kernels::matmul_f32_tiled, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
KernelInfo { name: "cache_tiled", func: kernels::matmul_f32_cache_tiled, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::matmul_f32_sse as MatMulF32Kernel, kernels::matmul_f32_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::matmul_f32_avx2 as MatMulF32Kernel, kernels::matmul_f32_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::matmul_f32_neon as MatMulF32Kernel, kernels::matmul_f32_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_matmul_f32(a: *const f32, b: *const f32, c: *mut f32, n: usize) {
if a.is_null() || b.is_null() || c.is_null() || n == 0 {
return;
}
let a = unsafe { std::slice::from_raw_parts(a, n * n) };
let b = unsafe { std::slice::from_raw_parts(b, n * n) };
let c = unsafe { std::slice::from_raw_parts_mut(c, n * n) };
c.fill(0.0);
let mut guard = match matmul_f32_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return,
};
guard.compute(a, b, c, n)
}
fn matmul_bias_relu_dispatch() -> &'static Mutex<Dispatch<MatMulBiasReluKernel>> {
static D: OnceLock<Mutex<Dispatch<MatMulBiasReluKernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("matmul_bias_relu_ffi", vec![
KernelInfo { name: "scalar", func: kernels::matmul_bias_relu_scalar, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::matmul_bias_relu_sse as MatMulBiasReluKernel, kernels::matmul_bias_relu_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::matmul_bias_relu_avx2 as MatMulBiasReluKernel, kernels::matmul_bias_relu_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::matmul_bias_relu_neon as MatMulBiasReluKernel, kernels::matmul_bias_relu_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_matmul_bias_relu(
a: *const f64,
w: *const f64,
bias: *const f64,
c: *mut f64,
m: usize,
k: usize,
n: usize,
) {
if a.is_null() || w.is_null() || bias.is_null() || c.is_null() || m == 0 || k == 0 || n == 0 {
return;
}
let a = unsafe { std::slice::from_raw_parts(a, m * k) };
let w = unsafe { std::slice::from_raw_parts(w, k * n) };
let bias = unsafe { std::slice::from_raw_parts(bias, n) };
let c = unsafe { std::slice::from_raw_parts_mut(c, m * n) };
c.fill(0.0);
let mut guard = match matmul_bias_relu_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return,
};
guard.compute(a, w, bias, c, m, k, n)
}
fn memchr_dispatch() -> &'static Mutex<Dispatch<MemchrKernel>> {
static D: OnceLock<Mutex<Dispatch<MemchrKernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("memchr_ffi", vec![
KernelInfo { name: "scalar", func: kernels::memchr_scalar, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::memchr_sse as MemchrKernel, kernels::memchr_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::memchr_avx2 as MemchrKernel, kernels::memchr_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::memchr_neon as MemchrKernel, kernels::memchr_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_memchr(byte: u8, data: *const u8, n: usize) -> isize {
if data.is_null() || n == 0 {
return -1;
}
let data = unsafe { std::slice::from_raw_parts(data, n) };
let mut guard = match memchr_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return -1,
};
match guard.compute(byte, data) {
Some(idx) => idx as isize,
None => -1,
}
}
fn dot_clamp_reduce_dispatch() -> &'static Mutex<Dispatch<DotKernel>> {
static D: OnceLock<Mutex<Dispatch<DotKernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("dot_clamp_reduce_ffi", vec![
KernelInfo { name: "scalar", func: kernels::dot_clamp_reduce_scalar as DotKernel, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::dot_clamp_reduce_sse as DotKernel, kernels::dot_clamp_reduce_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::dot_clamp_reduce_avx2 as DotKernel, kernels::dot_clamp_reduce_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::dot_clamp_reduce_neon as DotKernel, kernels::dot_clamp_reduce_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_fused_dot_clamp_reduce(a: *const f64, b: *const f64, n: usize) -> f64 {
if a.is_null() || b.is_null() || n == 0 {
return 0.0;
}
let a = unsafe { std::slice::from_raw_parts(a, n) };
let b = unsafe { std::slice::from_raw_parts(b, n) };
let mut guard = match dot_clamp_reduce_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return 0.0,
};
guard.compute(a, b)
}
fn sort_f64_dispatch() -> &'static Mutex<Dispatch<SortF64Kernel>> {
static D: OnceLock<Mutex<Dispatch<SortF64Kernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("sort_f64_ffi", vec![
KernelInfo { name: "scalar", func: kernels::sort_f64_scalar, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::sort_f64_sse as SortF64Kernel, kernels::sort_f64_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::sort_f64_avx2 as SortF64Kernel, kernels::sort_f64_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::sort_f64_neon as SortF64Kernel, kernels::sort_f64_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_sort_f64(a: *const f64, out: *mut f64, n: usize) {
if a.is_null() || out.is_null() || n == 0 {
return;
}
let a = unsafe { std::slice::from_raw_parts(a, n) };
let out = unsafe { std::slice::from_raw_parts_mut(out, n) };
out.copy_from_slice(a);
let mut guard = match sort_f64_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return,
};
guard.compute(out)
}
fn gemv_f64_dispatch() -> &'static Mutex<Dispatch<GemvF64Kernel>> {
static D: OnceLock<Mutex<Dispatch<GemvF64Kernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("gemv_f64_ffi", vec![
KernelInfo { name: "scalar", func: kernels::gemv_f64_scalar, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::gemv_f64_sse as GemvF64Kernel, kernels::gemv_f64_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::gemv_f64_avx2 as GemvF64Kernel, kernels::gemv_f64_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::gemv_f64_neon as GemvF64Kernel, kernels::gemv_f64_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_gemv_f64(
alpha: f64,
a: *const f64,
x: *const f64,
beta: f64,
y: *mut f64,
rows: usize,
cols: usize,
lda: usize,
) {
if a.is_null() || x.is_null() || y.is_null() || rows == 0 || cols == 0 || lda == 0 {
return;
}
let a = unsafe { std::slice::from_raw_parts(a, rows * lda) };
let x = unsafe { std::slice::from_raw_parts(x, cols) };
let y = unsafe { std::slice::from_raw_parts_mut(y, rows) };
let mut guard = match gemv_f64_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return,
};
guard.compute(alpha, a, x, beta, y, rows, cols)
}
fn conv1d_f64_dispatch() -> &'static Mutex<Dispatch<BinaryOpKernel>> {
static D: OnceLock<Mutex<Dispatch<BinaryOpKernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("conv1d_f64_ffi", vec![
KernelInfo { name: "scalar", func: kernels::conv1d_f64_scalar as BinaryOpKernel, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::conv1d_f64_sse as BinaryOpKernel, kernels::conv1d_f64_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::conv1d_f64_avx2 as BinaryOpKernel, kernels::conv1d_f64_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::conv1d_f64_neon as BinaryOpKernel, kernels::conv1d_f64_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_conv1d_f64(
signal: *const f64,
kernel: *const f64,
out: *mut f64,
signal_len: usize,
kernel_len: usize,
) {
if signal.is_null() || kernel.is_null() || out.is_null() || signal_len == 0 || kernel_len == 0 {
return;
}
let signal = unsafe { std::slice::from_raw_parts(signal, signal_len) };
let kernel = unsafe { std::slice::from_raw_parts(kernel, kernel_len) };
let out_len = signal_len.saturating_sub(kernel_len) + 1;
let out = unsafe { std::slice::from_raw_parts_mut(out, out_len) };
let mut guard = match conv1d_f64_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return,
};
guard.compute(signal, kernel, out)
}
fn layernorm_f64_dispatch() -> &'static Mutex<Dispatch<LayerNormF64Kernel>> {
static D: OnceLock<Mutex<Dispatch<LayerNormF64Kernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("layernorm_f64_ffi", vec![
KernelInfo { name: "scalar", func: kernels::layer_norm_f64_scalar, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::layer_norm_f64_sse as LayerNormF64Kernel, kernels::layer_norm_f64_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::layer_norm_f64_avx2 as LayerNormF64Kernel, kernels::layer_norm_f64_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::layer_norm_f64_neon as LayerNormF64Kernel, kernels::layer_norm_f64_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_layernorm_f64(
input: *const f64,
gamma: *const f64,
beta: *const f64,
epsilon: f64,
output: *mut f64,
n: usize,
) {
if input.is_null() || gamma.is_null() || beta.is_null() || output.is_null() || n == 0 {
return;
}
let input = unsafe { std::slice::from_raw_parts(input, n) };
let gamma = unsafe { std::slice::from_raw_parts(gamma, n) };
let beta = unsafe { std::slice::from_raw_parts(beta, n) };
let output = unsafe { std::slice::from_raw_parts_mut(output, n) };
let mut guard = match layernorm_f64_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return,
};
guard.compute(input, gamma, beta, epsilon, output)
}
fn cumsum_f64_dispatch() -> &'static Mutex<Dispatch<SoftmaxKernel>> {
static D: OnceLock<Mutex<Dispatch<SoftmaxKernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("cumsum_f64_ffi", vec![
KernelInfo { name: "scalar", func: kernels::cumsum_f64_scalar as SoftmaxKernel, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::cumsum_f64_sse as SoftmaxKernel, kernels::cumsum_f64_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::cumsum_f64_avx2 as SoftmaxKernel, kernels::cumsum_f64_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::cumsum_f64_neon as SoftmaxKernel, kernels::cumsum_f64_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_cumsum_f64(input: *const f64, output: *mut f64, n: usize) {
if input.is_null() || output.is_null() || n == 0 {
return;
}
let input = unsafe { std::slice::from_raw_parts(input, n) };
let output = unsafe { std::slice::from_raw_parts_mut(output, n) };
let mut guard = match cumsum_f64_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return,
};
guard.compute(input, output)
}
fn sin_batch_f64_dispatch() -> &'static Mutex<Dispatch<SoftmaxKernel>> {
static D: OnceLock<Mutex<Dispatch<SoftmaxKernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("sin_batch_f64_ffi", vec![
KernelInfo { name: "scalar", func: kernels::sin_batch_f64_scalar as SoftmaxKernel, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::sin_batch_f64_sse as SoftmaxKernel, kernels::sin_batch_f64_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::sin_batch_f64_avx2 as SoftmaxKernel, kernels::sin_batch_f64_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::sin_batch_f64_neon as SoftmaxKernel, kernels::sin_batch_f64_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_sin_batch_f64(input: *const f64, output: *mut f64, n: usize) {
if input.is_null() || output.is_null() || n == 0 {
return;
}
let input = unsafe { std::slice::from_raw_parts(input, n) };
let output = unsafe { std::slice::from_raw_parts_mut(output, n) };
let mut guard = match sin_batch_f64_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return,
};
guard.compute(input, output)
}
fn cos_batch_f64_dispatch() -> &'static Mutex<Dispatch<SoftmaxKernel>> {
static D: OnceLock<Mutex<Dispatch<SoftmaxKernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("cos_batch_f64_ffi", vec![
KernelInfo { name: "scalar", func: kernels::cos_batch_f64_scalar as SoftmaxKernel, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::cos_batch_f64_sse as SoftmaxKernel, kernels::cos_batch_f64_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::cos_batch_f64_avx2 as SoftmaxKernel, kernels::cos_batch_f64_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::cos_batch_f64_neon as SoftmaxKernel, kernels::cos_batch_f64_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_cos_batch_f64(input: *const f64, output: *mut f64, n: usize) {
if input.is_null() || output.is_null() || n == 0 {
return;
}
let input = unsafe { std::slice::from_raw_parts(input, n) };
let output = unsafe { std::slice::from_raw_parts_mut(output, n) };
let mut guard = match cos_batch_f64_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return,
};
guard.compute(input, output)
}
fn tan_batch_f64_dispatch() -> &'static Mutex<Dispatch<SoftmaxKernel>> {
static D: OnceLock<Mutex<Dispatch<SoftmaxKernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("tan_batch_f64_ffi", vec![
KernelInfo { name: "scalar", func: kernels::tan_batch_f64_scalar as SoftmaxKernel, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::tan_batch_f64_sse as SoftmaxKernel, kernels::tan_batch_f64_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::tan_batch_f64_avx2 as SoftmaxKernel, kernels::tan_batch_f64_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::tan_batch_f64_neon as SoftmaxKernel, kernels::tan_batch_f64_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_tan_batch_f64(input: *const f64, output: *mut f64, n: usize) {
if input.is_null() || output.is_null() || n == 0 {
return;
}
let input = unsafe { std::slice::from_raw_parts(input, n) };
let output = unsafe { std::slice::from_raw_parts_mut(output, n) };
let mut guard = match tan_batch_f64_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return,
};
guard.compute(input, output)
}
fn sort_f32_dispatch() -> &'static Mutex<Dispatch<SortF32Kernel>> {
static D: OnceLock<Mutex<Dispatch<SortF32Kernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("sort_f32_ffi", vec![
KernelInfo { name: "scalar", func: kernels::sort_f32_scalar, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::sort_f32_sse as SortF32Kernel, kernels::sort_f32_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::sort_f32_avx2 as SortF32Kernel, kernels::sort_f32_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::sort_f32_neon as SortF32Kernel, kernels::sort_f32_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_sort_f32(a: *const f32, out: *mut f32, n: usize) {
if a.is_null() || out.is_null() || n == 0 {
return;
}
let a = unsafe { std::slice::from_raw_parts(a, n) };
let out = unsafe { std::slice::from_raw_parts_mut(out, n) };
out.copy_from_slice(a);
let mut guard = match sort_f32_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return,
};
guard.compute(out)
}
fn gemv_f32_dispatch() -> &'static Mutex<Dispatch<GemvF32Kernel>> {
static D: OnceLock<Mutex<Dispatch<GemvF32Kernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("gemv_f32_ffi", vec![
KernelInfo { name: "scalar", func: kernels::gemv_f32_scalar, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::gemv_f32_sse as GemvF32Kernel, kernels::gemv_f32_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::gemv_f32_avx2 as GemvF32Kernel, kernels::gemv_f32_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::gemv_f32_neon as GemvF32Kernel, kernels::gemv_f32_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_gemv_f32(
alpha: f32,
a: *const f32,
x: *const f32,
beta: f32,
y: *mut f32,
rows: usize,
cols: usize,
lda: usize,
) {
if a.is_null() || x.is_null() || y.is_null() || rows == 0 || cols == 0 || lda == 0 {
return;
}
let a = unsafe { std::slice::from_raw_parts(a, rows * lda) };
let x = unsafe { std::slice::from_raw_parts(x, cols) };
let y = unsafe { std::slice::from_raw_parts_mut(y, rows) };
let mut guard = match gemv_f32_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return,
};
guard.compute(alpha, a, x, beta, y, rows, cols)
}
fn conv1d_f32_dispatch() -> &'static Mutex<Dispatch<AddF32Kernel>> {
static D: OnceLock<Mutex<Dispatch<AddF32Kernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("conv1d_f32_ffi", vec![
KernelInfo { name: "scalar", func: kernels::conv1d_f32_scalar as AddF32Kernel, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::conv1d_f32_sse as AddF32Kernel, kernels::conv1d_f32_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::conv1d_f32_avx2 as AddF32Kernel, kernels::conv1d_f32_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::conv1d_f32_neon as AddF32Kernel, kernels::conv1d_f32_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_conv1d_f32(
signal: *const f32,
kernel: *const f32,
out: *mut f32,
signal_len: usize,
kernel_len: usize,
) {
if signal.is_null() || kernel.is_null() || out.is_null() || signal_len == 0 || kernel_len == 0 {
return;
}
let signal = unsafe { std::slice::from_raw_parts(signal, signal_len) };
let kernel = unsafe { std::slice::from_raw_parts(kernel, kernel_len) };
let out_len = signal_len.saturating_sub(kernel_len) + 1;
let out = unsafe { std::slice::from_raw_parts_mut(out, out_len) };
let mut guard = match conv1d_f32_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return,
};
guard.compute(signal, kernel, out)
}
fn layernorm_f32_dispatch() -> &'static Mutex<Dispatch<LayerNormF32Kernel>> {
static D: OnceLock<Mutex<Dispatch<LayerNormF32Kernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("layernorm_f32_ffi", vec![
KernelInfo { name: "scalar", func: kernels::layer_norm_f32_scalar, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::layer_norm_f32_sse as LayerNormF32Kernel, kernels::layer_norm_f32_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::layer_norm_f32_avx2 as LayerNormF32Kernel, kernels::layer_norm_f32_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::layer_norm_f32_neon as LayerNormF32Kernel, kernels::layer_norm_f32_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_layernorm_f32(
input: *const f32,
gamma: *const f32,
beta: *const f32,
epsilon: f32,
output: *mut f32,
n: usize,
) {
if input.is_null() || gamma.is_null() || beta.is_null() || output.is_null() || n == 0 {
return;
}
let input = unsafe { std::slice::from_raw_parts(input, n) };
let gamma = unsafe { std::slice::from_raw_parts(gamma, n) };
let beta = unsafe { std::slice::from_raw_parts(beta, n) };
let output = unsafe { std::slice::from_raw_parts_mut(output, n) };
let mut guard = match layernorm_f32_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return,
};
guard.compute(input, gamma, beta, epsilon, output)
}
fn cumsum_f32_dispatch() -> &'static Mutex<Dispatch<UnaryOpF32Kernel>> {
static D: OnceLock<Mutex<Dispatch<UnaryOpF32Kernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("cumsum_f32_ffi", vec![
KernelInfo { name: "scalar", func: kernels::cumsum_f32_scalar as UnaryOpF32Kernel, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::cumsum_f32_sse as UnaryOpF32Kernel, kernels::cumsum_f32_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::cumsum_f32_avx2 as UnaryOpF32Kernel, kernels::cumsum_f32_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::cumsum_f32_neon as UnaryOpF32Kernel, kernels::cumsum_f32_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_cumsum_f32(input: *const f32, output: *mut f32, n: usize) {
if input.is_null() || output.is_null() || n == 0 {
return;
}
let input = unsafe { std::slice::from_raw_parts(input, n) };
let output = unsafe { std::slice::from_raw_parts_mut(output, n) };
let mut guard = match cumsum_f32_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return,
};
guard.compute(input, output)
}
fn sin_batch_f32_dispatch() -> &'static Mutex<Dispatch<UnaryOpF32Kernel>> {
static D: OnceLock<Mutex<Dispatch<UnaryOpF32Kernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("sin_batch_f32_ffi", vec![
KernelInfo { name: "scalar", func: kernels::sin_batch_f32_scalar as UnaryOpF32Kernel, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::sin_batch_f32_sse as UnaryOpF32Kernel, kernels::sin_batch_f32_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::sin_batch_f32_avx2 as UnaryOpF32Kernel, kernels::sin_batch_f32_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::sin_batch_f32_neon as UnaryOpF32Kernel, kernels::sin_batch_f32_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_sin_batch_f32(input: *const f32, output: *mut f32, n: usize) {
if input.is_null() || output.is_null() || n == 0 {
return;
}
let input = unsafe { std::slice::from_raw_parts(input, n) };
let output = unsafe { std::slice::from_raw_parts_mut(output, n) };
let mut guard = match sin_batch_f32_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return,
};
guard.compute(input, output)
}
fn cos_batch_f32_dispatch() -> &'static Mutex<Dispatch<UnaryOpF32Kernel>> {
static D: OnceLock<Mutex<Dispatch<UnaryOpF32Kernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("cos_batch_f32_ffi", vec![
KernelInfo { name: "scalar", func: kernels::cos_batch_f32_scalar as UnaryOpF32Kernel, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::cos_batch_f32_sse as UnaryOpF32Kernel, kernels::cos_batch_f32_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::cos_batch_f32_avx2 as UnaryOpF32Kernel, kernels::cos_batch_f32_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::cos_batch_f32_neon as UnaryOpF32Kernel, kernels::cos_batch_f32_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_cos_batch_f32(input: *const f32, output: *mut f32, n: usize) {
if input.is_null() || output.is_null() || n == 0 {
return;
}
let input = unsafe { std::slice::from_raw_parts(input, n) };
let output = unsafe { std::slice::from_raw_parts_mut(output, n) };
let mut guard = match cos_batch_f32_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return,
};
guard.compute(input, output)
}
fn tan_batch_f32_dispatch() -> &'static Mutex<Dispatch<UnaryOpF32Kernel>> {
static D: OnceLock<Mutex<Dispatch<UnaryOpF32Kernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("tan_batch_f32_ffi", vec![
KernelInfo { name: "scalar", func: kernels::tan_batch_f32_scalar as UnaryOpF32Kernel, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
#[cfg(target_arch = "x86_64")]
add_sse_kernel("SSE", kernels::tan_batch_f32_sse as UnaryOpF32Kernel, kernels::tan_batch_f32_sse_supported, 2),
#[cfg(target_arch = "x86_64")]
add_avx2_kernel("AVX2", kernels::tan_batch_f32_avx2 as UnaryOpF32Kernel, kernels::tan_batch_f32_avx2_supported, 3),
#[cfg(target_arch = "aarch64")]
add_neon_kernel("NEON", kernels::tan_batch_f32_neon as UnaryOpF32Kernel, kernels::tan_batch_f32_neon_supported, 2),
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_tan_batch_f32(input: *const f32, output: *mut f32, n: usize) {
if input.is_null() || output.is_null() || n == 0 {
return;
}
let input = unsafe { std::slice::from_raw_parts(input, n) };
let output = unsafe { std::slice::from_raw_parts_mut(output, n) };
let mut guard = match tan_batch_f32_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return,
};
guard.compute(input, output)
}
fn fused_dot_clamp_reduce_f32_dispatch() -> &'static Mutex<Dispatch<DotF32Kernel>> {
static D: OnceLock<Mutex<Dispatch<DotF32Kernel>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new("fused_dot_clamp_reduce_f32_ffi", vec![
KernelInfo { name: "scalar", func: kernels::dot_f32_scalar as DotF32Kernel, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
]))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_fused_dot_clamp_reduce_f32(a: *const f32, b: *const f32, n: usize) -> f32 {
if a.is_null() || b.is_null() || n == 0 {
return 0.0;
}
let a = unsafe { std::slice::from_raw_parts(a, n) };
let b = unsafe { std::slice::from_raw_parts(b, n) };
let mut guard = match fused_dot_clamp_reduce_f32_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return 0.0,
};
guard.compute(a, b)
}
macro_rules! make_int_dispatch {
($name:ident, $kt:ty, $scalar:expr, $neon:expr, $neon_supported:expr) => {
fn $name() -> &'static Mutex<Dispatch<$kt>> {
static D: OnceLock<Mutex<Dispatch<$kt>>> = OnceLock::new();
D.get_or_init(|| {
let mut k: Vec<KernelInfo<$kt>> = vec![
KernelInfo { name: "scalar", func: $scalar, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
];
#[cfg(target_arch = "aarch64")]
k.push(add_neon_kernel("NEON", $neon as $kt, $neon_supported, 2));
Mutex::new(Dispatch::new(stringify!($name), k))
})
}
};
($name:ident, $kt:ty, $scalar:expr) => {
fn $name() -> &'static Mutex<Dispatch<$kt>> {
static D: OnceLock<Mutex<Dispatch<$kt>>> = OnceLock::new();
D.get_or_init(|| {
Mutex::new(Dispatch::new(stringify!($name), vec![
KernelInfo { name: "scalar", func: $scalar, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
]))
})
}
};
}
make_int_dispatch!(dot_i32_dispatch, DotI32Kernel, kernels::dot_i32_scalar, kernels::dot_i32_neon, kernels::dot_i32_neon_supported);
make_int_dispatch!(dot_i64_dispatch, DotI64Kernel, kernels::dot_i64_scalar, kernels::dot_i64_neon, kernels::dot_i64_neon_supported);
make_int_dispatch!(dot_i16_dispatch, DotI16Kernel, kernels::dot_i16_scalar, kernels::dot_i16_neon, kernels::dot_i16_neon_supported);
make_int_dispatch!(dot_i8_dispatch, DotI8Kernel, kernels::dot_i8_scalar, kernels::dot_i8_neon, kernels::dot_i8_neon_supported);
#[no_mangle] pub unsafe extern "C" fn himada_dot_i32(a: *const i32, b: *const i32, n: usize) -> i32 {
if a.is_null() || b.is_null() || n == 0 { return 0; }
let a = unsafe { std::slice::from_raw_parts(a, n) };
let b = unsafe { std::slice::from_raw_parts(b, n) };
let mut guard = match dot_i32_dispatch().lock() { Ok(l) => l, Err(_) => return 0 };
guard.compute(a, b)
}
#[no_mangle] pub unsafe extern "C" fn himada_dot_i64(a: *const i64, b: *const i64, n: usize) -> i64 {
if a.is_null() || b.is_null() || n == 0 { return 0; }
let a = unsafe { std::slice::from_raw_parts(a, n) };
let b = unsafe { std::slice::from_raw_parts(b, n) };
let mut guard = match dot_i64_dispatch().lock() { Ok(l) => l, Err(_) => return 0 };
guard.compute(a, b)
}
#[no_mangle] pub unsafe extern "C" fn himada_dot_i16(a: *const i16, b: *const i16, n: usize) -> i16 {
if a.is_null() || b.is_null() || n == 0 { return 0; }
let a = unsafe { std::slice::from_raw_parts(a, n) };
let b = unsafe { std::slice::from_raw_parts(b, n) };
let mut guard = match dot_i16_dispatch().lock() { Ok(l) => l, Err(_) => return 0 };
guard.compute(a, b)
}
#[no_mangle] pub unsafe extern "C" fn himada_dot_i8(a: *const i8, b: *const i8, n: usize) -> i8 {
if a.is_null() || b.is_null() || n == 0 { return 0; }
let a = unsafe { std::slice::from_raw_parts(a, n) };
let b = unsafe { std::slice::from_raw_parts(b, n) };
let mut guard = match dot_i8_dispatch().lock() { Ok(l) => l, Err(_) => return 0 };
guard.compute(a, b)
}
make_int_dispatch!(reduce_sum_i32_dispatch, ReduceI32Kernel, kernels::reduce_sum_i32_scalar, kernels::reduce_sum_i32_neon, kernels::reduce_sum_i32_neon_supported);
make_int_dispatch!(reduce_sum_i64_dispatch, ReduceI64Kernel, kernels::reduce_sum_i64_scalar, kernels::reduce_sum_i64_neon, kernels::reduce_sum_i64_neon_supported);
make_int_dispatch!(reduce_sum_i16_dispatch, ReduceI16Kernel, kernels::reduce_sum_i16_scalar, kernels::reduce_sum_i16_neon, kernels::reduce_sum_i16_neon_supported);
make_int_dispatch!(reduce_sum_i8_dispatch, ReduceI8Kernel, kernels::reduce_sum_i8_scalar, kernels::reduce_sum_i8_neon, kernels::reduce_sum_i8_neon_supported);
#[no_mangle] pub unsafe extern "C" fn himada_reduce_sum_i32(a: *const i32, n: usize) -> i32 {
if a.is_null() || n == 0 { return 0; }
let a = unsafe { std::slice::from_raw_parts(a, n) };
let mut guard = match reduce_sum_i32_dispatch().lock() { Ok(l) => l, Err(_) => return 0 };
guard.compute(a)
}
#[no_mangle] pub unsafe extern "C" fn himada_reduce_sum_i64(a: *const i64, n: usize) -> i64 {
if a.is_null() || n == 0 { return 0; }
let a = unsafe { std::slice::from_raw_parts(a, n) };
let mut guard = match reduce_sum_i64_dispatch().lock() { Ok(l) => l, Err(_) => return 0 };
guard.compute(a)
}
#[no_mangle] pub unsafe extern "C" fn himada_reduce_sum_i16(a: *const i16, n: usize) -> i16 {
if a.is_null() || n == 0 { return 0; }
let a = unsafe { std::slice::from_raw_parts(a, n) };
let mut guard = match reduce_sum_i16_dispatch().lock() { Ok(l) => l, Err(_) => return 0 };
guard.compute(a)
}
#[no_mangle] pub unsafe extern "C" fn himada_reduce_sum_i8(a: *const i8, n: usize) -> i8 {
if a.is_null() || n == 0 { return 0; }
let a = unsafe { std::slice::from_raw_parts(a, n) };
let mut guard = match reduce_sum_i8_dispatch().lock() { Ok(l) => l, Err(_) => return 0 };
guard.compute(a)
}
make_int_dispatch!(reduce_max_i32_dispatch, ReduceI32Kernel, kernels::reduce_max_i32_scalar as ReduceI32Kernel, kernels::reduce_max_i32_neon as ReduceI32Kernel, kernels::reduce_max_i32_neon_supported);
make_int_dispatch!(reduce_max_i64_dispatch, ReduceI64Kernel, kernels::reduce_max_i64_scalar as ReduceI64Kernel, kernels::reduce_max_i64_neon as ReduceI64Kernel, kernels::reduce_max_i64_neon_supported);
make_int_dispatch!(reduce_max_i16_dispatch, ReduceI16Kernel, kernels::reduce_max_i16_scalar as ReduceI16Kernel, kernels::reduce_max_i16_neon as ReduceI16Kernel, kernels::reduce_max_i16_neon_supported);
make_int_dispatch!(reduce_max_i8_dispatch, ReduceI8Kernel, kernels::reduce_max_i8_scalar as ReduceI8Kernel, kernels::reduce_max_i8_neon as ReduceI8Kernel, kernels::reduce_max_i8_neon_supported);
#[no_mangle] pub unsafe extern "C" fn himada_reduce_max_i32(a: *const i32, n: usize) -> i32 {
if a.is_null() || n == 0 { return 0; }
let a = unsafe { std::slice::from_raw_parts(a, n) };
let mut guard = match reduce_max_i32_dispatch().lock() { Ok(l) => l, Err(_) => return 0 };
guard.compute(a)
}
#[no_mangle] pub unsafe extern "C" fn himada_reduce_max_i64(a: *const i64, n: usize) -> i64 {
if a.is_null() || n == 0 { return 0; }
let a = unsafe { std::slice::from_raw_parts(a, n) };
let mut guard = match reduce_max_i64_dispatch().lock() { Ok(l) => l, Err(_) => return 0 };
guard.compute(a)
}
#[no_mangle] pub unsafe extern "C" fn himada_reduce_max_i16(a: *const i16, n: usize) -> i16 {
if a.is_null() || n == 0 { return 0; }
let a = unsafe { std::slice::from_raw_parts(a, n) };
let mut guard = match reduce_max_i16_dispatch().lock() { Ok(l) => l, Err(_) => return 0 };
guard.compute(a)
}
#[no_mangle] pub unsafe extern "C" fn himada_reduce_max_i8(a: *const i8, n: usize) -> i8 {
if a.is_null() || n == 0 { return 0; }
let a = unsafe { std::slice::from_raw_parts(a, n) };
let mut guard = match reduce_max_i8_dispatch().lock() { Ok(l) => l, Err(_) => return 0 };
guard.compute(a)
}
make_int_dispatch!(abs_max_i32_dispatch, ReduceI32Kernel, kernels::abs_max_i32_scalar as ReduceI32Kernel, kernels::abs_max_i32_neon as ReduceI32Kernel, kernels::abs_max_i32_neon_supported);
make_int_dispatch!(abs_max_i64_dispatch, ReduceI64Kernel, kernels::abs_max_i64_scalar as ReduceI64Kernel, kernels::abs_max_i64_neon as ReduceI64Kernel, kernels::abs_max_i64_neon_supported);
make_int_dispatch!(abs_max_i16_dispatch, ReduceI16Kernel, kernels::abs_max_i16_scalar as ReduceI16Kernel, kernels::abs_max_i16_neon as ReduceI16Kernel, kernels::abs_max_i16_neon_supported);
make_int_dispatch!(abs_max_i8_dispatch, ReduceI8Kernel, kernels::abs_max_i8_scalar as ReduceI8Kernel, kernels::abs_max_i8_neon as ReduceI8Kernel, kernels::abs_max_i8_neon_supported);
#[no_mangle] pub unsafe extern "C" fn himada_abs_max_i32(a: *const i32, n: usize) -> i32 {
if a.is_null() || n == 0 { return 0; }
let a = unsafe { std::slice::from_raw_parts(a, n) };
let mut guard = match abs_max_i32_dispatch().lock() { Ok(l) => l, Err(_) => return 0 };
guard.compute(a)
}
#[no_mangle] pub unsafe extern "C" fn himada_abs_max_i64(a: *const i64, n: usize) -> i64 {
if a.is_null() || n == 0 { return 0; }
let a = unsafe { std::slice::from_raw_parts(a, n) };
let mut guard = match abs_max_i64_dispatch().lock() { Ok(l) => l, Err(_) => return 0 };
guard.compute(a)
}
#[no_mangle] pub unsafe extern "C" fn himada_abs_max_i16(a: *const i16, n: usize) -> i16 {
if a.is_null() || n == 0 { return 0; }
let a = unsafe { std::slice::from_raw_parts(a, n) };
let mut guard = match abs_max_i16_dispatch().lock() { Ok(l) => l, Err(_) => return 0 };
guard.compute(a)
}
#[no_mangle] pub unsafe extern "C" fn himada_abs_max_i8(a: *const i8, n: usize) -> i8 {
if a.is_null() || n == 0 { return 0; }
let a = unsafe { std::slice::from_raw_parts(a, n) };
let mut guard = match abs_max_i8_dispatch().lock() { Ok(l) => l, Err(_) => return 0 };
guard.compute(a)
}
make_int_dispatch!(argmax_i32_dispatch, ArgmaxI32Kernel, kernels::argmax_i32_scalar, kernels::argmax_i32_neon, kernels::argmax_i32_neon_supported);
make_int_dispatch!(argmax_i64_dispatch, ArgmaxI64Kernel, kernels::argmax_i64_scalar, kernels::argmax_i64_neon, kernels::argmax_i64_neon_supported);
make_int_dispatch!(argmax_i16_dispatch, ArgmaxI16Kernel, kernels::argmax_i16_scalar, kernels::argmax_i16_neon, kernels::argmax_i16_neon_supported);
make_int_dispatch!(argmax_i8_dispatch, ArgmaxI8Kernel, kernels::argmax_i8_scalar, kernels::argmax_i8_neon, kernels::argmax_i8_neon_supported);
#[no_mangle] pub unsafe extern "C" fn himada_argmax_i32(a: *const i32, n: usize) -> usize {
if a.is_null() || n == 0 { return 0; }
let a = unsafe { std::slice::from_raw_parts(a, n) };
let mut guard = match argmax_i32_dispatch().lock() { Ok(l) => l, Err(_) => return 0 };
guard.compute(a)
}
#[no_mangle] pub unsafe extern "C" fn himada_argmax_i64(a: *const i64, n: usize) -> usize {
if a.is_null() || n == 0 { return 0; }
let a = unsafe { std::slice::from_raw_parts(a, n) };
let mut guard = match argmax_i64_dispatch().lock() { Ok(l) => l, Err(_) => return 0 };
guard.compute(a)
}
#[no_mangle] pub unsafe extern "C" fn himada_argmax_i16(a: *const i16, n: usize) -> usize {
if a.is_null() || n == 0 { return 0; }
let a = unsafe { std::slice::from_raw_parts(a, n) };
let mut guard = match argmax_i16_dispatch().lock() { Ok(l) => l, Err(_) => return 0 };
guard.compute(a)
}
#[no_mangle] pub unsafe extern "C" fn himada_argmax_i8(a: *const i8, n: usize) -> usize {
if a.is_null() || n == 0 { return 0; }
let a = unsafe { std::slice::from_raw_parts(a, n) };
let mut guard = match argmax_i8_dispatch().lock() { Ok(l) => l, Err(_) => return 0 };
guard.compute(a)
}
make_int_dispatch!(hadamard_product_i32_dispatch, BinaryI32Kernel, kernels::hadamard_product_i32_scalar, kernels::hadamard_product_i32_neon, kernels::hadamard_product_i32_neon_supported);
make_int_dispatch!(hadamard_product_i64_dispatch, BinaryI64Kernel, kernels::hadamard_product_i64_scalar, kernels::hadamard_product_i64_neon, kernels::hadamard_product_i64_neon_supported);
make_int_dispatch!(hadamard_product_i16_dispatch, BinaryI16Kernel, kernels::hadamard_product_i16_scalar, kernels::hadamard_product_i16_neon, kernels::hadamard_product_i16_neon_supported);
make_int_dispatch!(hadamard_product_i8_dispatch, BinaryI8Kernel, kernels::hadamard_product_i8_scalar, kernels::hadamard_product_i8_neon, kernels::hadamard_product_i8_neon_supported);
#[no_mangle] pub unsafe extern "C" fn himada_hadamard_product_i32(a: *const i32, b: *const i32, out: *mut i32, n: usize) {
if a.is_null() || b.is_null() || out.is_null() || n == 0 { return; }
let a = unsafe { std::slice::from_raw_parts(a, n) }; let b = unsafe { std::slice::from_raw_parts(b, n) }; let out = unsafe { std::slice::from_raw_parts_mut(out, n) };
let mut guard = match hadamard_product_i32_dispatch().lock() { Ok(l) => l, Err(_) => return };
guard.compute(a, b, out)
}
#[no_mangle] pub unsafe extern "C" fn himada_hadamard_product_i64(a: *const i64, b: *const i64, out: *mut i64, n: usize) {
if a.is_null() || b.is_null() || out.is_null() || n == 0 { return; }
let a = unsafe { std::slice::from_raw_parts(a, n) }; let b = unsafe { std::slice::from_raw_parts(b, n) }; let out = unsafe { std::slice::from_raw_parts_mut(out, n) };
let mut guard = match hadamard_product_i64_dispatch().lock() { Ok(l) => l, Err(_) => return };
guard.compute(a, b, out)
}
#[no_mangle] pub unsafe extern "C" fn himada_hadamard_product_i16(a: *const i16, b: *const i16, out: *mut i16, n: usize) {
if a.is_null() || b.is_null() || out.is_null() || n == 0 { return; }
let a = unsafe { std::slice::from_raw_parts(a, n) }; let b = unsafe { std::slice::from_raw_parts(b, n) }; let out = unsafe { std::slice::from_raw_parts_mut(out, n) };
let mut guard = match hadamard_product_i16_dispatch().lock() { Ok(l) => l, Err(_) => return };
guard.compute(a, b, out)
}
#[no_mangle] pub unsafe extern "C" fn himada_hadamard_product_i8(a: *const i8, b: *const i8, out: *mut i8, n: usize) {
if a.is_null() || b.is_null() || out.is_null() || n == 0 { return; }
let a = unsafe { std::slice::from_raw_parts(a, n) }; let b = unsafe { std::slice::from_raw_parts(b, n) }; let out = unsafe { std::slice::from_raw_parts_mut(out, n) };
let mut guard = match hadamard_product_i8_dispatch().lock() { Ok(l) => l, Err(_) => return };
guard.compute(a, b, out)
}
make_int_dispatch!(add_i32_dispatch, BinaryI32Kernel, kernels::add_i32_scalar, kernels::add_i32_neon, kernels::add_i32_neon_supported);
make_int_dispatch!(add_i64_dispatch, BinaryI64Kernel, kernels::add_i64_scalar, kernels::add_i64_neon, kernels::add_i64_neon_supported);
make_int_dispatch!(add_i16_dispatch, BinaryI16Kernel, kernels::add_i16_scalar, kernels::add_i16_neon, kernels::add_i16_neon_supported);
make_int_dispatch!(add_i8_dispatch, BinaryI8Kernel, kernels::add_i8_scalar, kernels::add_i8_neon, kernels::add_i8_neon_supported);
#[no_mangle] pub unsafe extern "C" fn himada_add_i32(a: *const i32, b: *const i32, out: *mut i32, n: usize) {
if a.is_null() || b.is_null() || out.is_null() || n == 0 { return; }
let a = unsafe { std::slice::from_raw_parts(a, n) }; let b = unsafe { std::slice::from_raw_parts(b, n) }; let out = unsafe { std::slice::from_raw_parts_mut(out, n) };
let mut guard = match add_i32_dispatch().lock() { Ok(l) => l, Err(_) => return };
guard.compute(a, b, out)
}
#[no_mangle] pub unsafe extern "C" fn himada_add_i64(a: *const i64, b: *const i64, out: *mut i64, n: usize) {
if a.is_null() || b.is_null() || out.is_null() || n == 0 { return; }
let a = unsafe { std::slice::from_raw_parts(a, n) }; let b = unsafe { std::slice::from_raw_parts(b, n) }; let out = unsafe { std::slice::from_raw_parts_mut(out, n) };
let mut guard = match add_i64_dispatch().lock() { Ok(l) => l, Err(_) => return };
guard.compute(a, b, out)
}
#[no_mangle] pub unsafe extern "C" fn himada_add_i16(a: *const i16, b: *const i16, out: *mut i16, n: usize) {
if a.is_null() || b.is_null() || out.is_null() || n == 0 { return; }
let a = unsafe { std::slice::from_raw_parts(a, n) }; let b = unsafe { std::slice::from_raw_parts(b, n) }; let out = unsafe { std::slice::from_raw_parts_mut(out, n) };
let mut guard = match add_i16_dispatch().lock() { Ok(l) => l, Err(_) => return };
guard.compute(a, b, out)
}
#[no_mangle] pub unsafe extern "C" fn himada_add_i8(a: *const i8, b: *const i8, out: *mut i8, n: usize) {
if a.is_null() || b.is_null() || out.is_null() || n == 0 { return; }
let a = unsafe { std::slice::from_raw_parts(a, n) }; let b = unsafe { std::slice::from_raw_parts(b, n) }; let out = unsafe { std::slice::from_raw_parts_mut(out, n) };
let mut guard = match add_i8_dispatch().lock() { Ok(l) => l, Err(_) => return };
guard.compute(a, b, out)
}
make_int_dispatch!(sub_i32_dispatch, BinaryI32Kernel, kernels::sub_i32_scalar, kernels::sub_i32_neon, kernels::sub_i32_neon_supported);
make_int_dispatch!(sub_i64_dispatch, BinaryI64Kernel, kernels::sub_i64_scalar, kernels::sub_i64_neon, kernels::sub_i64_neon_supported);
make_int_dispatch!(sub_i16_dispatch, BinaryI16Kernel, kernels::sub_i16_scalar, kernels::sub_i16_neon, kernels::sub_i16_neon_supported);
make_int_dispatch!(sub_i8_dispatch, BinaryI8Kernel, kernels::sub_i8_scalar, kernels::sub_i8_neon, kernels::sub_i8_neon_supported);
#[no_mangle] pub unsafe extern "C" fn himada_sub_i32(a: *const i32, b: *const i32, out: *mut i32, n: usize) {
if a.is_null() || b.is_null() || out.is_null() || n == 0 { return; }
let a = unsafe { std::slice::from_raw_parts(a, n) }; let b = unsafe { std::slice::from_raw_parts(b, n) }; let out = unsafe { std::slice::from_raw_parts_mut(out, n) };
let mut guard = match sub_i32_dispatch().lock() { Ok(l) => l, Err(_) => return };
guard.compute(a, b, out)
}
#[no_mangle] pub unsafe extern "C" fn himada_sub_i64(a: *const i64, b: *const i64, out: *mut i64, n: usize) {
if a.is_null() || b.is_null() || out.is_null() || n == 0 { return; }
let a = unsafe { std::slice::from_raw_parts(a, n) }; let b = unsafe { std::slice::from_raw_parts(b, n) }; let out = unsafe { std::slice::from_raw_parts_mut(out, n) };
let mut guard = match sub_i64_dispatch().lock() { Ok(l) => l, Err(_) => return };
guard.compute(a, b, out)
}
#[no_mangle] pub unsafe extern "C" fn himada_sub_i16(a: *const i16, b: *const i16, out: *mut i16, n: usize) {
if a.is_null() || b.is_null() || out.is_null() || n == 0 { return; }
let a = unsafe { std::slice::from_raw_parts(a, n) }; let b = unsafe { std::slice::from_raw_parts(b, n) }; let out = unsafe { std::slice::from_raw_parts_mut(out, n) };
let mut guard = match sub_i16_dispatch().lock() { Ok(l) => l, Err(_) => return };
guard.compute(a, b, out)
}
#[no_mangle] pub unsafe extern "C" fn himada_sub_i8(a: *const i8, b: *const i8, out: *mut i8, n: usize) {
if a.is_null() || b.is_null() || out.is_null() || n == 0 { return; }
let a = unsafe { std::slice::from_raw_parts(a, n) }; let b = unsafe { std::slice::from_raw_parts(b, n) }; let out = unsafe { std::slice::from_raw_parts_mut(out, n) };
let mut guard = match sub_i8_dispatch().lock() { Ok(l) => l, Err(_) => return };
guard.compute(a, b, out)
}
make_int_dispatch!(mul_i32_dispatch, BinaryI32Kernel, kernels::mul_i32_scalar, kernels::mul_i32_neon, kernels::mul_i32_neon_supported);
make_int_dispatch!(mul_i64_dispatch, BinaryI64Kernel, kernels::mul_i64_scalar, kernels::mul_i64_neon, kernels::mul_i64_neon_supported);
make_int_dispatch!(mul_i16_dispatch, BinaryI16Kernel, kernels::mul_i16_scalar, kernels::mul_i16_neon, kernels::mul_i16_neon_supported);
make_int_dispatch!(mul_i8_dispatch, BinaryI8Kernel, kernels::mul_i8_scalar, kernels::mul_i8_neon, kernels::mul_i8_neon_supported);
#[no_mangle] pub unsafe extern "C" fn himada_mul_i32(a: *const i32, b: *const i32, out: *mut i32, n: usize) {
if a.is_null() || b.is_null() || out.is_null() || n == 0 { return; }
let a = unsafe { std::slice::from_raw_parts(a, n) }; let b = unsafe { std::slice::from_raw_parts(b, n) }; let out = unsafe { std::slice::from_raw_parts_mut(out, n) };
let mut guard = match mul_i32_dispatch().lock() { Ok(l) => l, Err(_) => return };
guard.compute(a, b, out)
}
#[no_mangle] pub unsafe extern "C" fn himada_mul_i64(a: *const i64, b: *const i64, out: *mut i64, n: usize) {
if a.is_null() || b.is_null() || out.is_null() || n == 0 { return; }
let a = unsafe { std::slice::from_raw_parts(a, n) }; let b = unsafe { std::slice::from_raw_parts(b, n) }; let out = unsafe { std::slice::from_raw_parts_mut(out, n) };
let mut guard = match mul_i64_dispatch().lock() { Ok(l) => l, Err(_) => return };
guard.compute(a, b, out)
}
#[no_mangle] pub unsafe extern "C" fn himada_mul_i16(a: *const i16, b: *const i16, out: *mut i16, n: usize) {
if a.is_null() || b.is_null() || out.is_null() || n == 0 { return; }
let a = unsafe { std::slice::from_raw_parts(a, n) }; let b = unsafe { std::slice::from_raw_parts(b, n) }; let out = unsafe { std::slice::from_raw_parts_mut(out, n) };
let mut guard = match mul_i16_dispatch().lock() { Ok(l) => l, Err(_) => return };
guard.compute(a, b, out)
}
#[no_mangle] pub unsafe extern "C" fn himada_mul_i8(a: *const i8, b: *const i8, out: *mut i8, n: usize) {
if a.is_null() || b.is_null() || out.is_null() || n == 0 { return; }
let a = unsafe { std::slice::from_raw_parts(a, n) }; let b = unsafe { std::slice::from_raw_parts(b, n) }; let out = unsafe { std::slice::from_raw_parts_mut(out, n) };
let mut guard = match mul_i8_dispatch().lock() { Ok(l) => l, Err(_) => return };
guard.compute(a, b, out)
}
fn sort_i32_dispatch() -> &'static Mutex<Dispatch<SortI32Kernel>> {
static D: OnceLock<Mutex<Dispatch<SortI32Kernel>>> = OnceLock::new();
D.get_or_init(|| {
let mut k: Vec<KernelInfo<SortI32Kernel>> = vec![
KernelInfo { name: "scalar", func: kernels::sort_i32_scalar, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
];
#[cfg(target_arch = "x86_64")]
k.push(add_sse_kernel("SSE", kernels::sort_i32_sse as SortI32Kernel, kernels::sort_i32_sse_supported, 2));
#[cfg(target_arch = "x86_64")]
k.push(add_avx2_kernel("AVX2", kernels::sort_i32_avx2 as SortI32Kernel, kernels::sort_i32_avx2_supported, 3));
#[cfg(target_arch = "aarch64")]
k.push(add_neon_kernel("NEON", kernels::sort_i32_neon as SortI32Kernel, kernels::sort_i32_neon_supported, 2));
Mutex::new(Dispatch::new("sort_i32_ffi", k))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_sort_i32(a: *const i32, out: *mut i32, n: usize) {
if a.is_null() || out.is_null() || n == 0 {
return;
}
let a = unsafe { std::slice::from_raw_parts(a, n) };
let out = unsafe { std::slice::from_raw_parts_mut(out, n) };
out.copy_from_slice(a);
let mut guard = match sort_i32_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return,
};
guard.compute(out)
}
fn cumsum_i32_dispatch() -> &'static Mutex<Dispatch<UnaryI32Kernel>> {
static D: OnceLock<Mutex<Dispatch<UnaryI32Kernel>>> = OnceLock::new();
D.get_or_init(|| {
let mut k: Vec<KernelInfo<UnaryI32Kernel>> = vec![
KernelInfo { name: "scalar", func: kernels::cumsum_i32_scalar as UnaryI32Kernel, is_supported: |_: &HardwareDNA| true, thermal_priority: 1 },
];
#[cfg(target_arch = "x86_64")]
k.push(add_sse_kernel("SSE", kernels::cumsum_i32_sse as UnaryI32Kernel, kernels::cumsum_i32_sse_supported, 2));
#[cfg(target_arch = "x86_64")]
k.push(add_avx2_kernel("AVX2", kernels::cumsum_i32_avx2 as UnaryI32Kernel, kernels::cumsum_i32_avx2_supported, 3));
#[cfg(target_arch = "aarch64")]
k.push(add_neon_kernel("NEON", kernels::cumsum_i32_neon as UnaryI32Kernel, kernels::cumsum_i32_neon_supported, 2));
Mutex::new(Dispatch::new("cumsum_i32_ffi", k))
})
}
#[no_mangle]
pub unsafe extern "C" fn himada_cumsum_i32(input: *const i32, output: *mut i32, n: usize) {
if input.is_null() || output.is_null() || n == 0 {
return;
}
let input = unsafe { std::slice::from_raw_parts(input, n) };
let output = unsafe { std::slice::from_raw_parts_mut(output, n) };
let mut guard = match cumsum_i32_dispatch().lock() {
Ok(lock) => lock,
Err(_) => return,
};
guard.compute(input, output)
}