Skip to content

Commit 5282a38

Browse files
committed
kernel/riscv64: add RVV LN TRSM kernel
1 parent 34f66e5 commit 5282a38

1 file changed

Lines changed: 333 additions & 0 deletions

File tree

Lines changed: 333 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,333 @@
1+
/***************************************************************************
2+
Copyright (c) 2022, 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 CONSEQUENTIAL
21+
DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
22+
SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
23+
CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
24+
OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE
25+
USE OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
26+
*****************************************************************************/
27+
28+
#include "common.h"
29+
30+
#if !defined(DOUBLE)
31+
#define VSETVL(n) __riscv_vsetvl_e32m2(n)
32+
#define FLOAT_V_T vfloat32m2_t
33+
#define FLOAT_VX2_T vfloat32m2x2_t
34+
#define VGET_VX2 __riscv_vget_v_f32m2x2_f32m2
35+
#define VSET_VX2 __riscv_vset_v_f32m2_f32m2x2
36+
#define VLEV_FLOAT __riscv_vle32_v_f32m2
37+
#define VSEV_FLOAT __riscv_vse32_v_f32m2
38+
#define VSSEG2_FLOAT __riscv_vsseg2e32_v_f32m2x2
39+
#define VLSEG2_FLOAT __riscv_vlseg2e32_v_f32m2x2
40+
#define VFMACCVF_FLOAT __riscv_vfmacc_vf_f32m2
41+
#define VFNMSACVF_FLOAT __riscv_vfnmsac_vf_f32m2
42+
#define VFMULVF_FLOAT __riscv_vfmul_vf_f32m2
43+
#else
44+
#define VSETVL(n) __riscv_vsetvl_e64m2(n)
45+
#define FLOAT_V_T vfloat64m2_t
46+
#define FLOAT_VX2_T vfloat64m2x2_t
47+
#define VGET_VX2 __riscv_vget_v_f64m2x2_f64m2
48+
#define VSET_VX2 __riscv_vset_v_f64m2_f64m2x2
49+
#define VLEV_FLOAT __riscv_vle64_v_f64m2
50+
#define VSEV_FLOAT __riscv_vse64_v_f64m2
51+
#define VSSEG2_FLOAT __riscv_vsseg2e64_v_f64m2x2
52+
#define VLSEG2_FLOAT __riscv_vlseg2e64_v_f64m2x2
53+
#define VFMVVF_FLOAT __riscv_vfmv_v_f_f64m2
54+
#define VFMACCVF_FLOAT __riscv_vfmacc_vf_f64m2
55+
#define VFNMSACVF_FLOAT __riscv_vfnmsac_vf_f64m2
56+
#define VFMULVF_FLOAT __riscv_vfmul_vf_f64m2
57+
#endif
58+
59+
60+
static FLOAT dm1 = -1.;
61+
62+
#ifdef CONJ
63+
#define GEMM_KERNEL GEMM_KERNEL_L
64+
#else
65+
#define GEMM_KERNEL GEMM_KERNEL_N
66+
#endif
67+
68+
#if GEMM_DEFAULT_UNROLL_N == 1
69+
#define GEMM_UNROLL_N_SHIFT 0
70+
#endif
71+
72+
#if GEMM_DEFAULT_UNROLL_N == 2
73+
#define GEMM_UNROLL_N_SHIFT 1
74+
#endif
75+
76+
#if GEMM_DEFAULT_UNROLL_N == 4
77+
#define GEMM_UNROLL_N_SHIFT 2
78+
#endif
79+
80+
#if GEMM_DEFAULT_UNROLL_N == 8
81+
#define GEMM_UNROLL_N_SHIFT 3
82+
#endif
83+
84+
#if GEMM_DEFAULT_UNROLL_N == 16
85+
#define GEMM_UNROLL_N_SHIFT 4
86+
#endif
87+
88+
// Optimizes the implementation in ../arm64/trsm_kernel_LN_sve.c
89+
90+
#ifndef COMPLEX
91+
92+
static inline void solve(BLASLONG m, BLASLONG n, FLOAT *a, FLOAT *b, FLOAT *c, BLASLONG ldc) {
93+
FLOAT aa;
94+
FLOAT* pc;
95+
96+
int i, j, k;
97+
98+
FLOAT_V_T va, vc;
99+
100+
size_t vl;
101+
102+
a += (m - 1) * m;
103+
b += (m - 1) * n;
104+
105+
for (i = m - 1; i >= 0; i--) {
106+
107+
aa = *(a + i);
108+
for (j = 0; j < n; j++) {
109+
FLOAT bb;
110+
111+
pc = c + j * ldc;
112+
bb = *(pc + i) * aa;
113+
*(b + j) = bb;
114+
*(pc + i) = bb;
115+
}
116+
117+
for (k = 0; k < i; k += vl) {
118+
vl = VSETVL(i - k);
119+
va = VLEV_FLOAT(a + k, vl);
120+
for (j = 0; j < n; j++) {
121+
pc = c + j * ldc;
122+
vc = VLEV_FLOAT(pc + k, vl);
123+
vc = VFNMSACVF_FLOAT(vc, *(b + j), va, vl);
124+
VSEV_FLOAT(pc + k, vc, vl);
125+
}
126+
}
127+
128+
a -= m;
129+
b += n;
130+
b -= 2 * n;
131+
}
132+
133+
}
134+
#else
135+
136+
static inline void solve(BLASLONG m, BLASLONG n, FLOAT *a, FLOAT *b, FLOAT *c, BLASLONG ldc) {
137+
138+
FLOAT aa1, aa2;
139+
FLOAT *pc;
140+
int i, j, k;
141+
142+
FLOAT_VX2_T vax2, vcx2;
143+
FLOAT_V_T va1, va2, vc1, vc2;
144+
size_t vl;
145+
BLASLONG ldc2 = ldc * 2;
146+
147+
a += (m - 1) * m * 2;
148+
b += (m - 1) * n * 2;
149+
150+
for (i = m - 1; i >= 0; i--) {
151+
152+
aa1 = *(a + i * 2 + 0);
153+
aa2 = *(a + i * 2 + 1);
154+
for (j = 0; j < n; j++) {
155+
FLOAT bb1, bb2, ss1, ss2;
156+
157+
pc = c + j * ldc2;
158+
bb1 = *(pc + i * 2 + 0);
159+
bb2 = *(pc + i * 2 + 1);
160+
#ifndef CONJ
161+
ss1 = aa1 * bb1 - aa2 * bb2;
162+
ss2 = aa1 * bb2 + aa2 * bb1;
163+
#else
164+
ss1 = aa1 * bb1 + aa2 * bb2;
165+
ss2 = aa1 * bb2 - aa2 * bb1;
166+
#endif
167+
*(b + j * 2 + 0) = ss1;
168+
*(b + j * 2 + 1) = ss2;
169+
*(pc + i * 2 + 0) = ss1;
170+
*(pc + i * 2 + 1) = ss2;
171+
}
172+
173+
for (k = 0; k < i; k += vl) {
174+
vl = VSETVL(i - k);
175+
vax2 = VLSEG2_FLOAT(a + k * 2, vl);
176+
va1 = VGET_VX2(vax2, 0);
177+
va2 = VGET_VX2(vax2, 1);
178+
for (j = 0; j < n; j++) {
179+
FLOAT ss1 = *(b + j * 2 + 0);
180+
FLOAT ss2 = *(b + j * 2 + 1);
181+
182+
pc = c + j * ldc2;
183+
vcx2 = VLSEG2_FLOAT(pc + k * 2, vl);
184+
vc1 = VGET_VX2(vcx2, 0);
185+
vc2 = VGET_VX2(vcx2, 1);
186+
#ifndef CONJ
187+
vc1 = VFMACCVF_FLOAT(vc1, ss2, va2, vl);
188+
vc1 = VFNMSACVF_FLOAT(vc1, ss1, va1, vl);
189+
vc2 = VFNMSACVF_FLOAT(vc2, ss1, va2, vl);
190+
vc2 = VFNMSACVF_FLOAT(vc2, ss2, va1, vl);
191+
#else
192+
vc1 = VFNMSACVF_FLOAT(vc1, ss2, va2, vl);
193+
vc1 = VFNMSACVF_FLOAT(vc1, ss1, va1, vl);
194+
vc2 = VFMACCVF_FLOAT(vc2, ss1, va2, vl);
195+
vc2 = VFNMSACVF_FLOAT(vc2, ss2, va1, vl);
196+
#endif
197+
vcx2 = VSET_VX2(vcx2, 0, vc1);
198+
vcx2 = VSET_VX2(vcx2, 1, vc2);
199+
VSSEG2_FLOAT(pc + k * 2, vcx2, vl);
200+
}
201+
}
202+
203+
a -= m * 2;
204+
b += n * 2;
205+
b -= 4 * n;
206+
}
207+
}
208+
209+
210+
#endif
211+
212+
int CNAME(BLASLONG m, BLASLONG n, BLASLONG k, FLOAT dummy1,
213+
#ifdef COMPLEX
214+
FLOAT dummy2,
215+
#endif
216+
FLOAT *a, FLOAT *b, FLOAT *c, BLASLONG ldc, BLASLONG offset){
217+
218+
BLASLONG i, j;
219+
FLOAT *aa, *cc;
220+
BLASLONG kk;
221+
222+
#ifndef COMPLEX
223+
#define PROCESS_LN_M_BLOCK(MB, NB) do { \
224+
if (k - kk > 0) { \
225+
GEMM_KERNEL((MB), (NB), k - kk, dm1, \
226+
aa + (MB) * kk * COMPSIZE, \
227+
b + (NB) * kk * COMPSIZE, \
228+
cc, ldc); \
229+
} \
230+
solve((MB), (NB), \
231+
aa + (kk - (MB)) * (MB) * COMPSIZE, \
232+
b + (kk - (MB)) * (NB) * COMPSIZE, \
233+
cc, ldc); \
234+
} while (0)
235+
#else
236+
#define PROCESS_LN_M_BLOCK(MB, NB) do { \
237+
if (k - kk > 0) { \
238+
GEMM_KERNEL((MB), (NB), k - kk, dm1, ZERO, \
239+
aa + (MB) * kk * COMPSIZE, \
240+
b + (NB) * kk * COMPSIZE, \
241+
cc, ldc); \
242+
} \
243+
solve((MB), (NB), \
244+
aa + (kk - (MB)) * (MB) * COMPSIZE, \
245+
b + (kk - (MB)) * (NB) * COMPSIZE, \
246+
cc, ldc); \
247+
} while (0)
248+
#endif
249+
250+
j = (n >> GEMM_UNROLL_N_SHIFT);
251+
252+
while (j > 0) {
253+
254+
kk = m + offset;
255+
256+
if (m & (GEMM_DEFAULT_UNROLL_M - 1)) {
257+
for (i = 1; i < GEMM_DEFAULT_UNROLL_M; i <<= 1) {
258+
if (m & i) {
259+
aa = a + ((m & ~(i - 1)) - i) * k * COMPSIZE;
260+
cc = c + ((m & ~(i - 1)) - i) * COMPSIZE;
261+
262+
PROCESS_LN_M_BLOCK(i, GEMM_UNROLL_N);
263+
kk -= i;
264+
}
265+
}
266+
}
267+
268+
i = (m & ~(GEMM_DEFAULT_UNROLL_M - 1));
269+
if (i > 0) {
270+
aa = a + (i - GEMM_DEFAULT_UNROLL_M) * k * COMPSIZE;
271+
cc = c + (i - GEMM_DEFAULT_UNROLL_M) * COMPSIZE;
272+
273+
do {
274+
PROCESS_LN_M_BLOCK(GEMM_DEFAULT_UNROLL_M, GEMM_UNROLL_N);
275+
276+
aa -= GEMM_DEFAULT_UNROLL_M * k * COMPSIZE;
277+
cc -= GEMM_DEFAULT_UNROLL_M * COMPSIZE;
278+
kk -= GEMM_DEFAULT_UNROLL_M;
279+
i -= GEMM_DEFAULT_UNROLL_M;
280+
} while (i > 0);
281+
}
282+
283+
b += GEMM_UNROLL_N * k * COMPSIZE;
284+
c += GEMM_UNROLL_N * ldc * COMPSIZE;
285+
j --;
286+
}
287+
288+
if (n & (GEMM_UNROLL_N - 1)) {
289+
290+
j = (GEMM_UNROLL_N >> 1);
291+
while (j > 0) {
292+
if (n & j) {
293+
294+
kk = m + offset;
295+
296+
if (m & (GEMM_DEFAULT_UNROLL_M - 1)) {
297+
for (i = 1; i < GEMM_DEFAULT_UNROLL_M; i <<= 1) {
298+
if (m & i) {
299+
aa = a + ((m & ~(i - 1)) - i) * k * COMPSIZE;
300+
cc = c + ((m & ~(i - 1)) - i) * COMPSIZE;
301+
302+
PROCESS_LN_M_BLOCK(i, j);
303+
kk -= i;
304+
}
305+
}
306+
}
307+
308+
i = (m & ~(GEMM_DEFAULT_UNROLL_M - 1));
309+
if (i > 0) {
310+
aa = a + (i - GEMM_DEFAULT_UNROLL_M) * k * COMPSIZE;
311+
cc = c + (i - GEMM_DEFAULT_UNROLL_M) * COMPSIZE;
312+
313+
do {
314+
PROCESS_LN_M_BLOCK(GEMM_DEFAULT_UNROLL_M, j);
315+
316+
aa -= GEMM_DEFAULT_UNROLL_M * k * COMPSIZE;
317+
cc -= GEMM_DEFAULT_UNROLL_M * COMPSIZE;
318+
kk -= GEMM_DEFAULT_UNROLL_M;
319+
i -= GEMM_DEFAULT_UNROLL_M;
320+
} while (i > 0);
321+
}
322+
323+
b += j * k * COMPSIZE;
324+
c += j * ldc * COMPSIZE;
325+
}
326+
j >>= 1;
327+
}
328+
}
329+
330+
return 0;
331+
332+
#undef PROCESS_LN_M_BLOCK
333+
}

0 commit comments

Comments
 (0)