sc_neurocore_engine/simd/
avx512.rs1#[cfg(target_arch = "x86_64")]
10use core::arch::x86_64::*;
11
12#[cfg(target_arch = "x86_64")]
13#[target_feature(enable = "avx512f,avx512vpopcntdq")]
14pub 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")]
35pub 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")]
76pub 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")]
109pub 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"))]
141pub 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")]
154pub 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"))]
169pub unsafe fn popcount_avx512(data: &[u64]) -> u64 {
174 crate::bitstream::popcount_words_portable(data)
175}
176
177#[cfg(not(target_arch = "x86_64"))]
178pub unsafe fn pack_avx512(bits: &[u8]) -> Vec<u64> {
183 crate::bitstream::pack_fast(bits).data
184}
185
186#[cfg(not(target_arch = "x86_64"))]
187pub 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"))]
199pub 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#[cfg(target_arch = "x86_64")]
216#[target_feature(enable = "avx512f")]
217pub 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")]
242pub 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")]
268pub 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")]
291pub 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")]
336pub 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 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 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 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 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}