Skip to main content

sc_neurocore_engine/simd/
avx512.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 — AVX512
8
9#[cfg(target_arch = "x86_64")]
10use core::arch::x86_64::*;
11
12#[cfg(target_arch = "x86_64")]
13#[target_feature(enable = "avx512f,avx512vpopcntdq")]
14/// Count set bits in 64-bit words using AVX-512 VPOPCNTDQ.
15///
16/// # Safety
17/// Caller must ensure the current CPU supports `avx512f` and `avx512vpopcntdq`.
18pub unsafe fn popcount_avx512(data: &[u64]) -> u64 {
19    let mut total = 0_u64;
20    let (chunks, remainder) = data.as_chunks::<16>();
21
22    for chunk in chunks {
23        let v0 = _mm512_loadu_si512(chunk.as_ptr() as *const __m512i);
24        let v1 = _mm512_loadu_si512(chunk.as_ptr().add(8) as *const __m512i);
25
26        total += _mm512_reduce_add_epi64(_mm512_popcnt_epi64(v0)) as u64;
27        total += _mm512_reduce_add_epi64(_mm512_popcnt_epi64(v1)) as u64;
28    }
29
30    total + crate::bitstream::popcount_words_portable(remainder)
31}
32
33#[cfg(target_arch = "x86_64")]
34#[target_feature(enable = "avx512f,avx512bw")]
35/// Pack u8 bits into u64 words using AVX-512 k-mask compare.
36///
37/// Processes 64 bytes per iteration where each compare result bit maps
38/// directly to one packed output bit.
39///
40/// # Safety
41/// Caller must ensure the current CPU supports `avx512f` and `avx512bw`.
42pub unsafe fn pack_avx512(bits: &[u8]) -> Vec<u64> {
43    let length = bits.len();
44    let words = length.div_ceil(64);
45    let mut data = vec![0_u64; words];
46    let full_words = length / 64;
47    let zero = _mm512_setzero_si512();
48
49    let (chunks, _) = data[..full_words].as_chunks_mut::<4>();
50    let mut word_idx = 0;
51    for chunk in chunks {
52        let base = word_idx * 64;
53        for i in 0..4 {
54            let v = _mm512_loadu_si512(bits.as_ptr().add(base + i * 64) as *const __m512i);
55            chunk[i] = _mm512_cmpneq_epi8_mask(v, zero);
56        }
57        word_idx += 4;
58    }
59
60    for i in word_idx..full_words {
61        let v = _mm512_loadu_si512(bits.as_ptr().add(i * 64) as *const __m512i);
62        data[i] = _mm512_cmpneq_epi8_mask(v, zero);
63    }
64
65    if full_words < words {
66        let tail_start = full_words * 64;
67        let tail = crate::bitstream::pack_fast(&bits[tail_start..]);
68        data[full_words] = tail.data.first().copied().unwrap_or(0);
69    }
70
71    data
72}
73
74#[cfg(target_arch = "x86_64")]
75#[target_feature(enable = "avx512f,avx512vpopcntdq")]
76/// Fused AND+popcount over packed words using AVX-512 VPOPCNTDQ.
77///
78/// # Safety
79/// Caller must ensure the current CPU supports `avx512f` and `avx512vpopcntdq`.
80pub unsafe fn fused_and_popcount_avx512(a: &[u64], b: &[u64]) -> u64 {
81    let len = a.len().min(b.len());
82    let mut total = 0_u64;
83    let (chunks_a, remainder_a) = a[..len].as_chunks::<16>();
84    let (chunks_b, remainder_b) = b[..len].as_chunks::<16>();
85
86    for (ca, cb) in chunks_a.iter().zip(chunks_b) {
87        let va0 = _mm512_loadu_si512(ca.as_ptr() as *const __m512i);
88        let vb0 = _mm512_loadu_si512(cb.as_ptr() as *const __m512i);
89        let va1 = _mm512_loadu_si512(ca.as_ptr().add(8) as *const __m512i);
90        let vb1 = _mm512_loadu_si512(cb.as_ptr().add(8) as *const __m512i);
91
92        let and0 = _mm512_and_si512(va0, vb0);
93        let and1 = _mm512_and_si512(va1, vb1);
94
95        total += _mm512_reduce_add_epi64(_mm512_popcnt_epi64(and0)) as u64;
96        total += _mm512_reduce_add_epi64(_mm512_popcnt_epi64(and1)) as u64;
97    }
98
99    total
100        + remainder_a
101            .iter()
102            .zip(remainder_b)
103            .map(|(&wa, &wb)| (wa & wb).count_ones() as u64)
104            .sum::<u64>()
105}
106
107#[cfg(target_arch = "x86_64")]
108#[target_feature(enable = "avx512f,avx512vpopcntdq")]
109/// Fused XOR+popcount over packed words using AVX-512 VPOPCNTDQ.
110///
111/// # Safety
112/// Caller must ensure the current CPU supports `avx512f` and `avx512vpopcntdq`.
113pub unsafe fn fused_xor_popcount_avx512(a: &[u64], b: &[u64]) -> u64 {
114    let len = a.len().min(b.len());
115    let mut total = 0_u64;
116    let (chunks_a, remainder_a) = a[..len].as_chunks::<16>();
117    let (chunks_b, remainder_b) = b[..len].as_chunks::<16>();
118
119    for (ca, cb) in chunks_a.iter().zip(chunks_b) {
120        let va0 = _mm512_loadu_si512(ca.as_ptr() as *const __m512i);
121        let vb0 = _mm512_loadu_si512(cb.as_ptr() as *const __m512i);
122        let va1 = _mm512_loadu_si512(ca.as_ptr().add(8) as *const __m512i);
123        let vb1 = _mm512_loadu_si512(cb.as_ptr().add(8) as *const __m512i);
124
125        let xor0 = _mm512_xor_si512(va0, vb0);
126        let xor1 = _mm512_xor_si512(va1, vb1);
127
128        total += _mm512_reduce_add_epi64(_mm512_popcnt_epi64(xor0)) as u64;
129        total += _mm512_reduce_add_epi64(_mm512_popcnt_epi64(xor1)) as u64;
130    }
131
132    total
133        + remainder_a
134            .iter()
135            .zip(remainder_b)
136            .map(|(&wa, &wb)| (wa ^ wb).count_ones() as u64)
137            .sum::<u64>()
138}
139
140#[cfg(not(target_arch = "x86_64"))]
141/// Fallback fused XOR+popcount when AVX-512 is unavailable on this architecture.
142///
143/// # Safety
144/// This function is marked unsafe for API parity with the AVX-512 variant.
145pub unsafe fn fused_xor_popcount_avx512(a: &[u64], b: &[u64]) -> u64 {
146    a.iter()
147        .zip(b.iter())
148        .map(|(&wa, &wb)| (wa ^ wb).count_ones() as u64)
149        .sum()
150}
151
152#[cfg(target_arch = "x86_64")]
153#[target_feature(enable = "avx512f,avx512bw")]
154/// Compare 64 random bytes against an unsigned threshold and return bit mask.
155///
156/// Bit `i` in the returned mask is 1 iff `buf[i] < threshold`.
157///
158/// # Safety
159/// Caller must ensure the current CPU supports `avx512f` and `avx512bw`.
160/// `buf` must have at least 64 elements.
161pub unsafe fn bernoulli_compare_avx512(buf: &[u8], threshold: u8) -> u64 {
162    assert!(buf.len() >= 64, "buffer must contain at least 64 bytes");
163    let data = _mm512_loadu_si512(buf.as_ptr() as *const __m512i);
164    let thresh = _mm512_set1_epi8(threshold as i8);
165    _mm512_cmplt_epu8_mask(data, thresh)
166}
167
168#[cfg(not(target_arch = "x86_64"))]
169/// Fallback popcount when AVX-512 is unavailable on this architecture.
170///
171/// # Safety
172/// This function is marked unsafe for API parity with the AVX-512 variant.
173pub unsafe fn popcount_avx512(data: &[u64]) -> u64 {
174    crate::bitstream::popcount_words_portable(data)
175}
176
177#[cfg(not(target_arch = "x86_64"))]
178/// Fallback pack when AVX-512 is unavailable on this architecture.
179///
180/// # Safety
181/// This function is marked unsafe for API parity with the AVX-512 variant.
182pub unsafe fn pack_avx512(bits: &[u8]) -> Vec<u64> {
183    crate::bitstream::pack_fast(bits).data
184}
185
186#[cfg(not(target_arch = "x86_64"))]
187/// Fallback fused AND+popcount when AVX-512 is unavailable on this architecture.
188///
189/// # Safety
190/// This function is marked unsafe for API parity with the AVX-512 variant.
191pub unsafe fn fused_and_popcount_avx512(a: &[u64], b: &[u64]) -> u64 {
192    a.iter()
193        .zip(b.iter())
194        .map(|(&wa, &wb)| (wa & wb).count_ones() as u64)
195        .sum()
196}
197
198#[cfg(not(target_arch = "x86_64"))]
199/// Fallback Bernoulli compare when AVX-512 is unavailable on this architecture.
200///
201/// # Safety
202/// This function is marked unsafe for API parity with the AVX-512 variant.
203pub unsafe fn bernoulli_compare_avx512(buf: &[u8], threshold: u8) -> u64 {
204    let mut mask = 0_u64;
205    for (bit, &rb) in buf.iter().take(64).enumerate() {
206        if rb < threshold {
207            mask |= 1_u64 << bit;
208        }
209    }
210    mask
211}
212
213// --- f64 SIMD operations (AVX-512: 8-wide f64) ---
214
215#[cfg(target_arch = "x86_64")]
216#[target_feature(enable = "avx512f")]
217/// Dot product of two f64 slices using AVX-512.
218///
219/// # Safety
220/// Caller must ensure the current CPU supports `avx512f`.
221pub unsafe fn dot_f64_avx512(a: &[f64], b: &[f64]) -> f64 {
222    let len = a.len().min(b.len());
223    let mut acc = _mm512_setzero_pd();
224    let (chunks_a, remainder_a) = a[..len].as_chunks::<8>();
225    let (chunks_b, remainder_b) = b[..len].as_chunks::<8>();
226
227    for (ca, cb) in chunks_a.iter().zip(chunks_b) {
228        let va = _mm512_loadu_pd(ca.as_ptr());
229        let vb = _mm512_loadu_pd(cb.as_ptr());
230        acc = _mm512_fmadd_pd(va, vb, acc);
231    }
232
233    let mut sum = _mm512_reduce_add_pd(acc);
234    for (&ra, &rb) in remainder_a.iter().zip(remainder_b) {
235        sum += ra * rb;
236    }
237    sum
238}
239
240#[cfg(target_arch = "x86_64")]
241#[target_feature(enable = "avx512f")]
242/// Maximum of f64 slice using AVX-512.
243///
244/// # Safety
245/// Caller must ensure the current CPU supports `avx512f`.
246pub unsafe fn max_f64_avx512(a: &[f64]) -> f64 {
247    if a.is_empty() {
248        return f64::NEG_INFINITY;
249    }
250    let mut vmax0 = _mm512_set1_pd(f64::NEG_INFINITY);
251    let mut vmax1 = _mm512_set1_pd(f64::NEG_INFINITY);
252    let (chunks, remainder) = a.as_chunks::<16>();
253
254    for chunk in chunks {
255        vmax0 = _mm512_max_pd(vmax0, _mm512_loadu_pd(chunk.as_ptr()));
256        vmax1 = _mm512_max_pd(vmax1, _mm512_loadu_pd(chunk.as_ptr().add(8)));
257    }
258
259    let mut m = _mm512_reduce_max_pd(_mm512_max_pd(vmax0, vmax1));
260    for &v in remainder {
261        m = m.max(v);
262    }
263    m
264}
265
266#[cfg(target_arch = "x86_64")]
267#[target_feature(enable = "avx512f")]
268/// Sum of f64 slice using AVX-512.
269///
270/// # Safety
271/// Caller must ensure the current CPU supports `avx512f`.
272pub unsafe fn sum_f64_avx512(a: &[f64]) -> f64 {
273    let mut acc0 = _mm512_setzero_pd();
274    let mut acc1 = _mm512_setzero_pd();
275    let (chunks, remainder) = a.as_chunks::<16>();
276
277    for chunk in chunks {
278        acc0 = _mm512_add_pd(acc0, _mm512_loadu_pd(chunk.as_ptr()));
279        acc1 = _mm512_add_pd(acc1, _mm512_loadu_pd(chunk.as_ptr().add(8)));
280    }
281
282    let mut sum = _mm512_reduce_add_pd(_mm512_add_pd(acc0, acc1));
283    for &v in remainder {
284        sum += v;
285    }
286    sum
287}
288
289#[cfg(target_arch = "x86_64")]
290#[target_feature(enable = "avx512f")]
291/// Scale f64 slice in-place: y[i] *= alpha, using AVX-512.
292///
293/// # Safety
294/// Caller must ensure the current CPU supports `avx512f`.
295pub unsafe fn scale_f64_avx512(alpha: f64, y: &mut [f64]) {
296    let valpha = _mm512_set1_pd(alpha);
297    let (chunks, remainder) = y.as_chunks_mut::<16>();
298
299    for chunk in chunks {
300        let v0 = _mm512_loadu_pd(chunk.as_ptr());
301        let v1 = _mm512_loadu_pd(chunk.as_ptr().add(8));
302        _mm512_storeu_pd(chunk.as_mut_ptr(), _mm512_mul_pd(v0, valpha));
303        _mm512_storeu_pd(chunk.as_mut_ptr().add(8), _mm512_mul_pd(v1, valpha));
304    }
305
306    for v in remainder {
307        *v *= alpha;
308    }
309}
310
311#[cfg(not(target_arch = "x86_64"))]
312pub unsafe fn dot_f64_avx512(a: &[f64], b: &[f64]) -> f64 {
313    let len = a.len().min(b.len());
314    a[..len].iter().zip(&b[..len]).map(|(&x, &y)| x * y).sum()
315}
316
317#[cfg(not(target_arch = "x86_64"))]
318pub unsafe fn max_f64_avx512(a: &[f64]) -> f64 {
319    a.iter().copied().fold(f64::NEG_INFINITY, f64::max)
320}
321
322#[cfg(not(target_arch = "x86_64"))]
323pub unsafe fn sum_f64_avx512(a: &[f64]) -> f64 {
324    a.iter().sum()
325}
326
327#[cfg(not(target_arch = "x86_64"))]
328pub unsafe fn scale_f64_avx512(alpha: f64, y: &mut [f64]) {
329    for v in y.iter_mut() {
330        *v *= alpha;
331    }
332}
333
334#[cfg(target_arch = "x86_64")]
335#[target_feature(enable = "avx512bw")]
336/// Compare 1024 random bytes against a threshold and return 16 u64 words.
337///
338/// # Safety
339/// Caller must ensure AVX-512 BW is available on the current CPU.
340pub unsafe fn bernoulli_compare_batch_avx512(buf: &[u8], threshold: u8, out: &mut [u64]) {
341    let v_thresh = _mm512_set1_epi8(threshold as i8);
342    for i in 0..16 {
343        let chunk = &buf[i * 64..(i + 1) * 64];
344        let v = _mm512_loadu_si512(chunk.as_ptr() as *const _);
345        // AVX-512 has direct unsigned comparison
346        out[i] = _mm512_cmplt_epu8_mask(v, v_thresh);
347    }
348}
349
350#[cfg(all(test, target_arch = "x86_64"))]
351mod tests {
352    use crate::bitstream::pack;
353
354    #[test]
355    fn pack_avx512_matches_pack() {
356        if !is_x86_feature_detected!("avx512bw") {
357            return;
358        }
359
360        let lengths = [
361            1_usize, 7, 31, 32, 33, 63, 64, 65, 127, 128, 129, 1024, 1031,
362        ];
363        for length in lengths {
364            let bits: Vec<u8> = (0..length)
365                .map(|i| if (i * 19 + 11) % 4 == 0 { 1 } else { 0 })
366                .collect();
367            // SAFETY: Runtime-guarded by feature detection in this test.
368            let got = unsafe { super::pack_avx512(&bits) };
369            let expected = pack(&bits).data;
370            assert_eq!(got, expected, "Mismatch at length={length}");
371        }
372    }
373
374    #[test]
375    fn fused_and_popcount_avx512_matches_scalar() {
376        if !is_x86_feature_detected!("avx512vpopcntdq") {
377            return;
378        }
379
380        let lengths = [1_usize, 7, 8, 15, 16, 17, 31, 32, 64, 128];
381        for len in lengths {
382            let a: Vec<u64> = (0..len)
383                .map(|i| (i as u64).wrapping_mul(0xD6E8_FD9D_5A2B_1C47) ^ 0x1357_9BDF_2468_ACE0)
384                .collect();
385            let b: Vec<u64> = (0..len)
386                .map(|i| (i as u64).wrapping_mul(0x94D0_49BB_1331_11EB) ^ 0xF0F0_0F0F_AAAA_5555)
387                .collect();
388
389            let expected: u64 = a
390                .iter()
391                .zip(b.iter())
392                .map(|(&wa, &wb)| (wa & wb).count_ones() as u64)
393                .sum();
394
395            // SAFETY: Runtime-guarded by feature detection in this test.
396            let got = unsafe { super::fused_and_popcount_avx512(&a, &b) };
397            assert_eq!(got, expected, "Mismatch at len={len}");
398        }
399    }
400
401    #[test]
402    fn bernoulli_compare_avx512_matches_scalar() {
403        if !is_x86_feature_detected!("avx512bw") {
404            return;
405        }
406
407        let buf: Vec<u8> = (0..64).map(|i| (i * 41 + 23) as u8).collect();
408        let thresholds = [0_u8, 1, 2, 17, 64, 127, 128, 200, 255];
409
410        for threshold in thresholds {
411            let expected = buf.iter().enumerate().fold(0_u64, |acc, (bit, &rb)| {
412                acc | (u64::from(rb < threshold) << bit)
413            });
414
415            // SAFETY: Runtime-guarded by feature detection in this test.
416            let got = unsafe { super::bernoulli_compare_avx512(&buf, threshold) };
417            assert_eq!(
418                got, expected,
419                "Mismatch for threshold={threshold} buf={buf:?}"
420            );
421        }
422    }
423}