1use super::{complex, modular, simd_ops::SimdOps};
2use hermes_simd_core::scalar::Scalar as ScalarTrait;
3use hermes_simd_core::sparse::{
4 BlockedCooData, CsrData, DenseWithMaskData, SellPData, ValidatedData,
5};
6use hermes_simd_core::view::SimdError;
7
8#[inline(always)]
10pub fn sum<T: SimdOps>(data: &[T]) -> T {
11 T::sum(data)
12}
13
14#[inline(always)]
18pub fn min<T: SimdOps>(data: &[T]) -> T {
19 T::min(data)
20}
21
22#[inline(always)]
26pub fn max<T: SimdOps>(data: &[T]) -> T {
27 T::max(data)
28}
29
30#[inline(always)]
32pub fn abs_sum<T: SimdOps>(data: &[T]) -> T {
33 T::abs_sum(data)
34}
35
36#[inline(always)]
38pub fn abs_max<T: SimdOps>(data: &[T]) -> T {
39 T::abs_max(data)
40}
41
42#[inline(always)]
44pub fn scale<T: SimdOps>(data: &mut [T], scalar: T) {
45 T::scale(data, scalar)
46}
47
48#[inline(always)]
50pub fn argmin<T: SimdOps>(data: &[T]) -> Option<(usize, T)> {
51 T::argmin(data)
52}
53
54#[inline(always)]
56pub fn argmax<T: SimdOps>(data: &[T]) -> Option<(usize, T)> {
57 T::argmax(data)
58}
59
60#[inline(always)]
62pub fn dot<T: SimdOps>(a: &[T], b: &[T]) -> Result<T, SimdError> {
63 T::dot(a, b)
64}
65
66#[inline(always)]
69pub fn axpy<T: SimdOps>(alpha: T, x: &[T], out: &mut [T]) -> Result<(), SimdError> {
70 T::axpy(alpha, x, out)
71}
72
73#[inline(always)]
79pub fn axpy_mul<T: SimdOps>(alpha: T, a: &[T], b: &[T], out: &mut [T]) -> Result<(), SimdError> {
80 T::axpy_mul(alpha, a, b, out)
81}
82
83#[inline(always)]
86pub fn axpy_rows<T: SimdOps>(
87 alphas: &[T],
88 x: &[T],
89 out: &mut [T],
90 row_stride: usize,
91 rows: usize,
92 cols: usize,
93) -> Result<(), SimdError> {
94 T::axpy_rows(alphas, x, out, row_stride, rows, cols)
95}
96
97#[inline(always)]
103pub fn axpy_rows_batch<T: SimdOps>(
104 alphas: &[T],
105 x_panel: &[T],
106 out: &mut [T],
107 row_stride: usize,
108 rows: usize,
109 depth: usize,
110 cols: usize,
111) -> Result<(), SimdError> {
112 T::axpy_rows_batch(alphas, x_panel, out, row_stride, rows, depth, cols)
113}
114
115#[inline(always)]
117pub fn elementwise_mul<T: SimdOps>(a: &[T], b: &[T], out: &mut [T]) -> Result<(), SimdError> {
118 T::elementwise_mul(a, b, out)
119}
120
121#[inline(always)]
123pub fn elementwise_add<T: SimdOps>(a: &[T], b: &[T], out: &mut [T]) -> Result<(), SimdError> {
124 T::elementwise_add(a, b, out)
125}
126
127#[inline(always)]
129pub fn elementwise_sub<T: SimdOps>(a: &[T], b: &[T], out: &mut [T]) -> Result<(), SimdError> {
130 T::elementwise_sub(a, b, out)
131}
132
133#[inline(always)]
135pub fn elementwise_div<T: SimdOps>(a: &[T], b: &[T], out: &mut [T]) -> Result<(), SimdError> {
136 T::elementwise_div(a, b, out)
137}
138
139#[inline]
141pub fn ntt_butterfly_stage_u64(
142 data: &mut [u64],
143 stage_len: usize,
144 twiddles: &[u64],
145 modulus: u64,
146) -> Result<(), SimdError> {
147 modular::ntt_butterfly_stage_u64(data, stage_len, twiddles, modulus)
148}
149
150#[inline(always)]
152pub fn masked_sum<T: SimdOps>(data: &[T], mask: &[bool]) -> T {
153 T::masked_sum(data, mask)
154}
155
156#[inline(always)]
158pub fn masked_dot<T: SimdOps>(a: &[T], b: &[T], mask: &[bool]) -> Result<T, SimdError> {
159 T::masked_dot(a, b, mask)
160}
161
162#[inline(always)]
164pub fn masked_add<T: SimdOps>(
165 a: &[T],
166 b: &[T],
167 mask: &[bool],
168 out: &mut [T],
169) -> Result<(), SimdError> {
170 T::masked_add(a, b, mask, out)
171}
172
173#[inline(always)]
179pub fn spmv_csr<T: SimdOps>(data: ValidatedData<CsrData<'_, T>>, x: &[T], y: &mut [T]) {
180 T::spmv_csr(data, x, y)
181}
182
183#[inline(always)]
190pub fn spmv_bcoo<T: SimdOps, const BM: usize, const BN: usize>(
191 data: ValidatedData<BlockedCooData<'_, T, BM, BN>>,
192 x: &[T],
193 y: &mut [T],
194) {
195 T::spmv_bcoo::<BM, BN>(data, x, y)
196}
197
198#[inline(always)]
200pub fn spmv_dense_masked<T: SimdOps>(data: DenseWithMaskData<'_, T>, x: &[T], y: &mut [T]) {
201 T::spmv_dense_masked(data, x, y)
202}
203
204#[inline(always)]
211pub fn spmv_sellp<T: SimdOps, const C: usize>(
212 data: ValidatedData<SellPData<'_, T, C>>,
213 x: &[T],
214 y: &mut [T],
215) {
216 T::spmv_sellp::<C>(data, x, y)
217}
218
219#[inline(always)]
221pub fn tiled_gemm<T: SimdOps>(
222 a: &[T],
223 b: &[T],
224 c: &mut [T],
225 m: usize,
226 n: usize,
227 k: usize,
228) -> Result<(), SimdError> {
229 T::tiled_gemm(a, b, c, m, n, k)
230}
231
232#[inline(always)]
242pub fn gemv<T: SimdOps>(
243 a: &[T],
244 x: &[T],
245 y: &mut [T],
246 nrows: usize,
247 ncols: usize,
248) -> Result<(), SimdError> {
249 T::gemv(a, x, y, nrows, ncols)
250}
251
252#[inline(always)]
263pub fn gemv_transpose<T: SimdOps>(
264 a: &[T],
265 x: &[T],
266 y: &mut [T],
267 nrows: usize,
268 ncols: usize,
269) -> Result<(), SimdError> {
270 T::gemv_transpose(a, x, y, nrows, ncols)
271}
272
273#[inline(always)]
282pub fn gemv_strided<T: SimdOps>(
283 a: &[T],
284 x: &[T],
285 y: &mut [T],
286 nrows: usize,
287 ncols: usize,
288 lda: usize,
289) -> Result<(), SimdError> {
290 T::gemv_strided(a, x, y, nrows, ncols, lda)
291}
292
293#[inline(always)]
301pub fn gemv_transpose_strided<T: SimdOps>(
302 a: &[T],
303 x: &[T],
304 y: &mut [T],
305 nrows: usize,
306 ncols: usize,
307 lda: usize,
308) -> Result<(), SimdError> {
309 T::gemv_transpose_strided(a, x, y, nrows, ncols, lda)
310}
311
312#[inline]
318pub fn interleaved_complex_mul_assign<T, A, const CONJ_B: bool>(
319 a: &mut [T],
320 b: &[T],
321) -> Result<(), SimdError>
322where
323 T: ScalarTrait + core::ops::Neg<Output = T>,
324 A: hermes_simd_core::arch::SimdArch + hermes_simd_core::kernel::SimdKernel<T>,
325{
326 complex::interleaved_complex_mul_assign::<T, A, CONJ_B>(a, b)
327}
328
329#[inline]
335pub fn interleaved_complex_dot<T, A, const CONJ_B: bool>(
336 a: &[T],
337 b: &[T],
338) -> Result<(T, T), SimdError>
339where
340 T: ScalarTrait + core::ops::Neg<Output = T>,
341 A: hermes_simd_core::arch::SimdArch + hermes_simd_core::kernel::SimdKernel<T>,
342{
343 complex::interleaved_complex_dot::<T, A, CONJ_B>(a, b)
344}
345
346#[inline]
348pub fn interleaved_complex_mul_assign_runtime<T, const CONJ_B: bool>(
349 a: &mut [T],
350 b: &[T],
351) -> Result<(), SimdError>
352where
353 T: SimdOps + core::ops::Neg<Output = T>,
354{
355 T::interleaved_complex_mul_assign::<CONJ_B>(a, b)
356}
357
358#[inline]
360pub fn interleaved_complex_dot_runtime<T, const CONJ_B: bool>(
361 a: &[T],
362 b: &[T],
363) -> Result<(T, T), SimdError>
364where
365 T: SimdOps + core::ops::Neg<Output = T>,
366{
367 T::interleaved_complex_dot::<CONJ_B>(a, b)
368}
369
370#[inline(always)]
372pub fn reduce_popcount<T: SimdOps>(data: &[T]) -> usize {
373 T::reduce_popcount(data)
374}
375
376#[inline(always)]
378pub fn reduce_popcount_and<T: SimdOps>(a: &[T], b: &[T]) -> Result<usize, SimdError> {
379 T::reduce_popcount_and(a, b)
380}
381
382#[inline(always)]
384pub fn reduce_popcount_or<T: SimdOps>(a: &[T], b: &[T]) -> Result<usize, SimdError> {
385 T::reduce_popcount_or(a, b)
386}
387
388#[inline(always)]
390pub fn reduce_popcount_xor<T: SimdOps>(a: &[T], b: &[T]) -> Result<usize, SimdError> {
391 T::reduce_popcount_xor(a, b)
392}