ruprim_host/simd/portable/
comparison_ops.rs1use super::*;
2
3#[derive(Clone, Copy)]
8pub enum CmpOp {
9 Gt,
10 Ge,
11 Lt,
12 Le,
13 Eq,
14 Ne,
15}
16
17#[inline]
18pub fn cmp_f32(a: &[f32], b: &[f32], out: &mut [u8], op: CmpOp) {
19 debug_assert_eq!(a.len(), b.len());
20 debug_assert_eq!(a.len(), out.len());
21
22 #[cfg(feature = "rayon")]
23 if a.len() >= PARALLEL_THRESHOLD {
24 cmp_f32_par(a, b, out, op);
25 return;
26 }
27
28 cmp_f32_seq(a, b, out, op);
29}
30
31#[inline]
38fn cmp_f32_seq(a: &[f32], b: &[f32], out: &mut [u8], op: CmpOp) {
39 match op {
40 CmpOp::Gt => {
41 for ((a, b), o) in a.iter().zip(b).zip(out.iter_mut()) {
42 *o = (*a > *b) as u8;
43 }
44 }
45 CmpOp::Ge => {
46 for ((a, b), o) in a.iter().zip(b).zip(out.iter_mut()) {
47 *o = (*a >= *b) as u8;
48 }
49 }
50 CmpOp::Lt => {
51 for ((a, b), o) in a.iter().zip(b).zip(out.iter_mut()) {
52 *o = (*a < *b) as u8;
53 }
54 }
55 CmpOp::Le => {
56 for ((a, b), o) in a.iter().zip(b).zip(out.iter_mut()) {
57 *o = (*a <= *b) as u8;
58 }
59 }
60 CmpOp::Eq => {
61 for ((a, b), o) in a.iter().zip(b).zip(out.iter_mut()) {
62 *o = (*a == *b) as u8;
63 }
64 }
65 CmpOp::Ne => {
66 for ((a, b), o) in a.iter().zip(b).zip(out.iter_mut()) {
67 *o = (*a != *b) as u8;
68 }
69 }
70 }
71}
72
73#[cfg(feature = "rayon")]
74fn cmp_f32_par(a: &[f32], b: &[f32], out: &mut [u8], op: CmpOp) {
75 out.par_chunks_mut(CHUNK_SIZE)
76 .enumerate()
77 .for_each(|(chunk_idx, out_chunk)| {
78 let start = chunk_idx * CHUNK_SIZE;
79 let end = (start + CHUNK_SIZE).min(a.len());
80 cmp_f32_seq(&a[start..end], &b[start..end], out_chunk, op);
81 });
82}
83
84#[inline]
85pub fn cmp_scalar_f32(a: &[f32], scalar: f32, out: &mut [u8], op: CmpOp) {
86 debug_assert_eq!(a.len(), out.len());
87
88 #[cfg(feature = "rayon")]
89 if a.len() >= PARALLEL_THRESHOLD {
90 cmp_scalar_f32_par(a, scalar, out, op);
91 return;
92 }
93
94 cmp_scalar_f32_seq(a, scalar, out, op);
95}
96
97#[inline]
100fn cmp_scalar_f32_seq(a: &[f32], scalar: f32, out: &mut [u8], op: CmpOp) {
101 match op {
102 CmpOp::Gt => {
103 for (a, o) in a.iter().zip(out.iter_mut()) {
104 *o = (*a > scalar) as u8;
105 }
106 }
107 CmpOp::Ge => {
108 for (a, o) in a.iter().zip(out.iter_mut()) {
109 *o = (*a >= scalar) as u8;
110 }
111 }
112 CmpOp::Lt => {
113 for (a, o) in a.iter().zip(out.iter_mut()) {
114 *o = (*a < scalar) as u8;
115 }
116 }
117 CmpOp::Le => {
118 for (a, o) in a.iter().zip(out.iter_mut()) {
119 *o = (*a <= scalar) as u8;
120 }
121 }
122 CmpOp::Eq => {
123 for (a, o) in a.iter().zip(out.iter_mut()) {
124 *o = (*a == scalar) as u8;
125 }
126 }
127 CmpOp::Ne => {
128 for (a, o) in a.iter().zip(out.iter_mut()) {
129 *o = (*a != scalar) as u8;
130 }
131 }
132 }
133}
134
135#[cfg(feature = "rayon")]
136fn cmp_scalar_f32_par(a: &[f32], scalar: f32, out: &mut [u8], op: CmpOp) {
137 out.par_chunks_mut(CHUNK_SIZE)
138 .enumerate()
139 .for_each(|(chunk_idx, out_chunk)| {
140 let start = chunk_idx * CHUNK_SIZE;
141 let end = (start + CHUNK_SIZE).min(a.len());
142 cmp_scalar_f32_seq(&a[start..end], scalar, out_chunk, op);
143 });
144}
145