@@ -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