Skip to main content

hermes_simd/dispatch/
scale.rs

1//! Generic runtime-dispatch in-place scale kernel.
2//!
3//! `scale_in_place<T>(data, scalar)` broadcasts `scalar` to all SIMD lanes
4//! then multiplies each chunk of `data` by it, covering the scalar tail
5//! element-by-element.
6
7use 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        // ── 4× unrolled SIMD loop — hides load/store latency ────────────────
30        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        // ── Remaining full SIMD vectors ─────────────────────────────────────
44        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    // Scalar tail
53    let simd_len = (len / lane_count) * lane_count;
54    for i in simd_len..len {
55        data[i] = data[i] * scalar;
56    }
57}