1#![allow(unsafe_code)]
50
51use crate::error::{AlgorithmError, Result};
52
53fn 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#[cfg(target_arch = "aarch64")]
76mod neon_impl {
77 use std::arch::aarch64::*;
78
79 #[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 #[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 #[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 #[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 #[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 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 #[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 #[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
209mod 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
248pub fn sqrt_f32(data: &[f32], out: &mut [f32]) -> Result<()> {
260 validate_unary(data, out)?;
261
262 #[cfg(target_arch = "aarch64")]
263 {
264 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
278pub fn ln_f32(data: &[f32], out: &mut [f32]) -> Result<()> {
287 validate_unary(data, out)?;
288
289 #[cfg(target_arch = "aarch64")]
290 {
291 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
305pub fn log10_f32(data: &[f32], out: &mut [f32]) -> Result<()> {
311 validate_unary(data, out)?;
312
313 #[cfg(target_arch = "aarch64")]
314 {
315 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
334pub fn log2_f32(data: &[f32], out: &mut [f32]) -> Result<()> {
340 validate_unary(data, out)?;
341
342 #[cfg(target_arch = "aarch64")]
343 {
344 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
363pub fn exp_f32(data: &[f32], out: &mut [f32]) -> Result<()> {
372 validate_unary(data, out)?;
373
374 #[cfg(target_arch = "aarch64")]
375 {
376 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
390pub 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
401pub 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
418pub 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
429pub 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
440pub 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
451pub 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
462pub 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
473pub 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
484pub 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
500pub 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
511pub 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
522pub 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
533pub fn abs_f32(data: &[f32], out: &mut [f32]) -> Result<()> {
541 validate_unary(data, out)?;
542
543 #[cfg(target_arch = "aarch64")]
544 {
545 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
559pub fn floor_f32(data: &[f32], out: &mut [f32]) -> Result<()> {
567 validate_unary(data, out)?;
568
569 #[cfg(target_arch = "aarch64")]
570 {
571 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
585pub fn ceil_f32(data: &[f32], out: &mut [f32]) -> Result<()> {
593 validate_unary(data, out)?;
594
595 #[cfg(target_arch = "aarch64")]
596 {
597 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
611pub fn round_f32(data: &[f32], out: &mut [f32]) -> Result<()> {
619 validate_unary(data, out)?;
620
621 #[cfg(target_arch = "aarch64")]
622 {
623 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
637pub fn fract_f32(data: &[f32], out: &mut [f32]) -> Result<()> {
643 validate_unary(data, out)?;
644 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 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 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 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}