use std::alloc::{Layout, alloc};
const SIMD_THRESHOLD: usize = 16;
#[unsafe(no_mangle)]
pub extern "C" fn jit_simd_add(a_ptr: *const f64, b_ptr: *const f64, len: u64) -> *mut f64 {
simd_binary_op(a_ptr, b_ptr, len as usize, |a, b| a + b)
}
#[unsafe(no_mangle)]
pub extern "C" fn jit_simd_sub(a_ptr: *const f64, b_ptr: *const f64, len: u64) -> *mut f64 {
simd_binary_op(a_ptr, b_ptr, len as usize, |a, b| a - b)
}
#[unsafe(no_mangle)]
pub extern "C" fn jit_simd_mul(a_ptr: *const f64, b_ptr: *const f64, len: u64) -> *mut f64 {
simd_binary_op(a_ptr, b_ptr, len as usize, |a, b| a * b)
}
#[unsafe(no_mangle)]
pub extern "C" fn jit_simd_div(a_ptr: *const f64, b_ptr: *const f64, len: u64) -> *mut f64 {
simd_binary_op(a_ptr, b_ptr, len as usize, |a, b| a / b)
}
#[unsafe(no_mangle)]
pub extern "C" fn jit_simd_max(a_ptr: *const f64, b_ptr: *const f64, len: u64) -> *mut f64 {
simd_binary_op(a_ptr, b_ptr, len as usize, |a, b| a.max(b))
}
#[unsafe(no_mangle)]
pub extern "C" fn jit_simd_min(a_ptr: *const f64, b_ptr: *const f64, len: u64) -> *mut f64 {
simd_binary_op(a_ptr, b_ptr, len as usize, |a, b| a.min(b))
}
#[unsafe(no_mangle)]
pub extern "C" fn jit_simd_add_scalar(a_ptr: *const f64, scalar: f64, len: u64) -> *mut f64 {
simd_scalar_op(a_ptr, scalar, len as usize, |a, s| a + s)
}
#[unsafe(no_mangle)]
pub extern "C" fn jit_simd_sub_scalar(a_ptr: *const f64, scalar: f64, len: u64) -> *mut f64 {
simd_scalar_op(a_ptr, scalar, len as usize, |a, s| a - s)
}
#[unsafe(no_mangle)]
pub extern "C" fn jit_simd_mul_scalar(a_ptr: *const f64, scalar: f64, len: u64) -> *mut f64 {
simd_scalar_op(a_ptr, scalar, len as usize, |a, s| a * s)
}
#[unsafe(no_mangle)]
pub extern "C" fn jit_simd_div_scalar(a_ptr: *const f64, scalar: f64, len: u64) -> *mut f64 {
simd_scalar_op(a_ptr, scalar, len as usize, |a, s| a / s)
}
#[unsafe(no_mangle)]
pub extern "C" fn jit_simd_gt(a_ptr: *const f64, b_ptr: *const f64, len: u64) -> *mut f64 {
simd_cmp_op(a_ptr, b_ptr, len as usize, |a, b| a > b)
}
#[unsafe(no_mangle)]
pub extern "C" fn jit_simd_lt(a_ptr: *const f64, b_ptr: *const f64, len: u64) -> *mut f64 {
simd_cmp_op(a_ptr, b_ptr, len as usize, |a, b| a < b)
}
#[unsafe(no_mangle)]
pub extern "C" fn jit_simd_gte(a_ptr: *const f64, b_ptr: *const f64, len: u64) -> *mut f64 {
simd_cmp_op(a_ptr, b_ptr, len as usize, |a, b| a >= b)
}
#[unsafe(no_mangle)]
pub extern "C" fn jit_simd_lte(a_ptr: *const f64, b_ptr: *const f64, len: u64) -> *mut f64 {
simd_cmp_op(a_ptr, b_ptr, len as usize, |a, b| a <= b)
}
#[unsafe(no_mangle)]
pub extern "C" fn jit_simd_eq(a_ptr: *const f64, b_ptr: *const f64, len: u64) -> *mut f64 {
simd_cmp_op(a_ptr, b_ptr, len as usize, |a, b| {
(a - b).abs() < f64::EPSILON
})
}
#[unsafe(no_mangle)]
pub extern "C" fn jit_simd_neq(a_ptr: *const f64, b_ptr: *const f64, len: u64) -> *mut f64 {
simd_cmp_op(a_ptr, b_ptr, len as usize, |a, b| {
(a - b).abs() >= f64::EPSILON
})
}
#[inline]
fn alloc_f64_buffer(len: usize) -> *mut f64 {
if len == 0 {
return std::ptr::null_mut();
}
let layout =
Layout::from_size_align(len * std::mem::size_of::<f64>(), 32).expect("Invalid layout");
unsafe { alloc(layout) as *mut f64 }
}
#[inline]
fn simd_binary_op<F>(a_ptr: *const f64, b_ptr: *const f64, len: usize, op: F) -> *mut f64
where
F: Fn(f64, f64) -> f64,
{
if a_ptr.is_null() || b_ptr.is_null() || len == 0 {
return std::ptr::null_mut();
}
let result = alloc_f64_buffer(len);
if result.is_null() {
return std::ptr::null_mut();
}
unsafe {
let a = std::slice::from_raw_parts(a_ptr, len);
let b = std::slice::from_raw_parts(b_ptr, len);
let out = std::slice::from_raw_parts_mut(result, len);
if len >= SIMD_THRESHOLD {
let chunks = len / 4;
for i in 0..chunks {
let idx = i * 4;
out[idx] = op(a[idx], b[idx]);
out[idx + 1] = op(a[idx + 1], b[idx + 1]);
out[idx + 2] = op(a[idx + 2], b[idx + 2]);
out[idx + 3] = op(a[idx + 3], b[idx + 3]);
}
for i in (chunks * 4)..len {
out[i] = op(a[i], b[i]);
}
} else {
for i in 0..len {
out[i] = op(a[i], b[i]);
}
}
}
result
}
#[inline]
fn simd_scalar_op<F>(a_ptr: *const f64, scalar: f64, len: usize, op: F) -> *mut f64
where
F: Fn(f64, f64) -> f64,
{
if a_ptr.is_null() || len == 0 {
return std::ptr::null_mut();
}
let result = alloc_f64_buffer(len);
if result.is_null() {
return std::ptr::null_mut();
}
unsafe {
let a = std::slice::from_raw_parts(a_ptr, len);
let out = std::slice::from_raw_parts_mut(result, len);
if len >= SIMD_THRESHOLD {
let chunks = len / 4;
for i in 0..chunks {
let idx = i * 4;
out[idx] = op(a[idx], scalar);
out[idx + 1] = op(a[idx + 1], scalar);
out[idx + 2] = op(a[idx + 2], scalar);
out[idx + 3] = op(a[idx + 3], scalar);
}
for i in (chunks * 4)..len {
out[i] = op(a[i], scalar);
}
} else {
for i in 0..len {
out[i] = op(a[i], scalar);
}
}
}
result
}
#[inline]
fn simd_cmp_op<F>(a_ptr: *const f64, b_ptr: *const f64, len: usize, op: F) -> *mut f64
where
F: Fn(f64, f64) -> bool,
{
if a_ptr.is_null() || b_ptr.is_null() || len == 0 {
return std::ptr::null_mut();
}
let result = alloc_f64_buffer(len);
if result.is_null() {
return std::ptr::null_mut();
}
unsafe {
let a = std::slice::from_raw_parts(a_ptr, len);
let b = std::slice::from_raw_parts(b_ptr, len);
let out = std::slice::from_raw_parts_mut(result, len);
if len >= SIMD_THRESHOLD {
let chunks = len / 4;
for i in 0..chunks {
let idx = i * 4;
out[idx] = if op(a[idx], b[idx]) { 1.0 } else { 0.0 };
out[idx + 1] = if op(a[idx + 1], b[idx + 1]) { 1.0 } else { 0.0 };
out[idx + 2] = if op(a[idx + 2], b[idx + 2]) { 1.0 } else { 0.0 };
out[idx + 3] = if op(a[idx + 3], b[idx + 3]) { 1.0 } else { 0.0 };
}
for i in (chunks * 4)..len {
out[i] = if op(a[i], b[i]) { 1.0 } else { 0.0 };
}
} else {
for i in 0..len {
out[i] = if op(a[i], b[i]) { 1.0 } else { 0.0 };
}
}
}
result
}
#[unsafe(no_mangle)]
pub extern "C" fn jit_simd_free(ptr: *mut f64, len: u64) {
if ptr.is_null() || len == 0 {
return;
}
let layout = Layout::from_size_align(len as usize * std::mem::size_of::<f64>(), 32)
.expect("Invalid layout");
unsafe {
std::alloc::dealloc(ptr as *mut u8, layout);
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_simd_add() {
let a = vec![1.0, 2.0, 3.0, 4.0];
let b = vec![10.0, 20.0, 30.0, 40.0];
let result = jit_simd_add(a.as_ptr(), b.as_ptr(), 4);
unsafe {
assert_eq!(*result, 11.0);
assert_eq!(*result.add(1), 22.0);
assert_eq!(*result.add(2), 33.0);
assert_eq!(*result.add(3), 44.0);
}
jit_simd_free(result, 4);
}
#[test]
fn test_simd_mul_large() {
let len = 1000;
let a: Vec<f64> = (0..len).map(|i| i as f64).collect();
let b: Vec<f64> = (0..len).map(|i| (i * 2) as f64).collect();
let result = jit_simd_mul(a.as_ptr(), b.as_ptr(), len as u64);
unsafe {
for i in 0..len {
assert_eq!(*result.add(i), (i * i * 2) as f64);
}
}
jit_simd_free(result, len as u64);
}
#[test]
fn test_simd_gt() {
let a = vec![5.0, 2.0, 8.0, 1.0];
let b = vec![3.0, 4.0, 8.0, 0.0];
let result = jit_simd_gt(a.as_ptr(), b.as_ptr(), 4);
unsafe {
assert_eq!(*result, 1.0); assert_eq!(*result.add(1), 0.0); assert_eq!(*result.add(2), 0.0); assert_eq!(*result.add(3), 1.0); }
jit_simd_free(result, 4);
}
}