sc_neurocore_engine/simd/
mod.rs1use rand::Rng;
15
16pub mod avx2;
17pub mod avx512;
18pub mod neon;
19pub mod rvv;
20pub mod sve;
21
22pub 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 let data = unsafe { avx512::pack_avx512(bits) };
31 return crate::bitstream::BitStreamTensor { data, length };
32 }
33 if is_x86_feature_detected!("avx2") {
34 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 let data = unsafe { sve::pack_sve(bits) };
44 return crate::bitstream::BitStreamTensor { data, length };
45 }
46
47 crate::bitstream::pack_fast(bits)
48}
49
50pub fn popcount_dispatch(data: &[u64]) -> u64 {
52 #[cfg(target_arch = "x86_64")]
53 {
54 if is_x86_feature_detected!("avx512vpopcntdq") {
55 return unsafe { avx512::popcount_avx512(data) };
57 }
58 if is_x86_feature_detected!("avx2") {
59 return unsafe { avx2::popcount_avx2(data) };
61 }
62 }
63
64 #[cfg(target_arch = "aarch64")]
65 {
66 #[cfg(target_feature = "sve")]
67 {
68 return unsafe { sve::popcount_sve(data) };
70 }
71 #[cfg(not(target_feature = "sve"))]
72 {
73 return unsafe { neon::popcount_neon(data) };
75 }
76 }
77
78 #[cfg(all(target_arch = "riscv64", target_feature = "v"))]
79 {
80 return unsafe { rvv::popcount_rvv(data) };
82 }
83
84 crate::bitstream::popcount_words_portable(data)
85}
86
87pub 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 return unsafe { avx512::fused_and_popcount_avx512(a, b) };
98 }
99 if is_x86_feature_detected!("avx2") {
100 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
135pub 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 return unsafe { avx512::fused_xor_popcount_avx512(a, b) };
146 }
147 if is_x86_feature_detected!("avx2") {
148 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
199pub 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
240pub 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
274pub 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
306pub 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
348pub fn hamming_distance_dispatch(a: &[u64], b: &[u64]) -> u64 {
350 fused_xor_popcount_dispatch(a, b)
351}
352
353pub 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
388pub 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
401pub 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 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 #[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}