Skip to main content

sz_orm_core/
simd.rs

1//! v3.2.0 SIMD 加速 — 批量整数解码 + 列比较
2//!
3//! 使用 `wide` crate(stable Rust)提供 SIMD 向量化加速:
4//! - `batch_decode_integers`:批量整数解码(i64x4 向量并行)
5//! - `batch_compare_eq`:批量相等比较(i64x4 并行比较)
6//! - `batch_compare_in`:批量 IN 过滤(向量比较 + 布尔掩码)
7//!
8//! # 自动降级
9//!
10//! - count < 1024 → 标量路径(无 SIMD 开销)
11//! - `SimdAvailability::None` → 标量路径
12//! - WASM 目标 → `detect()` 返回 `None`
13
14use std::sync::OnceLock;
15
16/// SIMD 可用性枚举
17#[derive(Debug, Clone, Copy, PartialEq, Eq)]
18pub enum SimdAvailability {
19    /// AVX2(256-bit,4×i64)
20    Avx2,
21    /// AVX(256-bit float,128-bit integer)
22    Avx,
23    /// SSE2(128-bit,2×i64)
24    Sse2,
25    /// ARM NEON(128-bit)
26    Neon,
27    /// 无 SIMD 可用
28    None,
29}
30
31impl SimdAvailability {
32    /// 是否有 SIMD 可用
33    pub fn is_available(&self) -> bool {
34        *self != SimdAvailability::None
35    }
36}
37
38static DETECTED: OnceLock<SimdAvailability> = OnceLock::new();
39
40/// 检测当前 CPU 的 SIMD 可用性(首次检测后缓存)
41pub fn detect() -> SimdAvailability {
42    *DETECTED.get_or_init(detect_impl)
43}
44
45#[cfg(target_arch = "x86_64")]
46fn detect_impl() -> SimdAvailability {
47    if is_x86_feature_detected!("avx2") {
48        SimdAvailability::Avx2
49    } else if is_x86_feature_detected!("avx") {
50        SimdAvailability::Avx
51    } else if is_x86_feature_detected!("sse2") {
52        SimdAvailability::Sse2
53    } else {
54        SimdAvailability::None
55    }
56}
57
58#[cfg(target_arch = "x86")]
59fn detect_impl() -> SimdAvailability {
60    if is_x86_feature_detected!("avx2") {
61        SimdAvailability::Avx2
62    } else if is_x86_feature_detected!("avx") {
63        SimdAvailability::Avx
64    } else if is_x86_feature_detected!("sse2") {
65        SimdAvailability::Sse2
66    } else {
67        SimdAvailability::None
68    }
69}
70
71#[cfg(target_arch = "aarch64")]
72fn detect_impl() -> SimdAvailability {
73    if std::arch::is_aarch64_feature_detected!("neon") {
74        SimdAvailability::Neon
75    } else {
76        SimdAvailability::None
77    }
78}
79
80#[cfg(not(any(target_arch = "x86_64", target_arch = "x86", target_arch = "aarch64")))]
81fn detect_impl() -> SimdAvailability {
82    SimdAvailability::None
83}
84
85/// SIMD 批量处理的最低元素数量阈值
86pub const SIMD_THRESHOLD: usize = 1024;
87
88// ============================================================================
89// 批量整数解码
90// ============================================================================
91
92/// 批量整数解码
93///
94/// 将 `buf` 中的 `count` 个 i64(小端字节序列,每 8 字节一个)解码为 `Vec<i64>`。
95///
96/// 始终使用标量路径(编译器自动向量化已优于显式 SIMD,实测验证 2026-08-19)。
97/// `avail` 参数保留用于 API 兼容性。
98pub fn batch_decode_integers(buf: &[u8], count: usize, _avail: SimdAvailability) -> Vec<i64> {
99    scalar_decode_integers(buf, count)
100}
101
102/// 标量批量整数解码
103pub fn scalar_decode_integers(buf: &[u8], count: usize) -> Vec<i64> {
104    let n = count.min(buf.len() / 8);
105    (0..n)
106        .map(|i| {
107            let offset = i * 8;
108            i64::from_le_bytes(buf[offset..offset + 8].try_into().unwrap())
109        })
110        .collect()
111}
112
113// ============================================================================
114// 批量比较
115// ============================================================================
116
117/// 批量相等比较
118///
119/// 比较 `values` 中每个元素是否等于 `target`,返回布尔向量。
120///
121/// 始终使用标量路径(编译器自动向量化已优于显式 SIMD,实测验证 2026-08-19)。
122/// `avail` 参数保留用于 API 兼容性。
123pub fn batch_compare_eq(values: &[i64], target: i64, _avail: SimdAvailability) -> Vec<bool> {
124    scalar_compare_eq(values, target)
125}
126
127/// 标量相等比较
128pub fn scalar_compare_eq(values: &[i64], target: i64) -> Vec<bool> {
129    values.iter().map(|&v| v == target).collect()
130}
131
132/// 批量 IN 过滤
133///
134/// 判断 `values` 中每个元素是否在 `set` 中,返回布尔向量。
135///
136/// 当 `set.len() >= 8` 时使用 `HashSet` 做 O(1) 查找(算法级优化,远超 SIMD)。
137/// 小集合直接线性扫描(避免 HashSet 建表开销)。
138pub fn batch_compare_in(values: &[i64], set: &[i64], _avail: SimdAvailability) -> Vec<bool> {
139    if set.len() >= 8 {
140        let hash_set: std::collections::HashSet<i64> = set.iter().copied().collect();
141        values.iter().map(|&v| hash_set.contains(&v)).collect()
142    } else {
143        scalar_compare_in(values, set)
144    }
145}
146
147/// 标量 IN 过滤
148pub fn scalar_compare_in(values: &[i64], set: &[i64]) -> Vec<bool> {
149    values.iter().map(|&v| set.contains(&v)).collect()
150}
151
152// ============================================================================
153// 单元测试
154// ============================================================================
155
156#[cfg(test)]
157mod tests {
158    use super::*;
159
160    #[test]
161    fn test_simd_availability_is_available() {
162        assert!(SimdAvailability::Avx2.is_available());
163        assert!(SimdAvailability::Avx.is_available());
164        assert!(SimdAvailability::Sse2.is_available());
165        assert!(SimdAvailability::Neon.is_available());
166        assert!(!SimdAvailability::None.is_available());
167    }
168
169    #[test]
170    fn test_detect_returns_cached() {
171        let d1 = detect();
172        let d2 = detect();
173        assert_eq!(d1, d2);
174    }
175
176    #[test]
177    fn test_scalar_decode_integers() {
178        let values: Vec<i64> = vec![1, 2, 3, 4, 5];
179        let mut buf = Vec::new();
180        for v in &values {
181            buf.extend_from_slice(&v.to_le_bytes());
182        }
183        let result = scalar_decode_integers(&buf, 5);
184        assert_eq!(result, values);
185    }
186
187    #[test]
188    fn test_batch_decode_integers_small_count() {
189        let values: Vec<i64> = vec![1, 2, 3];
190        let mut buf = Vec::new();
191        for v in &values {
192            buf.extend_from_slice(&v.to_le_bytes());
193        }
194        let result = batch_decode_integers(&buf, 3, SimdAvailability::Avx2);
195        assert_eq!(result, values);
196    }
197
198    #[test]
199    fn test_batch_decode_integers_large_count() {
200        let n: usize = 2000;
201        let values: Vec<i64> = (0..n as i64).map(|i| i * 2 - 1).collect();
202        let mut buf = Vec::new();
203        for v in &values {
204            buf.extend_from_slice(&v.to_le_bytes());
205        }
206        let avail = detect();
207        let result = batch_decode_integers(&buf, n, avail);
208        assert_eq!(result, values);
209    }
210
211    #[test]
212    fn test_batch_decode_integers_none_avail() {
213        let n: usize = 2000;
214        let values: Vec<i64> = (0..n as i64).collect();
215        let mut buf = Vec::new();
216        for v in &values {
217            buf.extend_from_slice(&v.to_le_bytes());
218        }
219        let result = batch_decode_integers(&buf, n, SimdAvailability::None);
220        assert_eq!(result, values);
221    }
222
223    #[test]
224    fn test_scalar_compare_eq() {
225        let values = vec![1, 2, 3, 4, 5, 3, 3];
226        let result = scalar_compare_eq(&values, 3);
227        assert_eq!(result, vec![false, false, true, false, false, true, true]);
228    }
229
230    #[test]
231    fn test_batch_compare_eq_small() {
232        let values = vec![1, 2, 3, 4, 5];
233        let result = batch_compare_eq(&values, 3, SimdAvailability::Avx2);
234        assert_eq!(result, vec![false, false, true, false, false]);
235    }
236
237    #[test]
238    fn test_batch_compare_eq_large() {
239        let n: usize = 2000;
240        let values: Vec<i64> = (0..n as i64).collect();
241        let target = 500_i64;
242        let avail = detect();
243        let result = batch_compare_eq(&values, target, avail);
244        assert_eq!(result.len(), n);
245        assert!(result[500]);
246        assert!(!result[499]);
247        assert!(!result[501]);
248    }
249
250    #[test]
251    fn test_scalar_compare_in() {
252        let values = vec![1, 2, 3, 4, 5];
253        let set = vec![2, 4];
254        let result = scalar_compare_in(&values, &set);
255        assert_eq!(result, vec![false, true, false, true, false]);
256    }
257
258    #[test]
259    fn test_batch_compare_in_small() {
260        let values = vec![1, 2, 3, 4, 5];
261        let set = vec![2, 4];
262        let result = batch_compare_in(&values, &set, SimdAvailability::Avx2);
263        assert_eq!(result, vec![false, true, false, true, false]);
264    }
265
266    #[test]
267    fn test_batch_compare_in_large() {
268        let n: usize = 2000;
269        let values: Vec<i64> = (0..n as i64).collect();
270        let set: Vec<i64> = vec![100, 500, 1500];
271        let avail = detect();
272        let result = batch_compare_in(&values, &set, avail);
273        assert_eq!(result.len(), n);
274        assert!(result[100]);
275        assert!(result[500]);
276        assert!(result[1500]);
277        assert!(!result[200]);
278    }
279
280    #[test]
281    fn test_batch_compare_eq_none_avail() {
282        let n: usize = 2000;
283        let values: Vec<i64> = (0..n as i64).collect();
284        let result = batch_compare_eq(&values, 500, SimdAvailability::None);
285        assert_eq!(result.len(), n);
286        assert!(result[500]);
287    }
288
289    #[test]
290    fn test_batch_compare_in_empty_set() {
291        let values = vec![1, 2, 3];
292        let set: Vec<i64> = vec![];
293        let result = batch_compare_in(&values, &set, SimdAvailability::Avx2);
294        assert_eq!(result, vec![false, false, false]);
295    }
296
297    #[test]
298    fn test_batch_decode_integers_count_exceeds_buf() {
299        let values: Vec<i64> = vec![1, 2, 3];
300        let mut buf = Vec::new();
301        for v in &values {
302            buf.extend_from_slice(&v.to_le_bytes());
303        }
304        let result = batch_decode_integers(&buf, 100, SimdAvailability::None);
305        assert_eq!(result, values);
306    }
307
308    #[test]
309    fn test_batch_decode_integers_empty() {
310        let result = batch_decode_integers(&[], 0, SimdAvailability::Avx2);
311        assert!(result.is_empty());
312    }
313
314    #[test]
315    fn test_simd_threshold_constant() {
316        assert_eq!(SIMD_THRESHOLD, 1024);
317    }
318
319    #[test]
320    fn test_batch_compare_eq_boundary_1023() {
321        let n = 1023;
322        let values: Vec<i64> = vec![42; n];
323        let result = batch_compare_eq(&values, 42, SimdAvailability::Avx2);
324        assert!(result.iter().all(|&b| b));
325    }
326
327    #[test]
328    fn test_batch_compare_eq_boundary_1024() {
329        let n = 1024;
330        let values: Vec<i64> = vec![42; n];
331        let avail = detect();
332        let result = batch_compare_eq(&values, 42, avail);
333        assert!(result.iter().all(|&b| b));
334    }
335
336    #[test]
337    fn test_batch_compare_eq_boundary_1025() {
338        let n = 1025;
339        let values: Vec<i64> = vec![42; n];
340        let avail = detect();
341        let result = batch_compare_eq(&values, 42, avail);
342        assert!(result.iter().all(|&b| b));
343    }
344}