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/// - `count >= 1024` 且 `avail != None` → SIMD 路径(wide::i64x4 向量批量解析)
97/// - `count < 1024` 或 `avail == None` → 标量降级
98pub fn batch_decode_integers(buf: &[u8], count: usize, avail: SimdAvailability) -> Vec<i64> {
99    if count >= SIMD_THRESHOLD && avail.is_available() {
100        simd_decode_integers(buf, count)
101    } else {
102        scalar_decode_integers(buf, count)
103    }
104}
105
106/// 标量批量整数解码
107pub fn scalar_decode_integers(buf: &[u8], count: usize) -> Vec<i64> {
108    let n = count.min(buf.len() / 8);
109    (0..n)
110        .map(|i| {
111            let offset = i * 8;
112            i64::from_le_bytes(buf[offset..offset + 8].try_into().unwrap())
113        })
114        .collect()
115}
116
117#[cfg(feature = "simd")]
118fn simd_decode_integers(buf: &[u8], count: usize) -> Vec<i64> {
119    use wide::i64x4;
120
121    let n = count.min(buf.len() / 8);
122    let mut result = Vec::with_capacity(n);
123
124    let chunk_count = n / 4;
125    let remainder = n % 4;
126
127    for chunk in 0..chunk_count {
128        let base = chunk * 4;
129        let v0 = i64::from_le_bytes(buf[base * 8..base * 8 + 8].try_into().unwrap());
130        let v1 = i64::from_le_bytes(buf[(base + 1) * 8..(base + 1) * 8 + 8].try_into().unwrap());
131        let v2 = i64::from_le_bytes(buf[(base + 2) * 8..(base + 2) * 8 + 8].try_into().unwrap());
132        let v3 = i64::from_le_bytes(buf[(base + 3) * 8..(base + 3) * 8 + 8].try_into().unwrap());
133
134        let vec = i64x4::from([v0, v1, v2, v3]);
135        let arr: [i64; 4] = vec.into();
136        result.extend_from_slice(&arr);
137    }
138
139    for i in 0..remainder {
140        let idx = chunk_count * 4 + i;
141        let offset = idx * 8;
142        result.push(i64::from_le_bytes(
143            buf[offset..offset + 8].try_into().unwrap(),
144        ));
145    }
146
147    result
148}
149
150#[cfg(not(feature = "simd"))]
151fn simd_decode_integers(buf: &[u8], count: usize) -> Vec<i64> {
152    scalar_decode_integers(buf, count)
153}
154
155// ============================================================================
156// 批量比较
157// ============================================================================
158
159/// 批量相等比较
160///
161/// 比较 `values` 中每个元素是否等于 `target`,返回布尔向量。
162///
163/// - `values.len() >= 1024` 且 `avail != None` → SIMD 路径
164/// - 否则 → 标量降级
165pub fn batch_compare_eq(values: &[i64], target: i64, avail: SimdAvailability) -> Vec<bool> {
166    if values.len() >= SIMD_THRESHOLD && avail.is_available() {
167        simd_compare_eq(values, target)
168    } else {
169        scalar_compare_eq(values, target)
170    }
171}
172
173/// 标量相等比较
174pub fn scalar_compare_eq(values: &[i64], target: i64) -> Vec<bool> {
175    values.iter().map(|&v| v == target).collect()
176}
177
178#[cfg(feature = "simd")]
179fn simd_compare_eq(values: &[i64], target: i64) -> Vec<bool> {
180    use wide::{i64x4, CmpEq};
181
182    let n = values.len();
183    let mut result = Vec::with_capacity(n);
184
185    let chunk_count = n / 4;
186    let remainder = n % 4;
187    let target_vec = i64x4::splat(target);
188
189    for chunk in 0..chunk_count {
190        let slice = &values[chunk * 4..chunk * 4 + 4];
191        let vec = i64x4::from([slice[0], slice[1], slice[2], slice[3]]);
192        let cmp = vec.cmp_eq(target_vec);
193        let mask: [i64; 4] = cmp.into();
194        for &m in &mask {
195            result.push(m != 0);
196        }
197    }
198
199    for i in 0..remainder {
200        let idx = chunk_count * 4 + i;
201        result.push(values[idx] == target);
202    }
203
204    result
205}
206
207#[cfg(not(feature = "simd"))]
208fn simd_compare_eq(values: &[i64], target: i64) -> Vec<bool> {
209    scalar_compare_eq(values, target)
210}
211
212/// 批量 IN 过滤
213///
214/// 判断 `values` 中每个元素是否在 `set` 中,返回布尔向量。
215///
216/// - `values.len() >= 1024` 且 `avail != None` → SIMD 路径
217/// - 否则 → 标量降级
218pub fn batch_compare_in(values: &[i64], set: &[i64], avail: SimdAvailability) -> Vec<bool> {
219    if values.len() >= SIMD_THRESHOLD && avail.is_available() {
220        simd_compare_in(values, set)
221    } else {
222        scalar_compare_in(values, set)
223    }
224}
225
226/// 标量 IN 过滤
227pub fn scalar_compare_in(values: &[i64], set: &[i64]) -> Vec<bool> {
228    values.iter().map(|&v| set.contains(&v)).collect()
229}
230
231#[cfg(feature = "simd")]
232fn simd_compare_in(values: &[i64], set: &[i64]) -> Vec<bool> {
233    use wide::{i64x4, CmpEq};
234
235    let n = values.len();
236    let mut result = Vec::with_capacity(n);
237
238    let chunk_count = n / 4;
239    let remainder = n % 4;
240
241    for chunk in 0..chunk_count {
242        let slice = &values[chunk * 4..chunk * 4 + 4];
243        let vec = i64x4::from([slice[0], slice[1], slice[2], slice[3]]);
244
245        let mut any_match = [false; 4];
246        for &s in set {
247            let target_vec = i64x4::splat(s);
248            let cmp = vec.cmp_eq(target_vec);
249            let mask: [i64; 4] = cmp.into();
250            for j in 0..4 {
251                if mask[j] != 0 {
252                    any_match[j] = true;
253                }
254            }
255        }
256        result.extend_from_slice(&any_match);
257    }
258
259    for i in 0..remainder {
260        let idx = chunk_count * 4 + i;
261        result.push(set.contains(&values[idx]));
262    }
263
264    result
265}
266
267#[cfg(not(feature = "simd"))]
268fn simd_compare_in(values: &[i64], set: &[i64]) -> Vec<bool> {
269    scalar_compare_in(values, set)
270}
271
272// ============================================================================
273// 单元测试
274// ============================================================================
275
276#[cfg(test)]
277mod tests {
278    use super::*;
279
280    #[test]
281    fn test_simd_availability_is_available() {
282        assert!(SimdAvailability::Avx2.is_available());
283        assert!(SimdAvailability::Avx.is_available());
284        assert!(SimdAvailability::Sse2.is_available());
285        assert!(SimdAvailability::Neon.is_available());
286        assert!(!SimdAvailability::None.is_available());
287    }
288
289    #[test]
290    fn test_detect_returns_cached() {
291        let d1 = detect();
292        let d2 = detect();
293        assert_eq!(d1, d2);
294    }
295
296    #[test]
297    fn test_scalar_decode_integers() {
298        let values: Vec<i64> = vec![1, 2, 3, 4, 5];
299        let mut buf = Vec::new();
300        for v in &values {
301            buf.extend_from_slice(&v.to_le_bytes());
302        }
303        let result = scalar_decode_integers(&buf, 5);
304        assert_eq!(result, values);
305    }
306
307    #[test]
308    fn test_batch_decode_integers_small_count() {
309        let values: Vec<i64> = vec![1, 2, 3];
310        let mut buf = Vec::new();
311        for v in &values {
312            buf.extend_from_slice(&v.to_le_bytes());
313        }
314        let result = batch_decode_integers(&buf, 3, SimdAvailability::Avx2);
315        assert_eq!(result, values);
316    }
317
318    #[test]
319    fn test_batch_decode_integers_large_count() {
320        let n: usize = 2000;
321        let values: Vec<i64> = (0..n as i64).map(|i| i * 2 - 1).collect();
322        let mut buf = Vec::new();
323        for v in &values {
324            buf.extend_from_slice(&v.to_le_bytes());
325        }
326        let avail = detect();
327        let result = batch_decode_integers(&buf, n, avail);
328        assert_eq!(result, values);
329    }
330
331    #[test]
332    fn test_batch_decode_integers_none_avail() {
333        let n: usize = 2000;
334        let values: Vec<i64> = (0..n as i64).collect();
335        let mut buf = Vec::new();
336        for v in &values {
337            buf.extend_from_slice(&v.to_le_bytes());
338        }
339        let result = batch_decode_integers(&buf, n, SimdAvailability::None);
340        assert_eq!(result, values);
341    }
342
343    #[test]
344    fn test_scalar_compare_eq() {
345        let values = vec![1, 2, 3, 4, 5, 3, 3];
346        let result = scalar_compare_eq(&values, 3);
347        assert_eq!(result, vec![false, false, true, false, false, true, true]);
348    }
349
350    #[test]
351    fn test_batch_compare_eq_small() {
352        let values = vec![1, 2, 3, 4, 5];
353        let result = batch_compare_eq(&values, 3, SimdAvailability::Avx2);
354        assert_eq!(result, vec![false, false, true, false, false]);
355    }
356
357    #[test]
358    fn test_batch_compare_eq_large() {
359        let n: usize = 2000;
360        let values: Vec<i64> = (0..n as i64).collect();
361        let target = 500_i64;
362        let avail = detect();
363        let result = batch_compare_eq(&values, target, avail);
364        assert_eq!(result.len(), n);
365        assert!(result[500]);
366        assert!(!result[499]);
367        assert!(!result[501]);
368    }
369
370    #[test]
371    fn test_scalar_compare_in() {
372        let values = vec![1, 2, 3, 4, 5];
373        let set = vec![2, 4];
374        let result = scalar_compare_in(&values, &set);
375        assert_eq!(result, vec![false, true, false, true, false]);
376    }
377
378    #[test]
379    fn test_batch_compare_in_small() {
380        let values = vec![1, 2, 3, 4, 5];
381        let set = vec![2, 4];
382        let result = batch_compare_in(&values, &set, SimdAvailability::Avx2);
383        assert_eq!(result, vec![false, true, false, true, false]);
384    }
385
386    #[test]
387    fn test_batch_compare_in_large() {
388        let n: usize = 2000;
389        let values: Vec<i64> = (0..n as i64).collect();
390        let set: Vec<i64> = vec![100, 500, 1500];
391        let avail = detect();
392        let result = batch_compare_in(&values, &set, avail);
393        assert_eq!(result.len(), n);
394        assert!(result[100]);
395        assert!(result[500]);
396        assert!(result[1500]);
397        assert!(!result[200]);
398    }
399
400    #[test]
401    fn test_batch_compare_eq_none_avail() {
402        let n: usize = 2000;
403        let values: Vec<i64> = (0..n as i64).collect();
404        let result = batch_compare_eq(&values, 500, SimdAvailability::None);
405        assert_eq!(result.len(), n);
406        assert!(result[500]);
407    }
408
409    #[test]
410    fn test_batch_compare_in_empty_set() {
411        let values = vec![1, 2, 3];
412        let set: Vec<i64> = vec![];
413        let result = batch_compare_in(&values, &set, SimdAvailability::Avx2);
414        assert_eq!(result, vec![false, false, false]);
415    }
416
417    #[test]
418    fn test_batch_decode_integers_count_exceeds_buf() {
419        let values: Vec<i64> = vec![1, 2, 3];
420        let mut buf = Vec::new();
421        for v in &values {
422            buf.extend_from_slice(&v.to_le_bytes());
423        }
424        let result = batch_decode_integers(&buf, 100, SimdAvailability::None);
425        assert_eq!(result, values);
426    }
427
428    #[test]
429    fn test_batch_decode_integers_empty() {
430        let result = batch_decode_integers(&[], 0, SimdAvailability::Avx2);
431        assert!(result.is_empty());
432    }
433
434    #[test]
435    fn test_simd_threshold_constant() {
436        assert_eq!(SIMD_THRESHOLD, 1024);
437    }
438
439    #[test]
440    fn test_batch_compare_eq_boundary_1023() {
441        let n = 1023;
442        let values: Vec<i64> = vec![42; n];
443        let result = batch_compare_eq(&values, 42, SimdAvailability::Avx2);
444        assert!(result.iter().all(|&b| b));
445    }
446
447    #[test]
448    fn test_batch_compare_eq_boundary_1024() {
449        let n = 1024;
450        let values: Vec<i64> = vec![42; n];
451        let avail = detect();
452        let result = batch_compare_eq(&values, 42, avail);
453        assert!(result.iter().all(|&b| b));
454    }
455
456    #[test]
457    fn test_batch_compare_eq_boundary_1025() {
458        let n = 1025;
459        let values: Vec<i64> = vec![42; n];
460        let avail = detect();
461        let result = batch_compare_eq(&values, 42, avail);
462        assert!(result.iter().all(|&b| b));
463    }
464}