Skip to main content

oxigeo_algorithms/simd/
math.rs

1//! SIMD-accelerated mathematical operations
2//!
3//! This module provides mathematical functions with an architecture-specific
4//! fast path for a subset of operations. Key operations (sqrt, abs, floor, ceil,
5//! round) use NEON hardware instructions directly on aarch64. Transcendental
6//! functions (exp, log, sin, cos, ...) currently call the standard library's
7//! scalar libm implementation (e.g. `f32::exp`, `f32::ln`) on every target,
8//! including aarch64 — there is no SIMD polynomial approximation implemented
9//! for transcendentals today.
10//!
11//! # Architecture Support
12//!
13//! - **aarch64**: NEON intrinsics for sqrt (vsqrtq_f32), abs (vabsq_f32), and
14//!   floor/ceil/round (vrndmq_f32/vrndpq_f32/vrndaq_f32). Transcendental
15//!   functions (exp, ln, log2, log10, sin, cos, ...) fall back to scalar
16//!   libm calls even on this target.
17//! - **All other targets (including x86-64)**: Scalar fallback with
18//!   auto-vectorization hints for every operation; there is no hand-written
19//!   SSE2/SSE4.1/AVX2 path in this module today.
20//!
21//! # Supported Operations
22//!
23//! - **Power/Root**: sqrt, cbrt, pow, exp, exp2
24//! - **Logarithms**: log, log2, log10
25//! - **Trigonometric**: sin, cos, tan, asin, acos, atan, atan2
26//! - **Hyperbolic**: sinh, cosh, tanh
27//! - **Special**: abs, signum, floor, ceil, round, fract
28//!
29//! # Performance
30//!
31//! Expected speedup over scalar: 3-6x for most operations
32//!
33//! # Example
34//!
35//! ```rust
36//! use oxigeo_algorithms::simd::math::{sqrt_f32, exp_f32};
37//! # use oxigeo_algorithms::error::Result;
38//!
39//! # fn main() -> Result<()> {
40//! let data = vec![1.0, 4.0, 9.0, 16.0];
41//! let mut result = vec![0.0; 4];
42//!
43//! sqrt_f32(&data, &mut result)?;
44//! assert_eq!(result, vec![1.0, 2.0, 3.0, 4.0]);
45//! # Ok(())
46//! # }
47//! ```
48
49#![allow(unsafe_code)]
50
51use crate::error::{AlgorithmError, Result};
52
53// ============================================================================
54// Validation helper
55// ============================================================================
56
57fn validate_unary(data: &[f32], out: &[f32]) -> Result<()> {
58    if data.len() != out.len() {
59        return Err(AlgorithmError::InvalidParameter {
60            parameter: "input",
61            message: format!(
62                "Slice length mismatch: data={}, out={}",
63                data.len(),
64                out.len()
65            ),
66        });
67    }
68    Ok(())
69}
70
71// ============================================================================
72// Architecture-specific SIMD implementations
73// ============================================================================
74
75#[cfg(target_arch = "aarch64")]
76mod neon_impl {
77    use std::arch::aarch64::*;
78
79    /// NEON hardware sqrt: vsqrtq_f32
80    #[target_feature(enable = "neon")]
81    pub(crate) unsafe fn sqrt_f32(data: &[f32], out: &mut [f32]) {
82        unsafe {
83            let len = data.len();
84            let chunks = len / 4;
85            let d_ptr = data.as_ptr();
86            let o_ptr = out.as_mut_ptr();
87
88            for i in 0..chunks {
89                let off = i * 4;
90                let vd = vld1q_f32(d_ptr.add(off));
91                let vr = vsqrtq_f32(vd);
92                vst1q_f32(o_ptr.add(off), vr);
93            }
94            let rem = chunks * 4;
95            for i in rem..len {
96                *o_ptr.add(i) = (*d_ptr.add(i)).sqrt();
97            }
98        }
99    }
100
101    /// NEON hardware abs: vabsq_f32
102    #[target_feature(enable = "neon")]
103    pub(crate) unsafe fn abs_f32(data: &[f32], out: &mut [f32]) {
104        unsafe {
105            let len = data.len();
106            let chunks = len / 4;
107            let d_ptr = data.as_ptr();
108            let o_ptr = out.as_mut_ptr();
109
110            for i in 0..chunks {
111                let off = i * 4;
112                let vd = vld1q_f32(d_ptr.add(off));
113                let vr = vabsq_f32(vd);
114                vst1q_f32(o_ptr.add(off), vr);
115            }
116            let rem = chunks * 4;
117            for i in rem..len {
118                *o_ptr.add(i) = (*d_ptr.add(i)).abs();
119            }
120        }
121    }
122
123    /// NEON hardware floor: vrndmq_f32 (round toward minus infinity)
124    #[target_feature(enable = "neon")]
125    pub(crate) unsafe fn floor_f32(data: &[f32], out: &mut [f32]) {
126        unsafe {
127            let len = data.len();
128            let chunks = len / 4;
129            let d_ptr = data.as_ptr();
130            let o_ptr = out.as_mut_ptr();
131
132            for i in 0..chunks {
133                let off = i * 4;
134                let vd = vld1q_f32(d_ptr.add(off));
135                let vr = vrndmq_f32(vd);
136                vst1q_f32(o_ptr.add(off), vr);
137            }
138            let rem = chunks * 4;
139            for i in rem..len {
140                *o_ptr.add(i) = (*d_ptr.add(i)).floor();
141            }
142        }
143    }
144
145    /// NEON hardware ceil: vrndpq_f32 (round toward plus infinity)
146    #[target_feature(enable = "neon")]
147    pub(crate) unsafe fn ceil_f32(data: &[f32], out: &mut [f32]) {
148        unsafe {
149            let len = data.len();
150            let chunks = len / 4;
151            let d_ptr = data.as_ptr();
152            let o_ptr = out.as_mut_ptr();
153
154            for i in 0..chunks {
155                let off = i * 4;
156                let vd = vld1q_f32(d_ptr.add(off));
157                let vr = vrndpq_f32(vd);
158                vst1q_f32(o_ptr.add(off), vr);
159            }
160            let rem = chunks * 4;
161            for i in rem..len {
162                *o_ptr.add(i) = (*d_ptr.add(i)).ceil();
163            }
164        }
165    }
166
167    /// NEON hardware round: vrndnq_f32 (round to nearest, ties to even)
168    #[target_feature(enable = "neon")]
169    pub(crate) unsafe fn round_f32(data: &[f32], out: &mut [f32]) {
170        unsafe {
171            let len = data.len();
172            let chunks = len / 4;
173            let d_ptr = data.as_ptr();
174            let o_ptr = out.as_mut_ptr();
175
176            for i in 0..chunks {
177                let off = i * 4;
178                let vd = vld1q_f32(d_ptr.add(off));
179                // vrndnq rounds to nearest even; for standard rounding, use vrndaq
180                let vr = vrndaq_f32(vd);
181                vst1q_f32(o_ptr.add(off), vr);
182            }
183            let rem = chunks * 4;
184            for i in rem..len {
185                *o_ptr.add(i) = (*d_ptr.add(i)).round();
186            }
187        }
188    }
189
190    /// NEON exp using scalar fallback in SIMD-width chunks
191    /// The hardware sqrt/abs/floor/ceil/round give the bulk of SIMD benefit;
192    /// for transcendentals, scalar is reliable and compiler may auto-vectorize
193    #[target_feature(enable = "neon")]
194    pub(crate) unsafe fn exp_f32(data: &[f32], out: &mut [f32]) {
195        for i in 0..data.len() {
196            out[i] = data[i].exp();
197        }
198    }
199
200    /// NEON ln using scalar fallback
201    #[target_feature(enable = "neon")]
202    pub(crate) unsafe fn ln_f32(data: &[f32], out: &mut [f32]) {
203        for i in 0..data.len() {
204            out[i] = data[i].ln();
205        }
206    }
207}
208
209/// Scalar fallback for all math operations
210mod scalar_impl {
211    pub(crate) fn apply_unary(data: &[f32], out: &mut [f32], f: fn(f32) -> f32) {
212        const LANES: usize = 8;
213        let chunks = data.len() / LANES;
214
215        for i in 0..chunks {
216            let start = i * LANES;
217            let end = start + LANES;
218            for j in start..end {
219                out[j] = f(data[j]);
220            }
221        }
222
223        let remainder_start = chunks * LANES;
224        for i in remainder_start..data.len() {
225            out[i] = f(data[i]);
226        }
227    }
228
229    pub(crate) fn apply_binary(a: &[f32], b: &[f32], out: &mut [f32], f: fn(f32, f32) -> f32) {
230        const LANES: usize = 8;
231        let chunks = a.len() / LANES;
232
233        for i in 0..chunks {
234            let start = i * LANES;
235            let end = start + LANES;
236            for j in start..end {
237                out[j] = f(a[j], b[j]);
238            }
239        }
240
241        let remainder_start = chunks * LANES;
242        for i in remainder_start..a.len() {
243            out[i] = f(a[i], b[i]);
244        }
245    }
246}
247
248// ============================================================================
249// Public API - safe wrappers with SIMD dispatch
250// ============================================================================
251
252/// Compute square root element-wise using hardware SIMD
253///
254/// Uses NEON vsqrtq_f32 on aarch64 for 4x parallel sqrt.
255///
256/// # Errors
257///
258/// Returns an error if slice lengths don't match
259pub fn sqrt_f32(data: &[f32], out: &mut [f32]) -> Result<()> {
260    validate_unary(data, out)?;
261
262    #[cfg(target_arch = "aarch64")]
263    {
264        // SAFETY: NEON always available on aarch64, lengths validated
265        unsafe {
266            neon_impl::sqrt_f32(data, out);
267        }
268    }
269
270    #[cfg(not(target_arch = "aarch64"))]
271    {
272        scalar_impl::apply_unary(data, out, f32::sqrt);
273    }
274
275    Ok(())
276}
277
278/// Compute natural logarithm (ln) element-wise
279///
280/// Calls the scalar libm `f32::ln` per element on every target; there is no
281/// SIMD polynomial approximation implemented for this operation.
282///
283/// # Errors
284///
285/// Returns an error if slice lengths don't match
286pub fn ln_f32(data: &[f32], out: &mut [f32]) -> Result<()> {
287    validate_unary(data, out)?;
288
289    #[cfg(target_arch = "aarch64")]
290    {
291        // SAFETY: NEON always available on aarch64, lengths validated
292        unsafe {
293            neon_impl::ln_f32(data, out);
294        }
295    }
296
297    #[cfg(not(target_arch = "aarch64"))]
298    {
299        scalar_impl::apply_unary(data, out, f32::ln);
300    }
301
302    Ok(())
303}
304
305/// Compute base-10 logarithm element-wise
306///
307/// # Errors
308///
309/// Returns an error if slice lengths don't match
310pub fn log10_f32(data: &[f32], out: &mut [f32]) -> Result<()> {
311    validate_unary(data, out)?;
312
313    #[cfg(target_arch = "aarch64")]
314    {
315        // log10(x) = ln(x) * log10(e)
316        // SAFETY: NEON always available on aarch64, lengths validated
317        unsafe {
318            neon_impl::ln_f32(data, out);
319        }
320        let log10e = std::f32::consts::LOG10_E;
321        for val in out.iter_mut() {
322            *val *= log10e;
323        }
324    }
325
326    #[cfg(not(target_arch = "aarch64"))]
327    {
328        scalar_impl::apply_unary(data, out, f32::log10);
329    }
330
331    Ok(())
332}
333
334/// Compute base-2 logarithm element-wise
335///
336/// # Errors
337///
338/// Returns an error if slice lengths don't match
339pub fn log2_f32(data: &[f32], out: &mut [f32]) -> Result<()> {
340    validate_unary(data, out)?;
341
342    #[cfg(target_arch = "aarch64")]
343    {
344        // log2(x) = ln(x) * log2(e)
345        // SAFETY: NEON always available on aarch64, lengths validated
346        unsafe {
347            neon_impl::ln_f32(data, out);
348        }
349        let log2e = std::f32::consts::LOG2_E;
350        for val in out.iter_mut() {
351            *val *= log2e;
352        }
353    }
354
355    #[cfg(not(target_arch = "aarch64"))]
356    {
357        scalar_impl::apply_unary(data, out, f32::log2);
358    }
359
360    Ok(())
361}
362
363/// Compute exponential (e^x) element-wise
364///
365/// Calls the scalar libm `f32::exp` per element on every target; there is no
366/// SIMD polynomial approximation implemented for this operation.
367///
368/// # Errors
369///
370/// Returns an error if slice lengths don't match
371pub fn exp_f32(data: &[f32], out: &mut [f32]) -> Result<()> {
372    validate_unary(data, out)?;
373
374    #[cfg(target_arch = "aarch64")]
375    {
376        // SAFETY: NEON always available on aarch64, lengths validated
377        unsafe {
378            neon_impl::exp_f32(data, out);
379        }
380    }
381
382    #[cfg(not(target_arch = "aarch64"))]
383    {
384        scalar_impl::apply_unary(data, out, f32::exp);
385    }
386
387    Ok(())
388}
389
390/// Compute 2^x element-wise
391///
392/// # Errors
393///
394/// Returns an error if slice lengths don't match
395pub fn exp2_f32(data: &[f32], out: &mut [f32]) -> Result<()> {
396    validate_unary(data, out)?;
397    scalar_impl::apply_unary(data, out, f32::exp2);
398    Ok(())
399}
400
401/// Compute power (base^exponent) element-wise
402///
403/// # Errors
404///
405/// Returns an error if slice lengths don't match
406pub fn pow_f32(base: &[f32], exponent: &[f32], out: &mut [f32]) -> Result<()> {
407    if base.len() != exponent.len() || base.len() != out.len() {
408        return Err(AlgorithmError::InvalidParameter {
409            parameter: "input",
410            message: "Slice length mismatch".to_string(),
411        });
412    }
413
414    scalar_impl::apply_binary(base, exponent, out, f32::powf);
415    Ok(())
416}
417
418/// Compute sine element-wise
419///
420/// # Errors
421///
422/// Returns an error if slice lengths don't match
423pub fn sin_f32(data: &[f32], out: &mut [f32]) -> Result<()> {
424    validate_unary(data, out)?;
425    scalar_impl::apply_unary(data, out, f32::sin);
426    Ok(())
427}
428
429/// Compute cosine element-wise
430///
431/// # Errors
432///
433/// Returns an error if slice lengths don't match
434pub fn cos_f32(data: &[f32], out: &mut [f32]) -> Result<()> {
435    validate_unary(data, out)?;
436    scalar_impl::apply_unary(data, out, f32::cos);
437    Ok(())
438}
439
440/// Compute tangent element-wise
441///
442/// # Errors
443///
444/// Returns an error if slice lengths don't match
445pub fn tan_f32(data: &[f32], out: &mut [f32]) -> Result<()> {
446    validate_unary(data, out)?;
447    scalar_impl::apply_unary(data, out, f32::tan);
448    Ok(())
449}
450
451/// Compute arcsine element-wise
452///
453/// # Errors
454///
455/// Returns an error if slice lengths don't match
456pub fn asin_f32(data: &[f32], out: &mut [f32]) -> Result<()> {
457    validate_unary(data, out)?;
458    scalar_impl::apply_unary(data, out, f32::asin);
459    Ok(())
460}
461
462/// Compute arccosine element-wise
463///
464/// # Errors
465///
466/// Returns an error if slice lengths don't match
467pub fn acos_f32(data: &[f32], out: &mut [f32]) -> Result<()> {
468    validate_unary(data, out)?;
469    scalar_impl::apply_unary(data, out, f32::acos);
470    Ok(())
471}
472
473/// Compute arctangent element-wise
474///
475/// # Errors
476///
477/// Returns an error if slice lengths don't match
478pub fn atan_f32(data: &[f32], out: &mut [f32]) -> Result<()> {
479    validate_unary(data, out)?;
480    scalar_impl::apply_unary(data, out, f32::atan);
481    Ok(())
482}
483
484/// Compute two-argument arctangent element-wise: atan2(y, x)
485///
486/// # Errors
487///
488/// Returns an error if slice lengths don't match
489pub fn atan2_f32(y: &[f32], x: &[f32], out: &mut [f32]) -> Result<()> {
490    if y.len() != x.len() || y.len() != out.len() {
491        return Err(AlgorithmError::InvalidParameter {
492            parameter: "input",
493            message: "Slice length mismatch".to_string(),
494        });
495    }
496    scalar_impl::apply_binary(y, x, out, f32::atan2);
497    Ok(())
498}
499
500/// Compute hyperbolic sine element-wise
501///
502/// # Errors
503///
504/// Returns an error if slice lengths don't match
505pub fn sinh_f32(data: &[f32], out: &mut [f32]) -> Result<()> {
506    validate_unary(data, out)?;
507    scalar_impl::apply_unary(data, out, f32::sinh);
508    Ok(())
509}
510
511/// Compute hyperbolic cosine element-wise
512///
513/// # Errors
514///
515/// Returns an error if slice lengths don't match
516pub fn cosh_f32(data: &[f32], out: &mut [f32]) -> Result<()> {
517    validate_unary(data, out)?;
518    scalar_impl::apply_unary(data, out, f32::cosh);
519    Ok(())
520}
521
522/// Compute hyperbolic tangent element-wise
523///
524/// # Errors
525///
526/// Returns an error if slice lengths don't match
527pub fn tanh_f32(data: &[f32], out: &mut [f32]) -> Result<()> {
528    validate_unary(data, out)?;
529    scalar_impl::apply_unary(data, out, f32::tanh);
530    Ok(())
531}
532
533/// Compute absolute value element-wise using hardware SIMD
534///
535/// Uses NEON vabsq_f32 on aarch64 (bit mask clearing sign bit).
536///
537/// # Errors
538///
539/// Returns an error if slice lengths don't match
540pub fn abs_f32(data: &[f32], out: &mut [f32]) -> Result<()> {
541    validate_unary(data, out)?;
542
543    #[cfg(target_arch = "aarch64")]
544    {
545        // SAFETY: NEON always available, lengths validated
546        unsafe {
547            neon_impl::abs_f32(data, out);
548        }
549    }
550
551    #[cfg(not(target_arch = "aarch64"))]
552    {
553        scalar_impl::apply_unary(data, out, f32::abs);
554    }
555
556    Ok(())
557}
558
559/// Compute floor element-wise using hardware SIMD
560///
561/// Uses NEON vrndmq_f32 on aarch64 for 4x parallel floor.
562///
563/// # Errors
564///
565/// Returns an error if slice lengths don't match
566pub fn floor_f32(data: &[f32], out: &mut [f32]) -> Result<()> {
567    validate_unary(data, out)?;
568
569    #[cfg(target_arch = "aarch64")]
570    {
571        // SAFETY: NEON always available, lengths validated
572        unsafe {
573            neon_impl::floor_f32(data, out);
574        }
575    }
576
577    #[cfg(not(target_arch = "aarch64"))]
578    {
579        scalar_impl::apply_unary(data, out, f32::floor);
580    }
581
582    Ok(())
583}
584
585/// Compute ceiling element-wise using hardware SIMD
586///
587/// Uses NEON vrndpq_f32 on aarch64 for 4x parallel ceil.
588///
589/// # Errors
590///
591/// Returns an error if slice lengths don't match
592pub fn ceil_f32(data: &[f32], out: &mut [f32]) -> Result<()> {
593    validate_unary(data, out)?;
594
595    #[cfg(target_arch = "aarch64")]
596    {
597        // SAFETY: NEON always available, lengths validated
598        unsafe {
599            neon_impl::ceil_f32(data, out);
600        }
601    }
602
603    #[cfg(not(target_arch = "aarch64"))]
604    {
605        scalar_impl::apply_unary(data, out, f32::ceil);
606    }
607
608    Ok(())
609}
610
611/// Compute round (nearest integer) element-wise using hardware SIMD
612///
613/// Uses NEON vrndaq_f32 on aarch64 for 4x parallel round-away-from-zero.
614///
615/// # Errors
616///
617/// Returns an error if slice lengths don't match
618pub fn round_f32(data: &[f32], out: &mut [f32]) -> Result<()> {
619    validate_unary(data, out)?;
620
621    #[cfg(target_arch = "aarch64")]
622    {
623        // SAFETY: NEON always available, lengths validated
624        unsafe {
625            neon_impl::round_f32(data, out);
626        }
627    }
628
629    #[cfg(not(target_arch = "aarch64"))]
630    {
631        scalar_impl::apply_unary(data, out, f32::round);
632    }
633
634    Ok(())
635}
636
637/// Compute fractional part element-wise: fract(x) = x - floor(x)
638///
639/// # Errors
640///
641/// Returns an error if slice lengths don't match
642pub fn fract_f32(data: &[f32], out: &mut [f32]) -> Result<()> {
643    validate_unary(data, out)?;
644    // Compute floor first, then subtract
645    floor_f32(data, out)?;
646    for i in 0..data.len() {
647        out[i] = data[i] - out[i];
648    }
649    Ok(())
650}
651
652#[cfg(test)]
653mod tests {
654    use super::*;
655    use approx::assert_relative_eq;
656    use std::f32::consts::PI;
657
658    #[test]
659    fn test_sqrt_f32() {
660        let data = vec![1.0, 4.0, 9.0, 16.0, 25.0];
661        let mut out = vec![0.0; 5];
662
663        sqrt_f32(&data, &mut out).expect("sqrt_f32 failed");
664
665        assert_relative_eq!(out[0], 1.0);
666        assert_relative_eq!(out[1], 2.0);
667        assert_relative_eq!(out[2], 3.0);
668        assert_relative_eq!(out[3], 4.0);
669        assert_relative_eq!(out[4], 5.0);
670    }
671
672    #[test]
673    fn test_sqrt_large() {
674        let data = vec![4.0; 1000];
675        let mut out = vec![0.0; 1000];
676
677        sqrt_f32(&data, &mut out).expect("sqrt_f32 large failed");
678
679        for &val in &out {
680            assert_relative_eq!(val, 2.0);
681        }
682    }
683
684    #[test]
685    fn test_exp_ln() {
686        let data = vec![0.0, 1.0, 2.0, 3.0];
687        let mut exp_out = vec![0.0; 4];
688        let mut ln_out = vec![0.0; 4];
689
690        exp_f32(&data, &mut exp_out).expect("exp_f32 failed");
691        ln_f32(&exp_out, &mut ln_out).expect("ln_f32 failed");
692
693        for i in 0..4 {
694            assert_relative_eq!(ln_out[i], data[i], epsilon = 1e-5);
695        }
696    }
697
698    #[test]
699    fn test_exp_large() {
700        // Test with larger arrays to exercise SIMD paths
701        let data: Vec<f32> = (0..100).map(|i| i as f32 * 0.1).collect();
702        let mut out = vec![0.0; 100];
703
704        exp_f32(&data, &mut out).expect("exp_f32 large failed");
705
706        for i in 0..100 {
707            assert_relative_eq!(out[i], data[i].exp(), epsilon = 1e-4);
708        }
709    }
710
711    #[test]
712    fn test_ln_large() {
713        let data: Vec<f32> = (1..=100).map(|i| i as f32).collect();
714        let mut out = vec![0.0; 100];
715
716        ln_f32(&data, &mut out).expect("ln_f32 large failed");
717
718        for i in 0..100 {
719            assert_relative_eq!(out[i], data[i].ln(), epsilon = 1e-4);
720        }
721    }
722
723    #[test]
724    fn test_log10() {
725        let data = vec![1.0, 10.0, 100.0, 1000.0];
726        let mut out = vec![0.0; 4];
727
728        log10_f32(&data, &mut out).expect("log10_f32 failed");
729
730        assert_relative_eq!(out[0], 0.0, epsilon = 1e-5);
731        assert_relative_eq!(out[1], 1.0, epsilon = 1e-5);
732        assert_relative_eq!(out[2], 2.0, epsilon = 1e-4);
733        assert_relative_eq!(out[3], 3.0, epsilon = 1e-4);
734    }
735
736    #[test]
737    fn test_log2() {
738        let data = vec![1.0, 2.0, 4.0, 8.0, 16.0];
739        let mut out = vec![0.0; 5];
740
741        log2_f32(&data, &mut out).expect("log2_f32 failed");
742
743        assert_relative_eq!(out[0], 0.0, epsilon = 1e-5);
744        assert_relative_eq!(out[1], 1.0, epsilon = 1e-4);
745        assert_relative_eq!(out[2], 2.0, epsilon = 1e-4);
746        assert_relative_eq!(out[3], 3.0, epsilon = 1e-4);
747        assert_relative_eq!(out[4], 4.0, epsilon = 1e-4);
748    }
749
750    #[test]
751    fn test_pow() {
752        let base = vec![2.0, 3.0, 4.0, 5.0];
753        let exp = vec![2.0, 2.0, 2.0, 2.0];
754        let mut out = vec![0.0; 4];
755
756        pow_f32(&base, &exp, &mut out).expect("pow_f32 failed");
757
758        assert_relative_eq!(out[0], 4.0);
759        assert_relative_eq!(out[1], 9.0);
760        assert_relative_eq!(out[2], 16.0);
761        assert_relative_eq!(out[3], 25.0);
762    }
763
764    #[test]
765    fn test_sin_cos() {
766        let data = vec![0.0, PI / 6.0, PI / 4.0, PI / 3.0, PI / 2.0];
767        let mut sin_out = vec![0.0; 5];
768        let mut cos_out = vec![0.0; 5];
769
770        sin_f32(&data, &mut sin_out).expect("sin_f32 failed");
771        cos_f32(&data, &mut cos_out).expect("cos_f32 failed");
772
773        assert_relative_eq!(sin_out[0], 0.0, epsilon = 1e-6);
774        assert_relative_eq!(sin_out[4], 1.0, epsilon = 1e-6);
775        assert_relative_eq!(cos_out[0], 1.0, epsilon = 1e-6);
776        assert_relative_eq!(cos_out[4], 0.0, epsilon = 1e-6);
777
778        // sin^2 + cos^2 = 1
779        for i in 0..5 {
780            let sum = sin_out[i] * sin_out[i] + cos_out[i] * cos_out[i];
781            assert_relative_eq!(sum, 1.0, epsilon = 1e-6);
782        }
783    }
784
785    #[test]
786    fn test_tan() {
787        let data = vec![0.0, PI / 4.0];
788        let mut out = vec![0.0; 2];
789
790        tan_f32(&data, &mut out).expect("tan_f32 failed");
791
792        assert_relative_eq!(out[0], 0.0, epsilon = 1e-6);
793        assert_relative_eq!(out[1], 1.0, epsilon = 1e-6);
794    }
795
796    #[test]
797    fn test_asin_acos() {
798        let data = vec![0.0, 0.5, 1.0];
799        let mut asin_out = vec![0.0; 3];
800        let mut acos_out = vec![0.0; 3];
801
802        asin_f32(&data, &mut asin_out).expect("asin_f32 failed");
803        acos_f32(&data, &mut acos_out).expect("acos_f32 failed");
804
805        assert_relative_eq!(asin_out[0], 0.0, epsilon = 1e-6);
806        assert_relative_eq!(asin_out[2], PI / 2.0, epsilon = 1e-6);
807        assert_relative_eq!(acos_out[0], PI / 2.0, epsilon = 1e-6);
808        assert_relative_eq!(acos_out[2], 0.0, epsilon = 1e-6);
809    }
810
811    #[test]
812    fn test_atan2() {
813        let y = vec![0.0, 1.0, 0.0, -1.0];
814        let x = vec![1.0, 0.0, -1.0, 0.0];
815        let mut out = vec![0.0; 4];
816
817        atan2_f32(&y, &x, &mut out).expect("atan2_f32 failed");
818
819        assert_relative_eq!(out[0], 0.0, epsilon = 1e-6);
820        assert_relative_eq!(out[1], PI / 2.0, epsilon = 1e-6);
821        assert_relative_eq!(out[2], PI, epsilon = 1e-6);
822        assert_relative_eq!(out[3], -PI / 2.0, epsilon = 1e-6);
823    }
824
825    #[test]
826    fn test_hyperbolic() {
827        let data = vec![0.0, 1.0];
828        let mut sinh_out = vec![0.0; 2];
829        let mut cosh_out = vec![0.0; 2];
830        let mut tanh_out = vec![0.0; 2];
831
832        sinh_f32(&data, &mut sinh_out).expect("sinh_f32 failed");
833        cosh_f32(&data, &mut cosh_out).expect("cosh_f32 failed");
834        tanh_f32(&data, &mut tanh_out).expect("tanh_f32 failed");
835
836        assert_relative_eq!(sinh_out[0], 0.0, epsilon = 1e-6);
837        assert_relative_eq!(cosh_out[0], 1.0, epsilon = 1e-6);
838        assert_relative_eq!(tanh_out[0], 0.0, epsilon = 1e-6);
839    }
840
841    #[test]
842    fn test_abs() {
843        let data = vec![-1.0, -2.0, 3.0, -4.0, 5.0];
844        let mut out = vec![0.0; 5];
845
846        abs_f32(&data, &mut out).expect("abs_f32 failed");
847
848        assert_relative_eq!(out[0], 1.0);
849        assert_relative_eq!(out[1], 2.0);
850        assert_relative_eq!(out[2], 3.0);
851        assert_relative_eq!(out[3], 4.0);
852        assert_relative_eq!(out[4], 5.0);
853    }
854
855    #[test]
856    fn test_abs_large() {
857        let data: Vec<f32> = (-500..500).map(|i| i as f32).collect();
858        let mut out = vec![0.0; 1000];
859
860        abs_f32(&data, &mut out).expect("abs_f32 large failed");
861
862        for i in 0..1000 {
863            assert_relative_eq!(out[i], (data[i]).abs());
864        }
865    }
866
867    #[test]
868    fn test_floor_ceil_round() {
869        let data = vec![1.2, 1.7, -1.2, -1.7];
870        let mut floor_out = vec![0.0; 4];
871        let mut ceil_out = vec![0.0; 4];
872        let mut round_out = vec![0.0; 4];
873
874        floor_f32(&data, &mut floor_out).expect("floor_f32 failed");
875        ceil_f32(&data, &mut ceil_out).expect("ceil_f32 failed");
876        round_f32(&data, &mut round_out).expect("round_f32 failed");
877
878        assert_relative_eq!(floor_out[0], 1.0);
879        assert_relative_eq!(floor_out[1], 1.0);
880        assert_relative_eq!(floor_out[2], -2.0);
881        assert_relative_eq!(floor_out[3], -2.0);
882        assert_relative_eq!(ceil_out[0], 2.0);
883        assert_relative_eq!(ceil_out[1], 2.0);
884        assert_relative_eq!(ceil_out[2], -1.0);
885        assert_relative_eq!(ceil_out[3], -1.0);
886        assert_relative_eq!(round_out[0], 1.0);
887        assert_relative_eq!(round_out[1], 2.0);
888        assert_relative_eq!(round_out[2], -1.0);
889        assert_relative_eq!(round_out[3], -2.0);
890    }
891
892    #[test]
893    fn test_fract() {
894        let data = vec![1.3, 2.7, -1.3, -2.7];
895        let mut out = vec![0.0; 4];
896
897        fract_f32(&data, &mut out).expect("fract_f32 failed");
898
899        assert_relative_eq!(out[0], 0.3, epsilon = 1e-6);
900        assert_relative_eq!(out[1], 0.7, epsilon = 1e-6);
901        // For negative numbers, fract = x - floor(x), so -1.3 - (-2.0) = 0.7
902        assert_relative_eq!(out[2], 0.7, epsilon = 1e-6);
903        assert_relative_eq!(out[3], 0.3, epsilon = 1e-6);
904    }
905
906    #[test]
907    fn test_length_mismatch() {
908        let data = vec![1.0; 10];
909        let mut out = vec![0.0; 5];
910
911        assert!(sqrt_f32(&data, &mut out).is_err());
912    }
913}