pub mod kernels {
#[cfg(feature = "simd")]
mod simd_impl {
use wide::f64x4;
pub fn scale_offset_f64x4(values: &mut [f64], scale: f64, offset: f64) {
let scale_v = f64x4::splat(scale);
let offset_v = f64x4::splat(offset);
let chunks = values.len() / 4;
for i in 0..chunks {
let base = i * 4;
let v = f64x4::new([
values[base],
values[base + 1],
values[base + 2],
values[base + 3],
]);
let result = scale_v * v + offset_v;
let arr: [f64; 4] = result.into();
values[base..base + 4].copy_from_slice(&arr);
}
for v in values[chunks * 4..].iter_mut() {
*v = scale * *v + offset;
}
}
pub fn deg_to_rad_batch(values: &mut [f64]) {
scale_offset_f64x4(values, core::f64::consts::PI / 180.0, 0.0);
}
pub fn rad_to_deg_batch(values: &mut [f64]) {
scale_offset_f64x4(values, 180.0 / core::f64::consts::PI, 0.0);
}
pub fn clamp_lat_batch(values: &mut [f64]) {
let min_v = f64x4::splat(-core::f64::consts::FRAC_PI_2);
let max_v = f64x4::splat(core::f64::consts::FRAC_PI_2);
let chunks = values.len() / 4;
for i in 0..chunks {
let base = i * 4;
let v = f64x4::new([
values[base],
values[base + 1],
values[base + 2],
values[base + 3],
]);
let clamped = v.max(min_v).min(max_v);
let arr: [f64; 4] = clamped.into();
values[base..base + 4].copy_from_slice(&arr);
}
for v in values[chunks * 4..].iter_mut() {
*v = v.clamp(-core::f64::consts::FRAC_PI_2, core::f64::consts::FRAC_PI_2);
}
}
}
#[cfg(not(feature = "simd"))]
mod simd_impl {
pub fn scale_offset_f64x4(values: &mut [f64], scale: f64, offset: f64) {
for v in values.iter_mut() {
*v = scale * *v + offset;
}
}
pub fn deg_to_rad_batch(values: &mut [f64]) {
for v in values.iter_mut() {
*v = v.to_radians();
}
}
pub fn rad_to_deg_batch(values: &mut [f64]) {
for v in values.iter_mut() {
*v = v.to_degrees();
}
}
pub fn clamp_lat_batch(values: &mut [f64]) {
for v in values.iter_mut() {
*v = v.clamp(-core::f64::consts::FRAC_PI_2, core::f64::consts::FRAC_PI_2);
}
}
}
pub use simd_impl::{clamp_lat_batch, deg_to_rad_batch, rad_to_deg_batch, scale_offset_f64x4};
}
#[cfg(test)]
mod tests {
use super::kernels::{clamp_lat_batch, deg_to_rad_batch, rad_to_deg_batch, scale_offset_f64x4};
#[cfg(feature = "no_std")]
use alloc::vec;
#[cfg(feature = "no_std")]
use alloc::vec::Vec;
#[test]
fn scale_offset_basic_values() {
let mut vals = vec![1.0_f64, 2.0, 3.0, 4.0, 5.0];
scale_offset_f64x4(&mut vals, 2.0, 1.0);
for (i, &expected) in [3.0_f64, 5.0, 7.0, 9.0, 11.0].iter().enumerate() {
assert!(
(vals[i] - expected).abs() < 1e-12,
"index {i}: got {}, expected {expected}",
vals[i]
);
}
}
#[test]
fn deg_to_rad_matches_scalar() {
let original = vec![0.0_f64, 45.0, 90.0, 135.0, 180.0];
let mut batched = original.clone();
deg_to_rad_batch(&mut batched);
for (i, (orig, batch)) in original.iter().zip(batched.iter()).enumerate() {
let scalar = orig.to_radians();
assert!(
(batch - scalar).abs() < 1e-12,
"index {i}: SIMD={batch}, scalar={scalar}"
);
}
}
#[test]
fn rad_to_deg_matches_scalar() {
let original = vec![
0.0_f64,
core::f64::consts::FRAC_PI_4,
core::f64::consts::FRAC_PI_2,
core::f64::consts::PI,
];
let mut batched = original.clone();
rad_to_deg_batch(&mut batched);
for (i, (orig, batch)) in original.iter().zip(batched.iter()).enumerate() {
let scalar = orig.to_degrees();
assert!(
(batch - scalar).abs() < 1e-12,
"index {i}: SIMD={batch}, scalar={scalar}"
);
}
}
#[test]
fn clamp_lat_clamps_correctly() {
let mut vals = vec![
-2.0_f64,
-core::f64::consts::FRAC_PI_2,
0.0,
core::f64::consts::FRAC_PI_2,
2.0,
];
clamp_lat_batch(&mut vals);
assert!((vals[0] - (-core::f64::consts::FRAC_PI_2)).abs() < 1e-12);
assert!((vals[1] - (-core::f64::consts::FRAC_PI_2)).abs() < 1e-12);
assert!(vals[2].abs() < 1e-12);
assert!((vals[3] - core::f64::consts::FRAC_PI_2).abs() < 1e-12);
assert!((vals[4] - core::f64::consts::FRAC_PI_2).abs() < 1e-12);
}
#[test]
fn scale_offset_empty_slice() {
let mut vals: Vec<f64> = vec![];
scale_offset_f64x4(&mut vals, 3.0, 7.0);
assert!(vals.is_empty());
}
#[test]
fn scale_offset_partial_chunk() {
let mut vals = vec![1.0_f64, 2.0, 3.0, 4.0, 5.0, 6.0];
scale_offset_f64x4(&mut vals, 3.0, 0.0);
for (i, &expected) in [3.0_f64, 6.0, 9.0, 12.0, 15.0, 18.0].iter().enumerate() {
assert!(
(vals[i] - expected).abs() < 1e-12,
"index {i}: got {}, expected {expected}",
vals[i]
);
}
}
}