1#[cfg(target_arch = "x86_64")]
10use core::arch::x86_64::*;
11
12#[cfg(target_arch = "x86_64")]
13#[target_feature(enable = "avx2")]
14pub 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")]
48pub 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")]
96pub 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")]
142pub 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"))]
187pub 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")]
200pub 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"))]
219pub unsafe fn popcount_avx2(data: &[u64]) -> u64 {
224 crate::bitstream::popcount_words_portable(data)
225}
226
227#[cfg(not(target_arch = "x86_64"))]
228pub unsafe fn pack_avx2(bits: &[u8]) -> Vec<u64> {
233 crate::bitstream::pack_fast(bits).data
234}
235
236#[cfg(not(target_arch = "x86_64"))]
237pub 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"))]
249pub 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#[cfg(target_arch = "x86_64")]
266#[target_feature(enable = "avx2,fma")]
267pub 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")]
295pub 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")]
322pub 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")]
346pub 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
394pub 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")]
404pub 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 (chunks, remainder) = scores.as_chunks_mut::<16>();
416 for chunk in chunks {
417 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 (chunks, remainder) = scores.as_chunks_mut::<16>();
442 for chunk in chunks {
443 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")]
461pub 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")]
509pub 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")]
542pub unsafe fn bernoulli_compare_batch_avx2(buf: &[u8], threshold: u8, out: &mut [u64]) {
547 let v_thresh = _mm256_set1_epi8(threshold as i8);
548 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 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")]
572pub 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")]
608pub 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 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 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 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}