1
0
mirror of https://github.com/opencv/opencv.git synced 2026-07-29 07:13:02 +04:00

dnn: keep fastGemmThin accumulators in registers on scalable-vector targets (RVV)

This commit is contained in:
Teddy-Yangjiale
2026-07-22 22:42:25 +08:00
parent c7b8fb28b6
commit 7cc9947809
2 changed files with 55 additions and 29 deletions
+7
View File
@@ -83,6 +83,13 @@ static const GemmParam_t test_matmul_configs[] = {
{ {16, 197, 64 }, {16, 64, 197} },
{ {16, 50, 64}, {16, 64, 50} },
{ {16, 50, 50}, {16, 50, 64} },
// transformer token generation cases (thin-M GEMV, fastGemmThin path)
{ {1, 1, 768}, {1, 768, 768} },
{ {1, 1, 768}, {1, 768, 3072} },
{ {1, 1, 3072}, {1, 3072, 768} },
{ {1, 2, 768}, {1, 768, 768} },
{ {1, 4, 768}, {1, 768, 768} },
};
struct GemmParamId
@@ -276,6 +276,45 @@ static inline int fast_gemm_thin_lanes() {
#endif
}
#if CV_SIMD_SCALABLE
// Apply the alpha/beta epilogue to one accumulator row.
static inline void fast_gemm_thin_store_row(float* c, v_float32 acc,
v_float32 v_alpha, float beta) {
if (beta == 0.f)
vx_store(c, v_mul(acc, v_alpha));
else if (beta == 1.f)
vx_store(c, v_fma(acc, v_alpha, vx_load(c)));
else
vx_store(c, v_fma(acc, v_alpha, v_mul(vx_load(c), vx_setall_f32(beta))));
}
// Multiply MR (1..4) rows of A by one packed-B strip. Sizeless vector types
// cannot form the acc[] array the CV_SIMD path uses, so the accumulators are
// named variables selected by the compile-time MR (the `if (MR > n)` branches
// fold away) and stay in registers for the whole K loop. Accumulation remains
// in increasing-k order, bit-exact with the previous scratch-buffer code.
template<int MR>
static inline void fast_gemm_thin_block(int K, const float* A, int lda0, int lda1,
const float* b_strip, int NR,
v_float32 v_alpha, float beta,
float* c_strip, int ldc) {
v_float32 c0 = vx_setzero_f32(), c1 = vx_setzero_f32();
v_float32 c2 = vx_setzero_f32(), c3 = vx_setzero_f32();
for (int k = 0; k < K; k++) {
v_float32 bv = vx_load(b_strip + k * NR);
const float* a = A + k * lda1;
c0 = v_fma(bv, vx_setall_f32(a[0]), c0);
if (MR > 1) c1 = v_fma(bv, vx_setall_f32(a[lda0]), c1);
if (MR > 2) c2 = v_fma(bv, vx_setall_f32(a[2 * lda0]), c2);
if (MR > 3) c3 = v_fma(bv, vx_setall_f32(a[3 * lda0]), c3);
}
fast_gemm_thin_store_row(c_strip, c0, v_alpha, beta);
if (MR > 1) fast_gemm_thin_store_row(c_strip + ldc, c1, v_alpha, beta);
if (MR > 2) fast_gemm_thin_store_row(c_strip + 2 * ldc, c2, v_alpha, beta);
if (MR > 3) fast_gemm_thin_store_row(c_strip + 3 * ldc, c3, v_alpha, beta);
}
#endif // CV_SIMD_SCALABLE
static inline void fast_gemm_thin_strip(int M, int K, float alpha,
const float* A, int lda0, int lda1,
const float* b_strip,
@@ -311,37 +350,17 @@ static inline void fast_gemm_thin_strip(int M, int K, float alpha,
}
}
#elif CV_SIMD_SCALABLE
// Scalable vector types (e.g. RVV) are sizeless and cannot form arrays;
// back the per-row accumulators with a scalar scratch buffer.
// Process the strip in register-resident row blocks; the common thin case
// M <= 4 (transformer token generation) is a single pass.
const int NR = VTraits<v_float32>::vlanes();
float acc_buf[FAST_GEMM_THIN_MAX_M * VTraits<v_float32>::max_nlanes];
for (int m = 0; m < M; m++) vx_store(acc_buf + m * NR, vx_setzero_f32());
for (int k = 0; k < K; k++) {
v_float32 bv = vx_load(b_strip + k * NR);
for (int m = 0; m < M; m++) {
v_float32 am = vx_setall_f32(A[m * lda0 + k * lda1]);
v_float32 acc_m = vx_load(acc_buf + m * NR);
vx_store(acc_buf + m * NR, v_fma(bv, am, acc_m));
}
}
const v_float32 v_alpha = vx_setall_f32(alpha);
if (beta == 0.f) {
for (int m = 0; m < M; m++)
vx_store(c_strip + m * ldc, v_mul(vx_load(acc_buf + m * NR), v_alpha));
} else if (beta == 1.f) {
for (int m = 0; m < M; m++) {
v_float32 cv = vx_load(c_strip + m * ldc);
vx_store(c_strip + m * ldc, v_fma(vx_load(acc_buf + m * NR), v_alpha, cv));
}
} else {
const v_float32 v_beta = vx_setall_f32(beta);
for (int m = 0; m < M; m++) {
v_float32 cv = vx_load(c_strip + m * ldc);
cv = v_mul(cv, v_beta);
vx_store(c_strip + m * ldc, v_fma(vx_load(acc_buf + m * NR), v_alpha, cv));
}
int m = 0;
for (; m + 4 <= M; m += 4)
fast_gemm_thin_block<4>(K, A + m * lda0, lda0, lda1, b_strip, NR, v_alpha, beta, c_strip + m * ldc, ldc);
switch (M - m) {
case 1: fast_gemm_thin_block<1>(K, A + m * lda0, lda0, lda1, b_strip, NR, v_alpha, beta, c_strip + m * ldc, ldc); break;
case 2: fast_gemm_thin_block<2>(K, A + m * lda0, lda0, lda1, b_strip, NR, v_alpha, beta, c_strip + m * ldc, ldc); break;
case 3: fast_gemm_thin_block<3>(K, A + m * lda0, lda0, lda1, b_strip, NR, v_alpha, beta, c_strip + m * ldc, ldc); break;
}
#else
const int NR = FAST_GEMM_THIN_SCALAR_NR;