Skip to main content

sc_neurocore_engine/simd/
mod.rs

1// SPDX-License-Identifier: AGPL-3.0-or-later
2// Commercial license available
3// © Concepts 1996–2026 Miroslav Šotek. All rights reserved.
4// © Code 2020–2026 Miroslav Šotek. All rights reserved.
5// ORCID: 0009-0009-3560-0851
6// Contact: www.anulum.li | protoscience@anulum.li
7// SC-NeuroCore — SIMD Popcount Dispatch
8
9//! # SIMD Popcount Dispatch
10//!
11//! Runtime CPU-feature dispatch for packed-bit popcount kernels.
12//! Supported backends: AVX-512, AVX2, ARM NEON, ARM SVE, RISC-V RVV.
13
14use rand::Rng;
15
16pub mod avx2;
17pub mod avx512;
18pub mod neon;
19pub mod rvv;
20pub mod sve;
21
22/// Pack u8 bits into u64 words using the best available SIMD path.
23pub fn pack_dispatch(bits: &[u8]) -> crate::bitstream::BitStreamTensor {
24    let length = bits.len();
25
26    #[cfg(target_arch = "x86_64")]
27    {
28        if is_x86_feature_detected!("avx512bw") {
29            // SAFETY: Guarded by runtime feature detection.
30            let data = unsafe { avx512::pack_avx512(bits) };
31            return crate::bitstream::BitStreamTensor { data, length };
32        }
33        if is_x86_feature_detected!("avx2") {
34            // SAFETY: Guarded by runtime feature detection.
35            let data = unsafe { avx2::pack_avx2(bits) };
36            return crate::bitstream::BitStreamTensor { data, length };
37        }
38    }
39
40    #[cfg(all(target_arch = "aarch64", target_feature = "sve"))]
41    {
42        // SAFETY: SVE target feature is compile-time guaranteed.
43        let data = unsafe { sve::pack_sve(bits) };
44        return crate::bitstream::BitStreamTensor { data, length };
45    }
46
47    crate::bitstream::pack_fast(bits)
48}
49
50/// Count set bits in packed `u64` words using the best available SIMD path.
51pub fn popcount_dispatch(data: &[u64]) -> u64 {
52    #[cfg(target_arch = "x86_64")]
53    {
54        if is_x86_feature_detected!("avx512vpopcntdq") {
55            // SAFETY: Guarded by runtime feature detection.
56            return unsafe { avx512::popcount_avx512(data) };
57        }
58        if is_x86_feature_detected!("avx2") {
59            // SAFETY: Guarded by runtime feature detection.
60            return unsafe { avx2::popcount_avx2(data) };
61        }
62    }
63
64    #[cfg(target_arch = "aarch64")]
65    {
66        #[cfg(target_feature = "sve")]
67        {
68            // SAFETY: SVE target feature is compile-time guaranteed.
69            return unsafe { sve::popcount_sve(data) };
70        }
71        #[cfg(not(target_feature = "sve"))]
72        {
73            // SAFETY: NEON is baseline on aarch64 targets.
74            return unsafe { neon::popcount_neon(data) };
75        }
76    }
77
78    #[cfg(all(target_arch = "riscv64", target_feature = "v"))]
79    {
80        // SAFETY: RVV target feature is compile-time guaranteed.
81        return unsafe { rvv::popcount_rvv(data) };
82    }
83
84    crate::bitstream::popcount_words_portable(data)
85}
86
87/// Fused AND+popcount dispatch using the best available SIMD path.
88pub fn fused_and_popcount_dispatch(a: &[u64], b: &[u64]) -> u64 {
89    let len = a.len().min(b.len());
90    let a = &a[..len];
91    let b = &b[..len];
92
93    #[cfg(target_arch = "x86_64")]
94    {
95        if is_x86_feature_detected!("avx512vpopcntdq") {
96            // SAFETY: Guarded by runtime feature detection.
97            return unsafe { avx512::fused_and_popcount_avx512(a, b) };
98        }
99        if is_x86_feature_detected!("avx2") {
100            // SAFETY: Guarded by runtime feature detection.
101            return unsafe { avx2::fused_and_popcount_avx2(a, b) };
102        }
103    }
104
105    #[cfg(target_arch = "aarch64")]
106    {
107        #[cfg(target_feature = "sve")]
108        {
109            return unsafe { sve::fused_and_popcount_sve(a, b) };
110        }
111    }
112
113    #[cfg(all(target_arch = "riscv64", target_feature = "v"))]
114    {
115        return unsafe { rvv::fused_and_popcount_rvv(a, b) };
116    }
117
118    let mut total = 0_u64;
119    let (chunks_a, remainder_a) = a.as_chunks::<4>();
120    let (chunks_b, remainder_b) = b.as_chunks::<4>();
121    for (ca, cb) in chunks_a.iter().zip(chunks_b) {
122        total += (ca[0] & cb[0]).count_ones() as u64;
123        total += (ca[1] & cb[1]).count_ones() as u64;
124        total += (ca[2] & cb[2]).count_ones() as u64;
125        total += (ca[3] & cb[3]).count_ones() as u64;
126    }
127    total += remainder_a
128        .iter()
129        .zip(remainder_b)
130        .map(|(&wa, &wb)| (wa & wb).count_ones() as u64)
131        .sum::<u64>();
132    total
133}
134
135/// Fused XOR+popcount dispatch using the best available SIMD path.
136pub fn fused_xor_popcount_dispatch(a: &[u64], b: &[u64]) -> u64 {
137    let len = a.len().min(b.len());
138    let a = &a[..len];
139    let b = &b[..len];
140
141    #[cfg(target_arch = "x86_64")]
142    {
143        if is_x86_feature_detected!("avx512vpopcntdq") {
144            // SAFETY: Guarded by runtime feature detection.
145            return unsafe { avx512::fused_xor_popcount_avx512(a, b) };
146        }
147        if is_x86_feature_detected!("avx2") {
148            // SAFETY: Guarded by runtime feature detection.
149            return unsafe { avx2::fused_xor_popcount_avx2(a, b) };
150        }
151        if is_x86_feature_detected!("avx") {
152            let mut total = 0_u64;
153            let (chunks_a, remainder_a) = a.as_chunks::<16>();
154            let (chunks_b, remainder_b) = b.as_chunks::<16>();
155            for (ca, cb) in chunks_a.iter().zip(chunks_b) {
156                for i in 0..16 {
157                    total += (ca[i] ^ cb[i]).count_ones() as u64;
158                }
159            }
160            total += remainder_a
161                .iter()
162                .zip(remainder_b)
163                .map(|(&wa, &wb)| (wa ^ wb).count_ones() as u64)
164                .sum::<u64>();
165            return total;
166        }
167    }
168
169    #[cfg(target_arch = "aarch64")]
170    {
171        #[cfg(target_feature = "sve")]
172        {
173            return unsafe { sve::fused_xor_popcount_sve(a, b) };
174        }
175    }
176
177    #[cfg(all(target_arch = "riscv64", target_feature = "v"))]
178    {
179        return unsafe { rvv::fused_xor_popcount_rvv(a, b) };
180    }
181
182    let mut total = 0_u64;
183    let (chunks_a, remainder_a) = a.as_chunks::<4>();
184    let (chunks_b, remainder_b) = b.as_chunks::<4>();
185    for (ca, cb) in chunks_a.iter().zip(chunks_b) {
186        total += (ca[0] ^ cb[0]).count_ones() as u64;
187        total += (ca[1] ^ cb[1]).count_ones() as u64;
188        total += (ca[2] ^ cb[2]).count_ones() as u64;
189        total += (ca[3] ^ cb[3]).count_ones() as u64;
190    }
191    total += remainder_a
192        .iter()
193        .zip(remainder_b)
194        .map(|(&wa, &wb)| (wa ^ wb).count_ones() as u64)
195        .sum::<u64>();
196    total
197}
198
199// --- f64 dispatch functions ---
200
201/// Dot product of two f64 slices using the best available SIMD path.
202pub fn dot_f64_dispatch(a: &[f64], b: &[f64]) -> f64 {
203    #[cfg(target_arch = "x86_64")]
204    {
205        if is_x86_feature_detected!("avx512f") {
206            return unsafe { avx512::dot_f64_avx512(a, b) };
207        }
208        if is_x86_feature_detected!("avx2") && is_x86_feature_detected!("fma") {
209            return unsafe { avx2::dot_f64_avx2(a, b) };
210        }
211        if is_x86_feature_detected!("avx") {
212            return unsafe { avx2::dot_f64_avx(a, b) };
213        }
214        if is_x86_feature_detected!("sse2") {
215            let len = a.len().min(b.len());
216            let mut sum = 0.0_f64;
217            let (chunks_a, remainder_a) = a[..len].as_chunks::<4>();
218            let (chunks_b, remainder_b) = b[..len].as_chunks::<4>();
219            for (ca, cb) in chunks_a.iter().zip(chunks_b) {
220                sum += ca[0] * cb[0] + ca[1] * cb[1] + ca[2] * cb[2] + ca[3] * cb[3];
221            }
222            sum += remainder_a
223                .iter()
224                .zip(remainder_b)
225                .map(|(x, y)| x * y)
226                .sum::<f64>();
227            return sum;
228        }
229    }
230
231    #[cfg(target_arch = "aarch64")]
232    {
233        return unsafe { neon::dot_f64_neon(a, b) };
234    }
235
236    let len = a.len().min(b.len());
237    a[..len].iter().zip(&b[..len]).map(|(&x, &y)| x * y).sum()
238}
239
240/// Maximum of f64 slice using the best available SIMD path.
241pub fn max_f64_dispatch(a: &[f64]) -> f64 {
242    #[cfg(target_arch = "x86_64")]
243    {
244        if is_x86_feature_detected!("avx512f") {
245            return unsafe { avx512::max_f64_avx512(a) };
246        }
247        if is_x86_feature_detected!("avx2") {
248            return unsafe { avx2::max_f64_avx2(a) };
249        }
250        if is_x86_feature_detected!("avx") {
251            return unsafe { avx2::max_f64_avx(a) };
252        }
253        if is_x86_feature_detected!("sse2") {
254            let mut m = f64::NEG_INFINITY;
255            let (chunks, remainder) = a.as_chunks::<4>();
256            for c in chunks {
257                m = m.max(c[0].max(c[1]).max(c[2].max(c[3])));
258            }
259            for &v in remainder {
260                m = m.max(v);
261            }
262            return m;
263        }
264    }
265
266    #[cfg(target_arch = "aarch64")]
267    {
268        return unsafe { neon::max_f64_neon(a) };
269    }
270
271    a.iter().copied().fold(f64::NEG_INFINITY, f64::max)
272}
273
274/// Sum of f64 slice using the best available SIMD path.
275pub fn sum_f64_dispatch(a: &[f64]) -> f64 {
276    #[cfg(target_arch = "x86_64")]
277    {
278        if is_x86_feature_detected!("avx512f") {
279            return unsafe { avx512::sum_f64_avx512(a) };
280        }
281        if is_x86_feature_detected!("avx2") {
282            return unsafe { avx2::sum_f64_avx2(a) };
283        }
284        if is_x86_feature_detected!("avx") {
285            return unsafe { avx2::sum_f64_avx(a) };
286        }
287        if is_x86_feature_detected!("sse2") {
288            let mut s = 0.0_f64;
289            let (chunks, remainder) = a.as_chunks::<4>();
290            for c in chunks {
291                s += c[0] + c[1] + c[2] + c[3];
292            }
293            s += remainder.iter().sum::<f64>();
294            return s;
295        }
296    }
297
298    #[cfg(target_arch = "aarch64")]
299    {
300        return unsafe { neon::sum_f64_neon(a) };
301    }
302
303    a.iter().sum()
304}
305
306/// Scale f64 slice in-place: y[i] *= alpha, using the best available SIMD path.
307pub fn scale_f64_dispatch(alpha: f64, y: &mut [f64]) {
308    #[cfg(target_arch = "x86_64")]
309    {
310        if is_x86_feature_detected!("avx512f") {
311            unsafe { avx512::scale_f64_avx512(alpha, y) };
312            return;
313        }
314        if is_x86_feature_detected!("avx2") {
315            unsafe { avx2::scale_f64_avx2(alpha, y) };
316            return;
317        }
318        if is_x86_feature_detected!("avx") {
319            unsafe { avx2::scale_f64_avx(alpha, y) };
320            return;
321        }
322        if is_x86_feature_detected!("sse2") {
323            let (chunks, remainder) = y.as_chunks_mut::<4>();
324            for c in chunks {
325                c[0] *= alpha;
326                c[1] *= alpha;
327                c[2] *= alpha;
328                c[3] *= alpha;
329            }
330            for v in remainder {
331                *v *= alpha;
332            }
333            return;
334        }
335    }
336
337    #[cfg(target_arch = "aarch64")]
338    {
339        unsafe { neon::scale_f64_neon(alpha, y) };
340        return;
341    }
342
343    for x in y.iter_mut() {
344        *x *= alpha;
345    }
346}
347
348/// Hamming distance between two packed bitstream slices.
349pub fn hamming_distance_dispatch(a: &[u64], b: &[u64]) -> u64 {
350    fused_xor_popcount_dispatch(a, b)
351}
352
353/// In-place softmax over an f64 slice (numerically stable).
354///
355/// Computes: subtract max → exp → normalize by sum.
356/// Uses SIMD dispatch for max-find, scaling, and sum reduction.
357pub fn softmax_inplace_f64_dispatch(scores: &mut [f64]) {
358    if scores.is_empty() {
359        return;
360    }
361
362    #[cfg(target_arch = "x86_64")]
363    {
364        if is_x86_feature_detected!("avx2") {
365            unsafe { avx2::softmax_inplace_f64_avx2(scores) };
366            return;
367        }
368    }
369
370    let max_val = max_f64_dispatch(scores);
371    let (chunks, remainder) = scores.as_chunks_mut::<4>();
372    for c in chunks {
373        c[0] = (c[0] - max_val).exp();
374        c[1] = (c[1] - max_val).exp();
375        c[2] = (c[2] - max_val).exp();
376        c[3] = (c[3] - max_val).exp();
377    }
378    for s in remainder {
379        *s = (*s - max_val).exp();
380    }
381
382    let exp_sum = sum_f64_dispatch(scores);
383    if exp_sum > 0.0 {
384        scale_f64_dispatch(1.0 / exp_sum, scores);
385    }
386}
387
388/// Fused encode+AND+popcount dispatch.
389///
390/// Delegates to the scalar-control implementation in `bitstream`,
391/// which already performs SIMD Bernoulli compare where available.
392pub fn encode_and_popcount_dispatch<R: Rng + ?Sized>(
393    weight_words: &[u64],
394    prob: f64,
395    length: usize,
396    rng: &mut R,
397) -> u64 {
398    crate::bitstream::encode_and_popcount(weight_words, prob, length, rng)
399}
400
401/// Batch compare 1024 bytes against threshold using best SIMD.
402pub fn bernoulli_compare_batch_1024(buf: &[u8], threshold: u8, out: &mut [u64]) {
403    #[cfg(target_arch = "x86_64")]
404    {
405        if is_x86_feature_detected!("avx512bw") {
406            return unsafe { avx512::bernoulli_compare_batch_avx512(buf, threshold, out) };
407        }
408        if is_x86_feature_detected!("avx2") {
409            return unsafe { avx2::bernoulli_compare_batch_avx2(buf, threshold, out) };
410        }
411    }
412
413    #[cfg(target_arch = "x86_64")]
414    {
415        // SSE2 fallback (available on all x86_64)
416        use core::arch::x86_64::*;
417        unsafe {
418            let v_thresh = _mm_set1_epi8(threshold as i8);
419            let bias = _mm_set1_epi8(i8::MIN);
420            let v_thresh_biased = _mm_xor_si128(v_thresh, bias);
421
422            for i in 0..16 {
423                let chunk = &buf[i * 64..(i + 1) * 64];
424                let mut word = 0_u64;
425                for j in 0..4 {
426                    let v = _mm_loadu_si128(chunk.as_ptr().add(j * 16) as *const __m128i);
427                    let v_biased = _mm_xor_si128(v, bias);
428                    let m = _mm_cmpgt_epi8(v_thresh_biased, v_biased);
429                    let mask = _mm_movemask_epi8(m) as u32;
430                    word |= (mask as u64) << (j * 16);
431                }
432                out[i] = word;
433            }
434        }
435    }
436
437    // Generic fallback: 16 scalar calls (non-x86_64 architectures)
438    #[cfg(not(target_arch = "x86_64"))]
439    for i in 0..16 {
440        out[i] =
441            crate::bitstream::simd_bernoulli_compare_exposed(&buf[i * 64..(i + 1) * 64], threshold);
442    }
443}