Skip to main content

embedded_nn/
simd.rs

1//! Target SIMD vectorization abstractions and acceleration hooks.
2
3/// Vectorized dot product for int8 slices with an input zero-point offset (`lhs_offset`).
4///
5/// Efficiently processes 4-element or 8-element chunks to trigger compiler auto-vectorization
6/// or ARM SIMD instructions (`SMLAD` on Cortex-M4/M7, `vmladavaq` on Cortex-M55/M85).
7#[inline]
8pub fn vec_dot_s8(lhs: &[i8], rhs: &[i8], lhs_offset: i32) -> i32 {
9    let len = lhs.len().min(rhs.len());
10    let mut acc: i32 = 0;
11
12    #[cfg(all(target_arch = "arm", target_feature = "dsp"))]
13    {
14        // Target ARM DSP hardware acceleration (Cortex-M4/M7 SMLAD)
15        let chunks = len / 2;
16        let remainder = len % 2;
17
18        let mut i = 0;
19        for _ in 0..chunks {
20            let l0 = lhs[i] as i32 + lhs_offset;
21            let l1 = lhs[i + 1] as i32 + lhs_offset;
22            let r0 = rhs[i] as i32;
23            let r1 = rhs[i + 1] as i32;
24
25            let val_l = (l0 & 0xFFFF) | ((l1 & 0xFFFF) << 16);
26            let val_r = (r0 & 0xFFFF) | ((r1 & 0xFFFF) << 16);
27
28            let res: i32;
29            unsafe {
30                core::arch::asm!(
31                    "smlad {0}, {1}, {2}, {3}",
32                    out(reg) res,
33                    in(reg) val_l,
34                    in(reg) val_r,
35                    in(reg) acc,
36                );
37            }
38            acc = res;
39            i += 2;
40        }
41
42        for j in 0..remainder {
43            let l = lhs[i + j] as i32 + lhs_offset;
44            let r = rhs[i + j] as i32;
45            acc += l * r;
46        }
47
48        return acc;
49    }
50
51    #[cfg(not(all(target_arch = "arm", target_feature = "dsp")))]
52    {
53        let chunks = len / 8;
54        let remainder = len % 8;
55
56        let mut i = 0;
57        for _ in 0..chunks {
58            let l0 = lhs[i] as i32 + lhs_offset;
59            let r0 = rhs[i] as i32;
60            let l1 = lhs[i + 1] as i32 + lhs_offset;
61            let r1 = rhs[i + 1] as i32;
62            let l2 = lhs[i + 2] as i32 + lhs_offset;
63            let r2 = rhs[i + 2] as i32;
64            let l3 = lhs[i + 3] as i32 + lhs_offset;
65            let r3 = rhs[i + 3] as i32;
66            let l4 = lhs[i + 4] as i32 + lhs_offset;
67            let r4 = rhs[i + 4] as i32;
68            let l5 = lhs[i + 5] as i32 + lhs_offset;
69            let r5 = rhs[i + 5] as i32;
70            let l6 = lhs[i + 6] as i32 + lhs_offset;
71            let r6 = rhs[i + 6] as i32;
72            let l7 = lhs[i + 7] as i32 + lhs_offset;
73            let r7 = rhs[i + 7] as i32;
74
75            acc += l0 * r0 + l1 * r1 + l2 * r2 + l3 * r3 + l4 * r4 + l5 * r5 + l6 * r6 + l7 * r7;
76            i += 8;
77        }
78
79        for j in 0..remainder {
80            let l = lhs[i + j] as i32 + lhs_offset;
81            let r = rhs[i + j] as i32;
82            acc += l * r;
83        }
84
85        acc
86    }
87}
88
89/// Vectorized dot product for int16 slices.
90#[inline]
91pub fn vec_dot_s16(lhs: &[i16], rhs: &[i16]) -> i64 {
92    let len = lhs.len().min(rhs.len());
93    let mut acc: i64 = 0;
94
95    let chunks = len / 4;
96    let remainder = len % 4;
97
98    let mut i = 0;
99    for _ in 0..chunks {
100        let l0 = lhs[i] as i64;
101        let r0 = rhs[i] as i64;
102        let l1 = lhs[i + 1] as i64;
103        let r1 = rhs[i + 1] as i64;
104        let l2 = lhs[i + 2] as i64;
105        let r2 = rhs[i + 2] as i64;
106        let l3 = lhs[i + 3] as i64;
107        let r3 = rhs[i + 3] as i64;
108
109        acc += l0 * r0 + l1 * r1 + l2 * r2 + l3 * r3;
110        i += 4;
111    }
112
113    for j in 0..remainder {
114        acc += (lhs[i + j] as i64) * (rhs[i + j] as i64);
115    }
116
117    acc
118}
119
120#[cfg(test)]
121mod tests {
122    use super::*;
123
124    #[test]
125    fn test_vec_dot_s8() {
126        let lhs = [1i8, 2i8, 3i8, 4i8, 5i8, 6i8, 7i8, 8i8, 9i8];
127        let rhs = [1i8, 1i8, 1i8, 1i8, 1i8, 1i8, 1i8, 1i8, 1i8];
128        assert_eq!(vec_dot_s8(&lhs, &rhs, 0), 45);
129        assert_eq!(vec_dot_s8(&lhs, &rhs, 1), 45 + 9);
130    }
131
132    #[test]
133    fn test_vec_dot_s16() {
134        let lhs = [10i16, 20i16, 30i16, 40i16, 50i16];
135        let rhs = [1i16, 2i16, 3i16, 4i16, 5i16];
136        assert_eq!(vec_dot_s16(&lhs, &rhs), 550);
137    }
138}