Skip to main content

sc_neurocore_engine/simd/
avx2.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 — AVX2
8
9#[cfg(target_arch = "x86_64")]
10use core::arch::x86_64::*;
11
12#[cfg(target_arch = "x86_64")]
13#[target_feature(enable = "avx2")]
14/// Count set bits in 64-bit words using AVX2.
15///
16/// # Safety
17/// Caller must ensure the current CPU supports `avx2`.
18pub unsafe fn popcount_avx2(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 = _mm256_loadu_si256(chunk.as_ptr() as *const __m256i);
24        let v1 = _mm256_loadu_si256(chunk.as_ptr().add(4) as *const __m256i);
25        let v2 = _mm256_loadu_si256(chunk.as_ptr().add(8) as *const __m256i);
26        let v3 = _mm256_loadu_si256(chunk.as_ptr().add(12) as *const __m256i);
27
28        let mut lanes = [0_u64; 16];
29        _mm256_storeu_si256(lanes.as_mut_ptr() as *mut __m256i, v0);
30        _mm256_storeu_si256(lanes.as_mut_ptr().add(4) as *mut __m256i, v1);
31        _mm256_storeu_si256(lanes.as_mut_ptr().add(8) as *mut __m256i, v2);
32        _mm256_storeu_si256(lanes.as_mut_ptr().add(12) as *mut __m256i, v3);
33
34        for &w in lanes.iter() {
35            total += w.count_ones() as u64;
36        }
37    }
38
39    total
40        + remainder
41            .iter()
42            .map(|&w| w.count_ones() as u64)
43            .sum::<u64>()
44}
45
46#[cfg(target_arch = "x86_64")]
47#[target_feature(enable = "avx2")]
48/// Pack u8 bits into u64 words using AVX2 movemask.
49///
50/// Processes 64 bytes into one u64 word by building two 32-bit masks.
51///
52/// # Safety
53/// Caller must ensure the current CPU supports `avx2`.
54pub unsafe fn pack_avx2(bits: &[u8]) -> Vec<u64> {
55    let length = bits.len();
56    let words = length.div_ceil(64);
57    let mut data = vec![0_u64; words];
58    let full_words = length / 64;
59    let zero = _mm256_setzero_si256();
60
61    let (chunks, _) = data[..full_words].as_chunks_mut::<4>();
62    let mut word_idx = 0;
63    for chunk in chunks {
64        let base = word_idx * 64;
65        for i in 0..4 {
66            let b = base + i * 64;
67            let lo = _mm256_loadu_si256(bits.as_ptr().add(b) as *const __m256i);
68            let hi = _mm256_loadu_si256(bits.as_ptr().add(b + 32) as *const __m256i);
69            let lo_mask = !(_mm256_movemask_epi8(_mm256_cmpeq_epi8(lo, zero)) as u32);
70            let hi_mask = !(_mm256_movemask_epi8(_mm256_cmpeq_epi8(hi, zero)) as u32);
71            chunk[i] = ((hi_mask as u64) << 32) | (lo_mask as u64);
72        }
73        word_idx += 4;
74    }
75
76    for i in word_idx..full_words {
77        let base = i * 64;
78        let lo = _mm256_loadu_si256(bits.as_ptr().add(base) as *const __m256i);
79        let hi = _mm256_loadu_si256(bits.as_ptr().add(base + 32) as *const __m256i);
80        let lo_mask = !(_mm256_movemask_epi8(_mm256_cmpeq_epi8(lo, zero)) as u32);
81        let hi_mask = !(_mm256_movemask_epi8(_mm256_cmpeq_epi8(hi, zero)) as u32);
82        data[i] = ((hi_mask as u64) << 32) | (lo_mask as u64);
83    }
84
85    if full_words < words {
86        let tail_start = full_words * 64;
87        let tail = crate::bitstream::pack_fast(&bits[tail_start..]);
88        data[full_words] = tail.data.first().copied().unwrap_or(0);
89    }
90
91    data
92}
93
94#[cfg(target_arch = "x86_64")]
95#[target_feature(enable = "avx2")]
96/// Fused AND+popcount over packed words using AVX2 for the AND stage.
97///
98/// # Safety
99/// Caller must ensure the current CPU supports `avx2`.
100pub unsafe fn fused_and_popcount_avx2(a: &[u64], b: &[u64]) -> u64 {
101    let len = a.len().min(b.len());
102    let mut total = 0_u64;
103    let (chunks_a, remainder_a) = a[..len].as_chunks::<16>();
104    let (chunks_b, remainder_b) = b[..len].as_chunks::<16>();
105
106    for (ca, cb) in chunks_a.iter().zip(chunks_b) {
107        let va0 = _mm256_loadu_si256(ca.as_ptr() as *const __m256i);
108        let vb0 = _mm256_loadu_si256(cb.as_ptr() as *const __m256i);
109        let va1 = _mm256_loadu_si256(ca.as_ptr().add(4) as *const __m256i);
110        let vb1 = _mm256_loadu_si256(cb.as_ptr().add(4) as *const __m256i);
111        let va2 = _mm256_loadu_si256(ca.as_ptr().add(8) as *const __m256i);
112        let vb2 = _mm256_loadu_si256(cb.as_ptr().add(8) as *const __m256i);
113        let va3 = _mm256_loadu_si256(ca.as_ptr().add(12) as *const __m256i);
114        let vb3 = _mm256_loadu_si256(cb.as_ptr().add(12) as *const __m256i);
115
116        let and0 = _mm256_and_si256(va0, vb0);
117        let and1 = _mm256_and_si256(va1, vb1);
118        let and2 = _mm256_and_si256(va2, vb2);
119        let and3 = _mm256_and_si256(va3, vb3);
120
121        let mut lanes = [0_u64; 16];
122        _mm256_storeu_si256(lanes.as_mut_ptr() as *mut __m256i, and0);
123        _mm256_storeu_si256(lanes.as_mut_ptr().add(4) as *mut __m256i, and1);
124        _mm256_storeu_si256(lanes.as_mut_ptr().add(8) as *mut __m256i, and2);
125        _mm256_storeu_si256(lanes.as_mut_ptr().add(12) as *mut __m256i, and3);
126
127        for &w in lanes.iter() {
128            total += w.count_ones() as u64;
129        }
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(target_arch = "x86_64")]
141#[target_feature(enable = "avx2")]
142/// Fused XOR+popcount over packed words using AVX2 for the XOR stage.
143///
144/// # Safety
145/// Caller must ensure the current CPU supports `avx2`.
146pub unsafe fn fused_xor_popcount_avx2(a: &[u64], b: &[u64]) -> u64 {
147    let len = a.len().min(b.len());
148    let mut total = 0_u64;
149    let (chunks_a, remainder_a) = a[..len].as_chunks::<16>();
150    let (chunks_b, remainder_b) = b[..len].as_chunks::<16>();
151
152    for (ca, cb) in chunks_a.iter().zip(chunks_b) {
153        let va0 = _mm256_loadu_si256(ca.as_ptr() as *const __m256i);
154        let vb0 = _mm256_loadu_si256(cb.as_ptr() as *const __m256i);
155        let va1 = _mm256_loadu_si256(ca.as_ptr().add(4) as *const __m256i);
156        let vb1 = _mm256_loadu_si256(cb.as_ptr().add(4) as *const __m256i);
157        let va2 = _mm256_loadu_si256(ca.as_ptr().add(8) as *const __m256i);
158        let vb2 = _mm256_loadu_si256(cb.as_ptr().add(8) as *const __m256i);
159        let va3 = _mm256_loadu_si256(ca.as_ptr().add(12) as *const __m256i);
160        let vb3 = _mm256_loadu_si256(cb.as_ptr().add(12) as *const __m256i);
161
162        let xor0 = _mm256_xor_si256(va0, vb0);
163        let xor1 = _mm256_xor_si256(va1, vb1);
164        let xor2 = _mm256_xor_si256(va2, vb2);
165        let xor3 = _mm256_xor_si256(va3, vb3);
166
167        let mut lanes = [0_u64; 16];
168        _mm256_storeu_si256(lanes.as_mut_ptr() as *mut __m256i, xor0);
169        _mm256_storeu_si256(lanes.as_mut_ptr().add(4) as *mut __m256i, xor1);
170        _mm256_storeu_si256(lanes.as_mut_ptr().add(8) as *mut __m256i, xor2);
171        _mm256_storeu_si256(lanes.as_mut_ptr().add(12) as *mut __m256i, xor3);
172
173        for &w in lanes.iter() {
174            total += w.count_ones() as u64;
175        }
176    }
177
178    total
179        + remainder_a
180            .iter()
181            .zip(remainder_b)
182            .map(|(&wa, &wb)| (wa ^ wb).count_ones() as u64)
183            .sum::<u64>()
184}
185
186#[cfg(not(target_arch = "x86_64"))]
187/// Fallback fused XOR+popcount when AVX2 is unavailable on this architecture.
188///
189/// # Safety
190/// This function is marked unsafe for API parity with the AVX2 variant.
191pub unsafe fn fused_xor_popcount_avx2(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(target_arch = "x86_64")]
199#[target_feature(enable = "avx2")]
200/// Compare 32 random bytes against an unsigned threshold and return bit mask.
201///
202/// Bit `i` in the returned mask is 1 iff `buf[i] < threshold`.
203///
204/// # Safety
205/// Caller must ensure the current CPU supports `avx2`.
206/// `buf` must have at least 32 elements.
207pub unsafe fn bernoulli_compare_avx2(buf: &[u8], threshold: u8) -> u32 {
208    assert!(buf.len() >= 32, "buffer must contain at least 32 bytes");
209
210    let data = _mm256_loadu_si256(buf.as_ptr() as *const __m256i);
211    let bias = _mm256_set1_epi8(i8::MIN);
212    let data_biased = _mm256_xor_si256(data, bias);
213    let thresh_biased = _mm256_set1_epi8((threshold ^ 0x80) as i8);
214    let lt = _mm256_cmpgt_epi8(thresh_biased, data_biased);
215    _mm256_movemask_epi8(lt) as u32
216}
217
218#[cfg(not(target_arch = "x86_64"))]
219/// Fallback popcount when AVX2 is unavailable on this architecture.
220///
221/// # Safety
222/// This function is marked unsafe for API parity with the AVX2 variant.
223pub unsafe fn popcount_avx2(data: &[u64]) -> u64 {
224    crate::bitstream::popcount_words_portable(data)
225}
226
227#[cfg(not(target_arch = "x86_64"))]
228/// Fallback pack when AVX2 is unavailable on this architecture.
229///
230/// # Safety
231/// This function is marked unsafe for API parity with the AVX2 variant.
232pub unsafe fn pack_avx2(bits: &[u8]) -> Vec<u64> {
233    crate::bitstream::pack_fast(bits).data
234}
235
236#[cfg(not(target_arch = "x86_64"))]
237/// Fallback fused AND+popcount when AVX2 is unavailable on this architecture.
238///
239/// # Safety
240/// This function is marked unsafe for API parity with the AVX2 variant.
241pub unsafe fn fused_and_popcount_avx2(a: &[u64], b: &[u64]) -> u64 {
242    a.iter()
243        .zip(b.iter())
244        .map(|(&wa, &wb)| (wa & wb).count_ones() as u64)
245        .sum()
246}
247
248#[cfg(not(target_arch = "x86_64"))]
249/// Fallback Bernoulli compare when AVX2 is unavailable on this architecture.
250///
251/// # Safety
252/// This function is marked unsafe for API parity with the AVX2 variant.
253pub unsafe fn bernoulli_compare_avx2(buf: &[u8], threshold: u8) -> u32 {
254    let mut mask = 0_u32;
255    for (bit, &rb) in buf.iter().take(32).enumerate() {
256        if rb < threshold {
257            mask |= 1_u32 << bit;
258        }
259    }
260    mask
261}
262
263// --- f64 SIMD operations (AVX2: 4-wide f64) ---
264
265#[cfg(target_arch = "x86_64")]
266#[target_feature(enable = "avx2,fma")]
267/// Dot product of two f64 slices using AVX2 FMA.
268///
269/// # Safety
270/// Caller must ensure the current CPU supports `avx2` and `fma`.
271pub unsafe fn dot_f64_avx2(a: &[f64], b: &[f64]) -> f64 {
272    let len = a.len().min(b.len());
273    let mut acc = _mm256_setzero_pd();
274    let (chunks_a, remainder_a) = a[..len].as_chunks::<4>();
275    let (chunks_b, remainder_b) = b[..len].as_chunks::<4>();
276
277    for (ca, cb) in chunks_a.iter().zip(chunks_b) {
278        let va = _mm256_loadu_pd(ca.as_ptr());
279        let vb = _mm256_loadu_pd(cb.as_ptr());
280        acc = _mm256_fmadd_pd(va, vb, acc);
281    }
282
283    let mut lanes = [0.0_f64; 4];
284    _mm256_storeu_pd(lanes.as_mut_ptr(), acc);
285    let mut sum = lanes[0] + lanes[1] + lanes[2] + lanes[3];
286
287    for (&ra, &rb) in remainder_a.iter().zip(remainder_b) {
288        sum += ra * rb;
289    }
290    sum
291}
292
293#[cfg(target_arch = "x86_64")]
294#[target_feature(enable = "avx2")]
295/// Maximum of f64 slice using AVX2.
296///
297/// # Safety
298/// Caller must ensure the current CPU supports `avx2`.
299pub unsafe fn max_f64_avx2(a: &[f64]) -> f64 {
300    if a.is_empty() {
301        return f64::NEG_INFINITY;
302    }
303    let mut vmax = _mm256_set1_pd(f64::NEG_INFINITY);
304    let (chunks, remainder) = a.as_chunks::<4>();
305
306    for chunk in chunks {
307        let va = _mm256_loadu_pd(chunk.as_ptr());
308        vmax = _mm256_max_pd(vmax, va);
309    }
310
311    let mut lanes = [0.0_f64; 4];
312    _mm256_storeu_pd(lanes.as_mut_ptr(), vmax);
313    let mut m = lanes[0].max(lanes[1]).max(lanes[2].max(lanes[3]));
314    for &v in remainder {
315        m = m.max(v);
316    }
317    m
318}
319
320#[cfg(target_arch = "x86_64")]
321#[target_feature(enable = "avx2")]
322/// Sum of f64 slice using AVX2.
323///
324/// # Safety
325/// Caller must ensure the current CPU supports `avx2`.
326pub unsafe fn sum_f64_avx2(a: &[f64]) -> f64 {
327    let mut acc = _mm256_setzero_pd();
328    let (chunks, remainder) = a.as_chunks::<4>();
329
330    for chunk in chunks {
331        let va = _mm256_loadu_pd(chunk.as_ptr());
332        acc = _mm256_add_pd(acc, va);
333    }
334
335    let mut lanes = [0.0_f64; 4];
336    _mm256_storeu_pd(lanes.as_mut_ptr(), acc);
337    let mut sum = lanes[0] + lanes[1] + lanes[2] + lanes[3];
338    for &v in remainder {
339        sum += v;
340    }
341    sum
342}
343
344#[cfg(target_arch = "x86_64")]
345#[target_feature(enable = "avx2")]
346/// Scale f64 slice in-place: y[i] *= alpha, using AVX2.
347///
348/// # Safety
349/// Caller must ensure the current CPU supports `avx2`.
350pub unsafe fn scale_f64_avx2(alpha: f64, y: &mut [f64]) {
351    let valpha = _mm256_set1_pd(alpha);
352    let (chunks, remainder) = y.as_chunks_mut::<16>();
353
354    for chunk in chunks {
355        let v0 = _mm256_loadu_pd(chunk.as_ptr());
356        let v1 = _mm256_loadu_pd(chunk.as_ptr().add(4));
357        let v2 = _mm256_loadu_pd(chunk.as_ptr().add(8));
358        let v3 = _mm256_loadu_pd(chunk.as_ptr().add(12));
359
360        _mm256_storeu_pd(chunk.as_mut_ptr(), _mm256_mul_pd(v0, valpha));
361        _mm256_storeu_pd(chunk.as_mut_ptr().add(4), _mm256_mul_pd(v1, valpha));
362        _mm256_storeu_pd(chunk.as_mut_ptr().add(8), _mm256_mul_pd(v2, valpha));
363        _mm256_storeu_pd(chunk.as_mut_ptr().add(12), _mm256_mul_pd(v3, valpha));
364    }
365
366    for v in remainder {
367        *v *= alpha;
368    }
369}
370
371#[cfg(not(target_arch = "x86_64"))]
372pub unsafe fn dot_f64_avx2(a: &[f64], b: &[f64]) -> f64 {
373    let len = a.len().min(b.len());
374    a[..len].iter().zip(&b[..len]).map(|(&x, &y)| x * y).sum()
375}
376
377#[cfg(not(target_arch = "x86_64"))]
378pub unsafe fn max_f64_avx2(a: &[f64]) -> f64 {
379    a.iter().copied().fold(f64::NEG_INFINITY, f64::max)
380}
381
382#[cfg(not(target_arch = "x86_64"))]
383pub unsafe fn sum_f64_avx2(a: &[f64]) -> f64 {
384    a.iter().sum()
385}
386
387#[cfg(not(target_arch = "x86_64"))]
388pub unsafe fn scale_f64_avx2(alpha: f64, y: &mut [f64]) {
389    for v in y.iter_mut() {
390        *v *= alpha;
391    }
392}
393
394/// Hamming distance between two packed bitstream slices using AVX2.
395///
396/// # Safety
397/// Caller must ensure the current CPU supports `avx2`.
398pub unsafe fn hamming_distance_avx2(a: &[u64], b: &[u64]) -> u64 {
399    fused_xor_popcount_avx2(a, b)
400}
401
402#[cfg(target_arch = "x86_64")]
403#[target_feature(enable = "avx2")]
404/// In-place softmax using AVX2 for max, sum, and scale steps.
405///
406/// # Safety
407/// Caller must ensure the current CPU supports `avx2`.
408pub unsafe fn softmax_inplace_f64_avx2(scores: &mut [f64]) {
409    if scores.is_empty() {
410        return;
411    }
412    let max_val = max_f64_avx2(scores);
413    // let v_max = _mm256_set1_pd(max_val);
414
415    let (chunks, remainder) = scores.as_chunks_mut::<16>();
416    for chunk in chunks {
417        // We still have to use scalar exp() because AVX2 does not have it in core::arch
418        // but we can unroll the subtractions and stores.
419        for i in 0..16 {
420            chunk[i] = (chunk[i] - max_val).exp();
421        }
422    }
423    for s in remainder {
424        *s = (*s - max_val).exp();
425    }
426
427    let exp_sum = sum_f64_avx2(scores);
428    if exp_sum > 0.0 {
429        scale_f64_avx2(1.0 / exp_sum, scores);
430    }
431}
432
433#[cfg(not(target_arch = "x86_64"))]
434pub unsafe fn softmax_inplace_f64_avx2(scores: &mut [f64]) {
435    if scores.is_empty() {
436        return;
437    }
438    let max_val = max_f64_avx2(scores);
439    // let v_max = _mm256_set1_pd(max_val);
440
441    let (chunks, remainder) = scores.as_chunks_mut::<16>();
442    for chunk in chunks {
443        // We still have to use scalar exp() because AVX2 does not have it in core::arch
444        // but we can unroll the subtractions and stores.
445        for i in 0..16 {
446            chunk[i] = (chunk[i] - max_val).exp();
447        }
448    }
449    for s in remainder {
450        *s = (*s - max_val).exp();
451    }
452
453    let exp_sum = sum_f64_avx2(scores);
454    if exp_sum > 0.0 {
455        scale_f64_avx2(1.0 / exp_sum, scores);
456    }
457}
458
459#[cfg(target_arch = "x86_64")]
460#[target_feature(enable = "avx")]
461/// Dot product of two f64 slices using AVX (no FMA).
462///
463/// # Safety
464/// Caller must ensure the current CPU supports `avx`.
465pub unsafe fn dot_f64_avx(a: &[f64], b: &[f64]) -> f64 {
466    let len = a.len().min(b.len());
467    let mut acc0 = _mm256_setzero_pd();
468    let mut acc1 = _mm256_setzero_pd();
469    let mut acc2 = _mm256_setzero_pd();
470    let mut acc3 = _mm256_setzero_pd();
471
472    let (chunks_a, remainder_a) = a[..len].as_chunks::<16>();
473    let (chunks_b, remainder_b) = b[..len].as_chunks::<16>();
474
475    for (ca, cb) in chunks_a.iter().zip(chunks_b) {
476        let va0 = _mm256_loadu_pd(ca.as_ptr());
477        let vb0 = _mm256_loadu_pd(cb.as_ptr());
478        acc0 = _mm256_add_pd(acc0, _mm256_mul_pd(va0, vb0));
479
480        let va1 = _mm256_loadu_pd(ca.as_ptr().add(4));
481        let vb1 = _mm256_loadu_pd(cb.as_ptr().add(4));
482        acc1 = _mm256_add_pd(acc1, _mm256_mul_pd(va1, vb1));
483
484        let va2 = _mm256_loadu_pd(ca.as_ptr().add(8));
485        let vb2 = _mm256_loadu_pd(cb.as_ptr().add(8));
486        acc2 = _mm256_add_pd(acc2, _mm256_mul_pd(va2, vb2));
487
488        let va3 = _mm256_loadu_pd(ca.as_ptr().add(12));
489        let vb3 = _mm256_loadu_pd(cb.as_ptr().add(12));
490        acc3 = _mm256_add_pd(acc3, _mm256_mul_pd(va3, vb3));
491    }
492
493    acc0 = _mm256_add_pd(acc0, acc1);
494    acc2 = _mm256_add_pd(acc2, acc3);
495    acc0 = _mm256_add_pd(acc0, acc2);
496
497    let mut lanes = [0.0_f64; 4];
498    _mm256_storeu_pd(lanes.as_mut_ptr(), acc0);
499    let mut sum = lanes[0] + lanes[1] + lanes[2] + lanes[3];
500
501    for (&ra, &rb) in remainder_a.iter().zip(remainder_b) {
502        sum += ra * rb;
503    }
504    sum
505}
506
507#[cfg(target_arch = "x86_64")]
508#[target_feature(enable = "avx")]
509/// Sum of f64 slice using AVX.
510///
511/// # Safety
512/// Caller must ensure AVX is available on the current CPU.
513pub unsafe fn sum_f64_avx(a: &[f64]) -> f64 {
514    let mut acc0 = _mm256_setzero_pd();
515    let mut acc1 = _mm256_setzero_pd();
516    let mut acc2 = _mm256_setzero_pd();
517    let mut acc3 = _mm256_setzero_pd();
518
519    let (chunks, remainder) = a.as_chunks::<16>();
520    for chunk in chunks {
521        acc0 = _mm256_add_pd(acc0, _mm256_loadu_pd(chunk.as_ptr()));
522        acc1 = _mm256_add_pd(acc1, _mm256_loadu_pd(chunk.as_ptr().add(4)));
523        acc2 = _mm256_add_pd(acc2, _mm256_loadu_pd(chunk.as_ptr().add(8)));
524        acc3 = _mm256_add_pd(acc3, _mm256_loadu_pd(chunk.as_ptr().add(12)));
525    }
526
527    acc0 = _mm256_add_pd(acc0, acc1);
528    acc2 = _mm256_add_pd(acc2, acc3);
529    acc0 = _mm256_add_pd(acc0, acc2);
530
531    let mut lanes = [0.0_f64; 4];
532    _mm256_storeu_pd(lanes.as_mut_ptr(), acc0);
533    let mut sum = lanes[0] + lanes[1] + lanes[2] + lanes[3];
534    for &v in remainder {
535        sum += v;
536    }
537    sum
538}
539
540#[cfg(target_arch = "x86_64")]
541#[target_feature(enable = "avx2")]
542/// Compare 1024 random bytes against a threshold and return 16 u64 words.
543///
544/// # Safety
545/// Caller must ensure AVX2 is available on the current CPU.
546pub unsafe fn bernoulli_compare_batch_avx2(buf: &[u8], threshold: u8, out: &mut [u64]) {
547    let v_thresh = _mm256_set1_epi8(threshold as i8);
548    // Note: epi8 comparison is signed. Using the xor 0x80 trick for unsigned.
549    let bias = _mm256_set1_epi8(i8::MIN);
550    let v_thresh_biased = _mm256_xor_si256(v_thresh, bias);
551
552    for i in 0..16 {
553        // Each loop iteration processes 64 bytes (2x 256-bit registers)
554        let chunk = &buf[i * 64..(i + 1) * 64];
555        let v0 = _mm256_loadu_si256(chunk.as_ptr() as *const __m256i);
556        let v1 = _mm256_loadu_si256(chunk.as_ptr().add(32) as *const __m256i);
557
558        let v0_biased = _mm256_xor_si256(v0, bias);
559        let v1_biased = _mm256_xor_si256(v1, bias);
560
561        let m0 = _mm256_cmpgt_epi8(v_thresh_biased, v0_biased);
562        let m1 = _mm256_cmpgt_epi8(v_thresh_biased, v1_biased);
563
564        let mask0 = _mm256_movemask_epi8(m0) as u32;
565        let mask1 = _mm256_movemask_epi8(m1) as u32;
566        out[i] = (mask0 as u64) | ((mask1 as u64) << 32);
567    }
568}
569
570#[cfg(target_arch = "x86_64")]
571#[target_feature(enable = "avx")]
572/// Maximum of f64 slice using AVX (v1).
573///
574/// # Safety
575/// Caller must ensure AVX is available on the current CPU.
576pub unsafe fn max_f64_avx(a: &[f64]) -> f64 {
577    if a.is_empty() {
578        return f64::NEG_INFINITY;
579    }
580    let mut max_vec0 = _mm256_set1_pd(f64::NEG_INFINITY);
581    let mut max_vec1 = _mm256_set1_pd(f64::NEG_INFINITY);
582    let mut max_vec2 = _mm256_set1_pd(f64::NEG_INFINITY);
583    let mut max_vec3 = _mm256_set1_pd(f64::NEG_INFINITY);
584
585    let (chunks, remainder) = a.as_chunks::<16>();
586    for chunk in chunks {
587        max_vec0 = _mm256_max_pd(max_vec0, _mm256_loadu_pd(chunk.as_ptr()));
588        max_vec1 = _mm256_max_pd(max_vec1, _mm256_loadu_pd(chunk.as_ptr().add(4)));
589        max_vec2 = _mm256_max_pd(max_vec2, _mm256_loadu_pd(chunk.as_ptr().add(8)));
590        max_vec3 = _mm256_max_pd(max_vec3, _mm256_loadu_pd(chunk.as_ptr().add(12)));
591    }
592
593    max_vec0 = _mm256_max_pd(max_vec0, max_vec1);
594    max_vec2 = _mm256_max_pd(max_vec2, max_vec3);
595    max_vec0 = _mm256_max_pd(max_vec0, max_vec2);
596
597    let mut lanes = [0.0_f64; 4];
598    _mm256_storeu_pd(lanes.as_mut_ptr(), max_vec0);
599    let mut m = lanes[0].max(lanes[1]).max(lanes[2].max(lanes[3]));
600    for &v in remainder {
601        m = m.max(v);
602    }
603    m
604}
605
606#[cfg(target_arch = "x86_64")]
607#[target_feature(enable = "avx")]
608/// Scale f64 slice using AVX (v1).
609///
610/// # Safety
611/// Caller must ensure AVX is available on the current CPU.
612pub unsafe fn scale_f64_avx(alpha: f64, y: &mut [f64]) {
613    let valpha = _mm256_set1_pd(alpha);
614    let (chunks, remainder) = y.as_chunks_mut::<16>();
615
616    for chunk in chunks {
617        let v0 = _mm256_loadu_pd(chunk.as_ptr());
618        let v1 = _mm256_loadu_pd(chunk.as_ptr().add(4));
619        let v2 = _mm256_loadu_pd(chunk.as_ptr().add(8));
620        let v3 = _mm256_loadu_pd(chunk.as_ptr().add(12));
621
622        _mm256_storeu_pd(chunk.as_mut_ptr(), _mm256_mul_pd(v0, valpha));
623        _mm256_storeu_pd(chunk.as_mut_ptr().add(4), _mm256_mul_pd(v1, valpha));
624        _mm256_storeu_pd(chunk.as_mut_ptr().add(8), _mm256_mul_pd(v2, valpha));
625        _mm256_storeu_pd(chunk.as_mut_ptr().add(12), _mm256_mul_pd(v3, valpha));
626    }
627
628    for v in remainder {
629        *v *= alpha;
630    }
631}
632
633#[cfg(all(test, target_arch = "x86_64"))]
634mod tests {
635    use crate::bitstream::pack;
636
637    #[test]
638    fn pack_avx2_matches_pack() {
639        if !is_x86_feature_detected!("avx2") {
640            return;
641        }
642
643        let lengths = [
644            1_usize, 7, 31, 32, 33, 63, 64, 65, 127, 128, 129, 1024, 1031,
645        ];
646        for length in lengths {
647            let bits: Vec<u8> = (0..length)
648                .map(|i| if (i * 17 + 5) % 3 == 0 { 1 } else { 0 })
649                .collect();
650            // SAFETY: Runtime-guarded by feature detection in this test.
651            let got = unsafe { super::pack_avx2(&bits) };
652            let expected = pack(&bits).data;
653            assert_eq!(got, expected, "Mismatch at length={length}");
654        }
655    }
656
657    #[test]
658    fn fused_and_popcount_avx2_matches_scalar() {
659        if !is_x86_feature_detected!("avx2") {
660            return;
661        }
662
663        let lengths = [1_usize, 7, 8, 15, 16, 17, 31, 32, 64, 128];
664        for len in lengths {
665            let a: Vec<u64> = (0..len)
666                .map(|i| (i as u64).wrapping_mul(0x9E37_79B9_7F4A_7C15) ^ 0xA5A5_A5A5_5A5A_5A5A)
667                .collect();
668            let b: Vec<u64> = (0..len)
669                .map(|i| (i as u64).wrapping_mul(0xC2B2_AE3D_27D4_EB4F) ^ 0x0F0F_F0F0_33CC_CC33)
670                .collect();
671
672            let expected: u64 = a
673                .iter()
674                .zip(b.iter())
675                .map(|(&wa, &wb)| (wa & wb).count_ones() as u64)
676                .sum();
677
678            // SAFETY: Runtime-guarded by feature detection in this test.
679            let got = unsafe { super::fused_and_popcount_avx2(&a, &b) };
680            assert_eq!(got, expected, "Mismatch at len={len}");
681        }
682    }
683
684    #[test]
685    fn dot_f64_avx2_matches_scalar() {
686        if !is_x86_feature_detected!("avx2") || !is_x86_feature_detected!("fma") {
687            return;
688        }
689        let a: Vec<f64> = (0..67).map(|i| i as f64 * 0.1).collect();
690        let b: Vec<f64> = (0..67).map(|i| (i as f64 * 0.3) - 5.0).collect();
691        let expected: f64 = a.iter().zip(&b).map(|(x, y)| x * y).sum();
692        let got = unsafe { super::dot_f64_avx2(&a, &b) };
693        assert!(
694            (got - expected).abs() < 1e-9,
695            "dot: got {got}, expected {expected}"
696        );
697    }
698
699    #[test]
700    fn max_f64_avx2_matches_scalar() {
701        if !is_x86_feature_detected!("avx2") {
702            return;
703        }
704        let a: Vec<f64> = (0..67).map(|i| (i as f64 * 7.3).sin()).collect();
705        let expected = a.iter().copied().fold(f64::NEG_INFINITY, f64::max);
706        let got = unsafe { super::max_f64_avx2(&a) };
707        assert!(
708            (got - expected).abs() < 1e-12,
709            "max: got {got}, expected {expected}"
710        );
711    }
712
713    #[test]
714    fn sum_f64_avx2_matches_scalar() {
715        if !is_x86_feature_detected!("avx2") {
716            return;
717        }
718        let a: Vec<f64> = (0..67).map(|i| i as f64 * 0.01).collect();
719        let expected: f64 = a.iter().sum();
720        let got = unsafe { super::sum_f64_avx2(&a) };
721        assert!(
722            (got - expected).abs() < 1e-9,
723            "sum: got {got}, expected {expected}"
724        );
725    }
726
727    #[test]
728    fn softmax_avx2_sums_to_one() {
729        if !is_x86_feature_detected!("avx2") {
730            return;
731        }
732        let mut scores: Vec<f64> = (0..67).map(|i| (i as f64 * 0.3) - 10.0).collect();
733        unsafe { super::softmax_inplace_f64_avx2(&mut scores) };
734        let sum: f64 = scores.iter().sum();
735        assert!(
736            (sum - 1.0).abs() < 1e-10,
737            "softmax must sum to 1.0, got {sum}"
738        );
739        assert!(scores.iter().all(|&s| s >= 0.0), "all values must be >= 0");
740    }
741
742    #[test]
743    fn bernoulli_compare_avx2_matches_scalar() {
744        if !is_x86_feature_detected!("avx2") {
745            return;
746        }
747
748        let buf: Vec<u8> = (0..32).map(|i| (i * 73 + 17) as u8).collect();
749        let thresholds = [0_u8, 1, 2, 17, 64, 127, 128, 200, 255];
750
751        for threshold in thresholds {
752            let expected = buf.iter().enumerate().fold(0_u32, |acc, (bit, &rb)| {
753                acc | (u32::from(rb < threshold) << bit)
754            });
755
756            // SAFETY: Runtime-guarded by feature detection in this test.
757            let got = unsafe { super::bernoulli_compare_avx2(&buf, threshold) };
758            assert_eq!(
759                got, expected,
760                "Mismatch for threshold={threshold} buf={buf:?}"
761            );
762        }
763    }
764
765    #[test]
766    fn dot_f64_avx_matches_scalar() {
767        if !is_x86_feature_detected!("avx") {
768            return;
769        }
770        let a: Vec<f64> = (0..67).map(|i| i as f64 * 0.1).collect();
771        let b: Vec<f64> = (0..67).map(|i| (i as f64 * 0.3) - 5.0).collect();
772        let expected: f64 = a.iter().zip(&b).map(|(x, y)| x * y).sum();
773        let got = unsafe { super::dot_f64_avx(&a, &b) };
774        assert!(
775            (got - expected).abs() < 1e-9,
776            "dot_avx: got {got}, expected {expected}"
777        );
778    }
779}