Skip to main content

ruprim_host/simd/portable/
comparison_ops.rs

1use super::*;
2
3// ============================================================================
4// f32 comparison ops
5// ============================================================================
6
7#[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/// Comparison kernel using simple loops that LLVM autovectorizes.
32///
33/// Autovectorization outperforms explicit SIMD here because comparisons
34/// produce u8 output from f32 input (4:1 size ratio). LLVM batches 16+
35/// comparisons and packs results into a single wide vector store, while
36/// explicit SIMD (store_as_bool) writes only `lanes` bytes per iteration.
37#[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/// Scalar comparison kernel using simple loops that LLVM autovectorizes.
98/// See `cmp_f32_seq` for rationale.
99#[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