1#[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#[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 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 a * b + c
64 }
65 #[inline(always)]
66 unsafe fn fnma(self, a: Self::Reg, b: Self::Reg, c: Self::Reg) -> Self::Reg {
67 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 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#[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#[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#[cfg(feature = "complex")]
203impl_complex_simd!(ScalarTok, f32, f32, 1);
204#[cfg(feature = "complex")]
205impl_complex_simd!(ScalarTok, f64, f64, 1);