Skip to content

Commit 632ef87

Browse files
authored
Merge pull request #5990 from jjerphan/wasm128-gemv
Add WASM SIMD128 SGEMV/DGEMV kernels
2 parents 4919204 + b47183a commit 632ef87

4 files changed

Lines changed: 455 additions & 4 deletions

File tree

kernel/wasm/KERNEL

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -89,13 +89,21 @@ DSWAPKERNEL = ../riscv64/swap.c
8989
CSWAPKERNEL = ../riscv64/zswap.c
9090
ZSWAPKERNEL = ../riscv64/zswap.c
9191

92+
ifndef SGEMVNKERNEL
9293
SGEMVNKERNEL = ../riscv64/gemv_n.c
94+
endif
95+
ifndef DGEMVNKERNEL
9396
DGEMVNKERNEL = ../riscv64/gemv_n.c
97+
endif
9498
CGEMVNKERNEL = ../riscv64/zgemv_n.c
9599
ZGEMVNKERNEL = ../riscv64/zgemv_n.c
96100

101+
ifndef SGEMVTKERNEL
97102
SGEMVTKERNEL = ../riscv64/gemv_t.c
103+
endif
104+
ifndef DGEMVTKERNEL
98105
DGEMVTKERNEL = ../riscv64/gemv_t.c
106+
endif
99107
CGEMVTKERNEL = ../riscv64/zgemv_t.c
100108
ZGEMVTKERNEL = ../riscv64/zgemv_t.c
101109

kernel/wasm/KERNEL.WASM128_GENERIC

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -85,13 +85,13 @@ DSWAPKERNEL = ../riscv64/swap.c
8585
CSWAPKERNEL = ../riscv64/zswap.c
8686
ZSWAPKERNEL = ../riscv64/zswap.c
8787

88-
SGEMVNKERNEL = ../riscv64/gemv_n.c
89-
DGEMVNKERNEL = ../riscv64/gemv_n.c
88+
SGEMVNKERNEL = gemv_n.c
89+
DGEMVNKERNEL = gemv_n.c
9090
CGEMVNKERNEL = ../riscv64/zgemv_n.c
9191
ZGEMVNKERNEL = ../riscv64/zgemv_n.c
9292

93-
SGEMVTKERNEL = ../riscv64/gemv_t.c
94-
DGEMVTKERNEL = ../riscv64/gemv_t.c
93+
SGEMVTKERNEL = gemv_t.c
94+
DGEMVTKERNEL = gemv_t.c
9595
CGEMVTKERNEL = ../riscv64/zgemv_t.c
9696
ZGEMVTKERNEL = ../riscv64/zgemv_t.c
9797

kernel/wasm/gemv_n.c

Lines changed: 196 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,196 @@
1+
/***************************************************************************
2+
Copyright (c) 2026, The OpenBLAS Project
3+
All rights reserved.
4+
Redistribution and use in source and binary forms, with or without
5+
modification, are permitted provided that the following conditions are
6+
met:
7+
1. Redistributions of source code must retain the above copyright
8+
notice, this list of conditions and the following disclaimer.
9+
2. Redistributions in binary form must reproduce the above copyright
10+
notice, this list of conditions and the following disclaimer in
11+
the documentation and/or other materials provided with the
12+
distribution.
13+
3. Neither the name of the OpenBLAS project nor the names of
14+
its contributors may be used to endorse or promote products
15+
derived from this software without specific prior written permission.
16+
THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
17+
AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
18+
IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE
19+
ARE DISCLAIMED. IN NO EVENT SHALL THE OPENBLAS PROJECT OR CONTRIBUTORS BE
20+
LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR
21+
CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF
22+
SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS
23+
INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN
24+
CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE)
25+
ARISING IN ANY WAY OUT OF THE USE OF THIS SOFTWARE, EVEN IF ADVISED OF THE
26+
POSSIBILITY OF SUCH DAMAGE.
27+
*****************************************************************************/
28+
29+
/*
30+
* WASM SIMD128 GEMV_N: y += alpha * A * x (AXPY each column of A into y).
31+
*
32+
* Compiled twice: SGEMV_N (-UDOUBLE) and DGEMV_N (-DDOUBLE).
33+
*
34+
* Four columns share one streaming of y. Inner loop unrolls four v128 lanes
35+
* (16 floats / 8 doubles). IEEE mul+add (not relaxed madd).
36+
* Non-unit inc_y stays scalar (no scatter).
37+
*/
38+
39+
#include "common.h"
40+
41+
#if defined(__wasm_simd128__)
42+
#include <wasm_simd128.h>
43+
44+
#ifdef DOUBLE
45+
#define GEMV_VLEN 2
46+
#define GEMV_SPLAT wasm_f64x2_splat
47+
#define GEMV_MUL wasm_f64x2_mul
48+
#define GEMV_ADD wasm_f64x2_add
49+
#else
50+
#define GEMV_VLEN 4
51+
#define GEMV_SPLAT wasm_f32x4_splat
52+
#define GEMV_MUL wasm_f32x4_mul
53+
#define GEMV_ADD wasm_f32x4_add
54+
#endif
55+
56+
#define GEMV_UNROLL 4
57+
#define GEMV_CHUNK (GEMV_VLEN * GEMV_UNROLL)
58+
#define GEMV_LOAD(p) wasm_v128_load((const void *)(p))
59+
#define GEMV_STORE(p, v) wasm_v128_store((void *)(p), (v))
60+
#define GEMV_MADD(y, a, x) GEMV_ADD((y), GEMV_MUL((a), (x)))
61+
62+
static void gemv_n_axpy4(BLASLONG m, const FLOAT *a0, const FLOAT *a1,
63+
const FLOAT *a2, const FLOAT *a3, FLOAT *y, FLOAT t0,
64+
FLOAT t1, FLOAT t2, FLOAT t3) {
65+
const v128_t v0 = GEMV_SPLAT(t0);
66+
const v128_t v1 = GEMV_SPLAT(t1);
67+
const v128_t v2 = GEMV_SPLAT(t2);
68+
const v128_t v3 = GEMV_SPLAT(t3);
69+
BLASLONG i = 0;
70+
const BLASLONG n_main = m & ~(BLASLONG)(GEMV_CHUNK - 1);
71+
72+
for (; i < n_main; i += GEMV_CHUNK) {
73+
v128_t acc;
74+
acc = GEMV_LOAD(y + i + 0 * GEMV_VLEN);
75+
acc = GEMV_MADD(acc, v0, GEMV_LOAD(a0 + i + 0 * GEMV_VLEN));
76+
acc = GEMV_MADD(acc, v1, GEMV_LOAD(a1 + i + 0 * GEMV_VLEN));
77+
acc = GEMV_MADD(acc, v2, GEMV_LOAD(a2 + i + 0 * GEMV_VLEN));
78+
acc = GEMV_MADD(acc, v3, GEMV_LOAD(a3 + i + 0 * GEMV_VLEN));
79+
GEMV_STORE(y + i + 0 * GEMV_VLEN, acc);
80+
81+
acc = GEMV_LOAD(y + i + 1 * GEMV_VLEN);
82+
acc = GEMV_MADD(acc, v0, GEMV_LOAD(a0 + i + 1 * GEMV_VLEN));
83+
acc = GEMV_MADD(acc, v1, GEMV_LOAD(a1 + i + 1 * GEMV_VLEN));
84+
acc = GEMV_MADD(acc, v2, GEMV_LOAD(a2 + i + 1 * GEMV_VLEN));
85+
acc = GEMV_MADD(acc, v3, GEMV_LOAD(a3 + i + 1 * GEMV_VLEN));
86+
GEMV_STORE(y + i + 1 * GEMV_VLEN, acc);
87+
88+
acc = GEMV_LOAD(y + i + 2 * GEMV_VLEN);
89+
acc = GEMV_MADD(acc, v0, GEMV_LOAD(a0 + i + 2 * GEMV_VLEN));
90+
acc = GEMV_MADD(acc, v1, GEMV_LOAD(a1 + i + 2 * GEMV_VLEN));
91+
acc = GEMV_MADD(acc, v2, GEMV_LOAD(a2 + i + 2 * GEMV_VLEN));
92+
acc = GEMV_MADD(acc, v3, GEMV_LOAD(a3 + i + 2 * GEMV_VLEN));
93+
GEMV_STORE(y + i + 2 * GEMV_VLEN, acc);
94+
95+
acc = GEMV_LOAD(y + i + 3 * GEMV_VLEN);
96+
acc = GEMV_MADD(acc, v0, GEMV_LOAD(a0 + i + 3 * GEMV_VLEN));
97+
acc = GEMV_MADD(acc, v1, GEMV_LOAD(a1 + i + 3 * GEMV_VLEN));
98+
acc = GEMV_MADD(acc, v2, GEMV_LOAD(a2 + i + 3 * GEMV_VLEN));
99+
acc = GEMV_MADD(acc, v3, GEMV_LOAD(a3 + i + 3 * GEMV_VLEN));
100+
GEMV_STORE(y + i + 3 * GEMV_VLEN, acc);
101+
}
102+
for (; i + GEMV_VLEN <= m; i += GEMV_VLEN) {
103+
v128_t acc = GEMV_LOAD(y + i);
104+
acc = GEMV_MADD(acc, v0, GEMV_LOAD(a0 + i));
105+
acc = GEMV_MADD(acc, v1, GEMV_LOAD(a1 + i));
106+
acc = GEMV_MADD(acc, v2, GEMV_LOAD(a2 + i));
107+
acc = GEMV_MADD(acc, v3, GEMV_LOAD(a3 + i));
108+
GEMV_STORE(y + i, acc);
109+
}
110+
for (; i < m; i++)
111+
y[i] += t0 * a0[i] + t1 * a1[i] + t2 * a2[i] + t3 * a3[i];
112+
}
113+
114+
static void gemv_n_axpy1(BLASLONG m, const FLOAT *a, FLOAT *y, FLOAT t) {
115+
const v128_t vt = GEMV_SPLAT(t);
116+
BLASLONG i = 0;
117+
const BLASLONG n_main = m & ~(BLASLONG)(GEMV_CHUNK - 1);
118+
119+
for (; i < n_main; i += GEMV_CHUNK) {
120+
v128_t acc;
121+
acc = GEMV_LOAD(y + i + 0 * GEMV_VLEN);
122+
GEMV_STORE(y + i + 0 * GEMV_VLEN, GEMV_MADD(acc, vt, GEMV_LOAD(a + i + 0 * GEMV_VLEN)));
123+
acc = GEMV_LOAD(y + i + 1 * GEMV_VLEN);
124+
GEMV_STORE(y + i + 1 * GEMV_VLEN, GEMV_MADD(acc, vt, GEMV_LOAD(a + i + 1 * GEMV_VLEN)));
125+
acc = GEMV_LOAD(y + i + 2 * GEMV_VLEN);
126+
GEMV_STORE(y + i + 2 * GEMV_VLEN, GEMV_MADD(acc, vt, GEMV_LOAD(a + i + 2 * GEMV_VLEN)));
127+
acc = GEMV_LOAD(y + i + 3 * GEMV_VLEN);
128+
GEMV_STORE(y + i + 3 * GEMV_VLEN, GEMV_MADD(acc, vt, GEMV_LOAD(a + i + 3 * GEMV_VLEN)));
129+
}
130+
for (; i + GEMV_VLEN <= m; i += GEMV_VLEN) {
131+
v128_t acc = GEMV_LOAD(y + i);
132+
GEMV_STORE(y + i, GEMV_MADD(acc, vt, GEMV_LOAD(a + i)));
133+
}
134+
for (; i < m; i++)
135+
y[i] += t * a[i];
136+
}
137+
138+
#undef GEMV_MADD
139+
#undef GEMV_STORE
140+
#undef GEMV_LOAD
141+
#undef GEMV_CHUNK
142+
#undef GEMV_UNROLL
143+
#undef GEMV_ADD
144+
#undef GEMV_MUL
145+
#undef GEMV_SPLAT
146+
#undef GEMV_VLEN
147+
#endif
148+
149+
int CNAME(BLASLONG m, BLASLONG n, BLASLONG dummy1, FLOAT alpha, FLOAT *a,
150+
BLASLONG lda, FLOAT *x, BLASLONG inc_x, FLOAT *y, BLASLONG inc_y,
151+
FLOAT *buffer) {
152+
BLASLONG i, j, ix, iy;
153+
FLOAT *a_ptr;
154+
FLOAT temp;
155+
156+
(void)dummy1;
157+
(void)buffer;
158+
159+
if (m < 1 || n < 1)
160+
return 0;
161+
if (alpha == (FLOAT)0.0)
162+
return 0;
163+
164+
#if defined(__wasm_simd128__)
165+
if (inc_y == 1) {
166+
ix = 0;
167+
j = 0;
168+
for (; j + 4 <= n; j += 4) {
169+
gemv_n_axpy4(m, a + j * lda, a + (j + 1) * lda, a + (j + 2) * lda,
170+
a + (j + 3) * lda, y, alpha * x[ix],
171+
alpha * x[ix + inc_x], alpha * x[ix + 2 * inc_x],
172+
alpha * x[ix + 3 * inc_x]);
173+
ix += 4 * inc_x;
174+
}
175+
for (; j < n; j++) {
176+
gemv_n_axpy1(m, a + j * lda, y, alpha * x[ix]);
177+
ix += inc_x;
178+
}
179+
return 0;
180+
}
181+
#endif
182+
183+
ix = 0;
184+
a_ptr = a;
185+
for (j = 0; j < n; j++) {
186+
temp = alpha * x[ix];
187+
iy = 0;
188+
for (i = 0; i < m; i++) {
189+
y[iy] += temp * a_ptr[i];
190+
iy += inc_y;
191+
}
192+
a_ptr += lda;
193+
ix += inc_x;
194+
}
195+
return 0;
196+
}

0 commit comments

Comments
 (0)