hermes_simd/dispatch/
scale.rs1use hermes_simd_core::{arch::SimdArch, kernel::SimdKernel, scalar::Scalar};
8use hermes_simd_macros::runtime_dispatch;
9
10#[runtime_dispatch(avx512f, avx2, neon, scalar)]
11pub(super) fn dispatch_scale_kernel<T, A>(data: &mut [T], scalar: T)
12where
13 T: Scalar,
14 A: SimdArch + SimdKernel<T>,
15{
16 let len = data.len();
17 if len == 0 {
18 return;
19 }
20 let lane_count = A::LANE_COUNT;
21 let unroll_factor = A::UNROLL_FACTOR;
22 let chunk_size = lane_count * unroll_factor;
23 let unrolled_simd_len = (len / chunk_size) * chunk_size;
24 let ptr = data.as_mut_ptr();
25
26 unsafe {
27 let vsplat = A::splat(scalar);
28
29 let mut i = 0usize;
31 while i < unrolled_simd_len {
32 let p0 = ptr.add(i);
33 let p1 = ptr.add(i + lane_count);
34 let p2 = ptr.add(i + lane_count * 2);
35 let p3 = ptr.add(i + lane_count * 3);
36 A::store_unaligned(p0, A::mul(A::load_unaligned(p0), vsplat));
37 A::store_unaligned(p1, A::mul(A::load_unaligned(p1), vsplat));
38 A::store_unaligned(p2, A::mul(A::load_unaligned(p2), vsplat));
39 A::store_unaligned(p3, A::mul(A::load_unaligned(p3), vsplat));
40 i += chunk_size;
41 }
42
43 let simd_len = (len / lane_count) * lane_count;
45 while i < simd_len {
46 let p = ptr.add(i);
47 A::store_unaligned(p, A::mul(A::load_unaligned(p), vsplat));
48 i += lane_count;
49 }
50 }
51
52 let simd_len = (len / lane_count) * lane_count;
54 for i in simd_len..len {
55 data[i] = data[i] * scalar;
56 }
57}