Skip to content

Commit 0569c7b

Browse files
implemented simd_gemm kernel for riscv vector extension
1 parent acc37a4 commit 0569c7b

1 file changed

Lines changed: 90 additions & 0 deletions

File tree

ggml/src/ggml-cpu/simd-gemm.h

Lines changed: 90 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -109,6 +109,96 @@ static void simd_gemm(
109109
C += N;
110110
}
111111
}
112+
#elif defined(GGML_SIMD) && defined(__riscv_v_intrinsic)
113+
// RM accumulators + 1 B vector = RM + 1 <= 8 => RM <= 7
114+
// Microkernel: C[RM x vl] += A[RM x K] * B[K x N]
115+
template <int RM>
116+
static inline void rvv_simd_gemm_ukernel(
117+
float * GGML_RESTRICT C,
118+
const float * GGML_RESTRICT A,
119+
const float * GGML_RESTRICT B,
120+
int K, int N, size_t vl)
121+
{
122+
static_assert(RM >= 1 && RM <= 7, "RM must be 1..7 for LMUL=4");
123+
124+
vfloat32m4_t acc_0 = __riscv_vle32_v_f32m4(C + 0 * N, vl);
125+
vfloat32m4_t acc_1, acc_2, acc_3, acc_4, acc_5, acc_6;
126+
if constexpr (RM > 1) acc_1 = __riscv_vle32_v_f32m4(C + 1 * N, vl);
127+
if constexpr (RM > 2) acc_2 = __riscv_vle32_v_f32m4(C + 2 * N, vl);
128+
if constexpr (RM > 3) acc_3 = __riscv_vle32_v_f32m4(C + 3 * N, vl);
129+
if constexpr (RM > 4) acc_4 = __riscv_vle32_v_f32m4(C + 4 * N, vl);
130+
if constexpr (RM > 5) acc_5 = __riscv_vle32_v_f32m4(C + 5 * N, vl);
131+
if constexpr (RM > 6) acc_6 = __riscv_vle32_v_f32m4(C + 6 * N, vl);
132+
133+
for (int kk = 0; kk < K; kk++) {
134+
vfloat32m4_t b_0 = __riscv_vle32_v_f32m4(B + kk * N, vl);
135+
136+
acc_0 = __riscv_vfmacc_vf_f32m4(acc_0, A[0 * K + kk], b_0, vl);
137+
if constexpr (RM > 1) acc_1 = __riscv_vfmacc_vf_f32m4(acc_1, A[1 * K + kk], b_0, vl);
138+
if constexpr (RM > 2) acc_2 = __riscv_vfmacc_vf_f32m4(acc_2, A[2 * K + kk], b_0, vl);
139+
if constexpr (RM > 3) acc_3 = __riscv_vfmacc_vf_f32m4(acc_3, A[3 * K + kk], b_0, vl);
140+
if constexpr (RM > 4) acc_4 = __riscv_vfmacc_vf_f32m4(acc_4, A[4 * K + kk], b_0, vl);
141+
if constexpr (RM > 5) acc_5 = __riscv_vfmacc_vf_f32m4(acc_5, A[5 * K + kk], b_0, vl);
142+
if constexpr (RM > 6) acc_6 = __riscv_vfmacc_vf_f32m4(acc_6, A[6 * K + kk], b_0, vl);
143+
}
144+
145+
__riscv_vse32_v_f32m4(C + 0 * N, acc_0, vl);
146+
if constexpr (RM > 1) __riscv_vse32_v_f32m4(C + 1 * N, acc_1, vl);
147+
if constexpr (RM > 2) __riscv_vse32_v_f32m4(C + 2 * N, acc_2, vl);
148+
if constexpr (RM > 3) __riscv_vse32_v_f32m4(C + 3 * N, acc_3, vl);
149+
if constexpr (RM > 4) __riscv_vse32_v_f32m4(C + 4 * N, acc_4, vl);
150+
if constexpr (RM > 5) __riscv_vse32_v_f32m4(C + 5 * N, acc_5, vl);
151+
if constexpr (RM > 6) __riscv_vse32_v_f32m4(C + 6 * N, acc_6, vl);
152+
}
153+
154+
template <int RM>
155+
static inline void rvv_simd_gemm_dispatch_tail(
156+
float * GGML_RESTRICT C,
157+
const float * GGML_RESTRICT A,
158+
const float * GGML_RESTRICT B,
159+
int K, int N, int KN, int remaining_rows)
160+
{
161+
if constexpr (RM > 0) {
162+
if (remaining_rows == RM) {
163+
int64_t jj = 0;
164+
for (; jj + KN <= N; jj += KN) {
165+
rvv_simd_gemm_ukernel<RM>(C + jj, A, B + jj, K, N, KN);
166+
}
167+
if (jj < N) {
168+
rvv_simd_gemm_ukernel<RM>(C + jj, A, B + jj, K, N, N - jj);
169+
}
170+
} else {
171+
rvv_simd_gemm_dispatch_tail<RM - 1>(C, A, B, K, N, KN, remaining_rows);
172+
}
173+
}
174+
}
175+
176+
static constexpr int GEMM_RM = 7;
177+
178+
// C[M x N] += A[M x K] * B[K x N]
179+
static void simd_gemm(
180+
float * GGML_RESTRICT C,
181+
const float * GGML_RESTRICT A,
182+
const float * GGML_RESTRICT B,
183+
int M, int K, int N)
184+
{
185+
const int KN = (int)__riscv_vlenb();
186+
int64_t ii = 0;
187+
for (; ii + GEMM_RM <= M; ii += GEMM_RM) {
188+
int64_t jj = 0;
189+
for (; jj + KN <= N; jj += KN) {
190+
rvv_simd_gemm_ukernel<GEMM_RM>(C + jj, A, B + jj, K, N, KN);
191+
}
192+
if (jj < N) {
193+
rvv_simd_gemm_ukernel<GEMM_RM>(C + jj, A, B + jj, K, N, N - jj);
194+
}
195+
A += GEMM_RM * K;
196+
C += GEMM_RM * N;
197+
}
198+
199+
int remaining_rows = M - ii;
200+
rvv_simd_gemm_dispatch_tail<GEMM_RM - 1>(C, A, B, K, N, KN, remaining_rows);
201+
}
112202

113203
#if defined(__GNUC__) && !defined(__clang__)
114204
#pragma GCC diagnostic pop

0 commit comments

Comments
 (0)