Skip to main content

gemmkit/simd/
scalar.rs

1//! Scalar (1-element) ISA token: no target features, no intrinsics, compiles anywhere
2//!
3//! `LANES = 1` for every `SimdOps` impl here, so a "register" is just the element
4//! itself. This is the portability floor and the Miri-checkable reference. The same
5//! generic microkernel that runs on AVX-512 runs unchanged on this token, one lane
6//! at a time
7
8#[cfg(feature = "half")]
9use half::{bf16, f16};
10
11#[cfg(any(feature = "half", feature = "int8"))]
12use super::KernelSimd;
13use super::{Simd, SimdOps};
14#[cfg(feature = "half")]
15use crate::scalar::NarrowFloat;
16
17/// The scalar (1-lane) ISA token, available on every target
18#[derive(Copy, Clone, Default)]
19pub struct ScalarTok;
20
21impl Simd for ScalarTok {
22    #[inline(always)]
23    unsafe fn vectorize<R>(self, f: impl FnOnce() -> R) -> R {
24        // No target feature to enable, so vectorize is just a call
25        f()
26    }
27}
28
29macro_rules! impl_scalar_ops {
30    ($t:ty) => {
31        impl SimdOps<$t> for ScalarTok {
32            type Reg = $t;
33            const LANES: usize = 1;
34
35            #[inline(always)]
36            unsafe fn zero(self) -> Self::Reg {
37                0.0
38            }
39            #[inline(always)]
40            unsafe fn splat(self, v: $t) -> Self::Reg {
41                v
42            }
43            #[inline(always)]
44            unsafe fn loadu(self, p: *const $t) -> Self::Reg {
45                unsafe { *p }
46            }
47            #[inline(always)]
48            unsafe fn storeu(self, p: *mut $t, v: Self::Reg) {
49                unsafe { *p = v }
50            }
51            #[inline(always)]
52            unsafe fn mul(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
53                a * b
54            }
55            #[inline(always)]
56            unsafe fn add(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
57                a + b
58            }
59            #[inline(always)]
60            unsafe fn mul_add(self, a: Self::Reg, b: Self::Reg, c: Self::Reg) -> Self::Reg {
61                // Plain multiply-then-add: the reference `mul_add` that every vector
62                // token's true FMA is checked against
63                a * b + c
64            }
65            #[inline(always)]
66            unsafe fn fnma(self, a: Self::Reg, b: Self::Reg, c: Self::Reg) -> Self::Reg {
67                // Plain `c - a*b`: the reference `fnma`, and the scalar model for the
68                // SoA complex kernel's `acc_re -= ai*bi` step
69                c - a * b
70            }
71            #[inline(always)]
72            unsafe fn reduce_sum(self, v: Self::Reg) -> $t {
73                v
74            }
75            #[inline(always)]
76            unsafe fn max(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
77                // `NaN > b` is always false, so a NaN `a` falls through to `b`, matching
78                // the trait's NaN-in-`a` contract
79                if a > b { a } else { b }
80            }
81            #[inline(always)]
82            unsafe fn min(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
83                if a < b { a } else { b }
84            }
85        }
86    };
87}
88
89impl_scalar_ops!(f32);
90impl_scalar_ops!(f64);
91
92// Mixed precision (scalar fallback): f16 and bf16 widen to f32 on load and round to
93// f16 or bf16 on store, one element at a time. `Reg` here is a bare `f32`
94#[cfg(feature = "half")]
95impl KernelSimd<f16, f16, f32, f16> for ScalarTok {
96    #[inline(always)]
97    unsafe fn load_lhs(self, p: *const f16) -> f32 {
98        unsafe { (*p).widen() }
99    }
100    #[inline(always)]
101    unsafe fn splat_rhs(self, v: f16) -> f32 {
102        v.widen()
103    }
104    #[inline(always)]
105    unsafe fn load_out(self, p: *const f16) -> f32 {
106        unsafe { (*p).widen() }
107    }
108    #[inline(always)]
109    unsafe fn store_out(self, p: *mut f16, v: f32) {
110        unsafe { *p = f16::narrow(v) }
111    }
112}
113
114// Integer GEMM (scalar fallback): i32 accumulator ops and the i8 -> i32 widen-load,
115// one element at a time. Wrapping arithmetic matches the modular semantics the
116// vector `mullo`/`add` intrinsics wrap under
117#[cfg(feature = "int8")]
118impl SimdOps<i32> for ScalarTok {
119    type Reg = i32;
120    const LANES: usize = 1;
121
122    #[inline(always)]
123    unsafe fn zero(self) -> i32 {
124        0
125    }
126    #[inline(always)]
127    unsafe fn splat(self, v: i32) -> i32 {
128        v
129    }
130    #[inline(always)]
131    unsafe fn loadu(self, p: *const i32) -> i32 {
132        unsafe { *p }
133    }
134    #[inline(always)]
135    unsafe fn storeu(self, p: *mut i32, v: i32) {
136        unsafe { *p = v }
137    }
138    #[inline(always)]
139    unsafe fn mul(self, a: i32, b: i32) -> i32 {
140        a.wrapping_mul(b)
141    }
142    #[inline(always)]
143    unsafe fn add(self, a: i32, b: i32) -> i32 {
144        a.wrapping_add(b)
145    }
146    #[inline(always)]
147    unsafe fn mul_add(self, a: i32, b: i32, c: i32) -> i32 {
148        a.wrapping_mul(b).wrapping_add(c)
149    }
150    #[inline(always)]
151    unsafe fn fnma(self, a: i32, b: i32, c: i32) -> i32 {
152        c.wrapping_sub(a.wrapping_mul(b))
153    }
154    #[inline(always)]
155    unsafe fn reduce_sum(self, v: i32) -> i32 {
156        v
157    }
158}
159
160#[cfg(feature = "int8")]
161impl KernelSimd<i8, i8, i32, i32> for ScalarTok {
162    #[inline(always)]
163    unsafe fn load_lhs(self, p: *const i8) -> i32 {
164        unsafe { *p as i32 }
165    }
166    #[inline(always)]
167    unsafe fn splat_rhs(self, v: i8) -> i32 {
168        v as i32
169    }
170    #[inline(always)]
171    unsafe fn load_out(self, p: *const i32) -> i32 {
172        unsafe { *p }
173    }
174    #[inline(always)]
175    unsafe fn store_out(self, p: *mut i32, v: i32) {
176        unsafe { *p = v }
177    }
178}
179
180#[cfg(feature = "half")]
181impl KernelSimd<bf16, bf16, f32, bf16> for ScalarTok {
182    #[inline(always)]
183    unsafe fn load_lhs(self, p: *const bf16) -> f32 {
184        unsafe { (*p).widen() }
185    }
186    #[inline(always)]
187    unsafe fn splat_rhs(self, v: bf16) -> f32 {
188        v.widen()
189    }
190    #[inline(always)]
191    unsafe fn load_out(self, p: *const bf16) -> f32 {
192        unsafe { (*p).widen() }
193    }
194    #[inline(always)]
195    unsafe fn store_out(self, p: *mut bf16, v: f32) {
196        unsafe { *p = bf16::narrow(v) }
197    }
198}
199
200// Complex (scalar fallback): LANES = 1, the real Reg is the scalar itself, and complex
201// GEMM routes through the shared soa_microkernel like every other token
202#[cfg(feature = "complex")]
203impl_complex_simd!(ScalarTok, f32, f32, 1);
204#[cfg(feature = "complex")]
205impl_complex_simd!(ScalarTok, f64, f64, 1);