From b97485577f3585d2bfbfb60da004a05b8d420e8e Mon Sep 17 00:00:00 2001 From: Arthur031221 Date: Fri, 2 Oct 2026 04:35:35 +0800 Subject: [PATCH 1/2] Fix complex scalars in strided batched GEMM --- CONTRIBUTORS.md | 3 +++ interface/gemm_batch_strided.c | 5 +++++ utest/test_extensions/test_cgemm.c | 18 ++++++++++++++++++ utest/test_extensions/test_zgemm.c | 18 ++++++++++++++++++ 4 files changed, 44 insertions(+) diff --git a/CONTRIBUTORS.md b/CONTRIBUTORS.md index 90601ccdb8..5553f9a877 100644 --- a/CONTRIBUTORS.md +++ b/CONTRIBUTORS.md @@ -289,3 +289,6 @@ hheei * Hugo Meiland * [2026-08-09] Add Cortex-A72 DGEMM 6x8 microkernel and blocking + +* Arthur031221 + * [2026-10-02] Fix complex scaling in strided batched GEMM. diff --git a/interface/gemm_batch_strided.c b/interface/gemm_batch_strided.c index 026eeb50e2..e9529b4a76 100644 --- a/interface/gemm_batch_strided.c +++ b/interface/gemm_batch_strided.c @@ -396,8 +396,13 @@ void CNAME(enum CBLAS_ORDER order, enum CBLAS_TRANSPOSE transa, enum CBLAS_TRANS args_array[i].lda=group_lda; args_array[i].ldb=group_ldb; args_array[i].ldc=group_ldc; +#ifdef COMPLEX + args_array[i].alpha=alpha; + args_array[i].beta=beta; +#else args_array[i].alpha=α args_array[i].beta=β +#endif #if defined(CBLAS) if (order == CblasColMajor) { diff --git a/utest/test_extensions/test_cgemm.c b/utest/test_extensions/test_cgemm.c index 15d64e372e..ba5223f0be 100644 --- a/utest/test_extensions/test_cgemm.c +++ b/utest/test_extensions/test_cgemm.c @@ -271,4 +271,22 @@ CTEST(cgemm, transa_conjnotransb) ASSERT_DBL_NEAR_TOL(0.0f, norm, SINGLE_EPS); } + +#ifndef NO_CBLAS +CTEST(cgemm, strided_batch_complex_scalars) +{ + float a[] = {1.0f, 2.0f}; + float b[] = {3.0f, 4.0f}; + float c[] = {5.0f, 6.0f}; + float alpha[] = {2.0f, 3.0f}; + float beta[] = {4.0f, -1.0f}; + + cblas_cgemm_batch_strided(CblasColMajor, CblasNoTrans, CblasNoTrans, + 1, 1, 1, alpha, a, 1, 2, b, 1, 2, + beta, c, 1, 2, 1); + + ASSERT_DBL_NEAR_TOL(-14.0, c[0], SINGLE_EPS); + ASSERT_DBL_NEAR_TOL(24.0, c[1], SINGLE_EPS); +} +#endif #endif diff --git a/utest/test_extensions/test_zgemm.c b/utest/test_extensions/test_zgemm.c index bd23ebca42..241088f1ad 100644 --- a/utest/test_extensions/test_zgemm.c +++ b/utest/test_extensions/test_zgemm.c @@ -271,4 +271,22 @@ CTEST(zgemm, transa_conjnotransb) ASSERT_DBL_NEAR_TOL(0.0, norm, DOUBLE_EPS); } + +#ifndef NO_CBLAS +CTEST(zgemm, strided_batch_complex_scalars) +{ + double a[] = {1.0, 2.0}; + double b[] = {3.0, 4.0}; + double c[] = {5.0, 6.0}; + double alpha[] = {2.0, 3.0}; + double beta[] = {4.0, -1.0}; + + cblas_zgemm_batch_strided(CblasColMajor, CblasNoTrans, CblasNoTrans, + 1, 1, 1, alpha, a, 1, 2, b, 1, 2, + beta, c, 1, 2, 1); + + ASSERT_DBL_NEAR_TOL(-14.0, c[0], DOUBLE_EPS); + ASSERT_DBL_NEAR_TOL(24.0, c[1], DOUBLE_EPS); +} +#endif #endif From e862d775c8012e45c44de2e174d8af0eb0a511b4 Mon Sep 17 00:00:00 2001 From: Arthur031221 Date: Sun, 4 Oct 2026 02:31:20 +0800 Subject: [PATCH 2/2] Resolve small-kernel addresses in strided batched GEMM --- interface/gemm_batch_strided.c | 8 ++++---- utest/test_extensions/test_cgemm.c | 10 ++++++++++ utest/test_extensions/test_zgemm.c | 10 ++++++++++ 3 files changed, 24 insertions(+), 4 deletions(-) diff --git a/interface/gemm_batch_strided.c b/interface/gemm_batch_strided.c index e9529b4a76..f711b77b6c 100644 --- a/interface/gemm_batch_strided.c +++ b/interface/gemm_batch_strided.c @@ -366,18 +366,18 @@ void CNAME(enum CBLAS_ORDER order, enum CBLAS_TRANSPOSE transa, enum CBLAS_TRANS #if !defined(COMPLEX) if(beta == 0.0){ group_mode=mode | BLAS_SMALL_B0_OPT; - group_small_matrix_opt_routine=(void *)(gemm_small_kernel_b0[(group_transb<<2)|group_transa]); + group_small_matrix_opt_routine=SMALL_KERNEL_ADDR(gemm_small_kernel_b0, ((group_transb<<2)|group_transa)); }else{ group_mode=mode | BLAS_SMALL_OPT; - group_small_matrix_opt_routine=(void *)(gemm_small_kernel[(group_transb<<2)|group_transa]); + group_small_matrix_opt_routine=SMALL_KERNEL_ADDR(gemm_small_kernel, ((group_transb<<2)|group_transa)); } #else if(beta[0] == 0.0 && beta[1] == 0.0){ group_mode=mode | BLAS_SMALL_B0_OPT; - group_small_matrix_opt_routine=(void *)(zgemm_small_kernel_b0[(group_transb<<2)|group_transa]); + group_small_matrix_opt_routine=SMALL_KERNEL_ADDR(zgemm_small_kernel_b0, ((group_transb<<2)|group_transa)); }else{ group_mode=mode | BLAS_SMALL_OPT; - group_small_matrix_opt_routine=(void *)(zgemm_small_kernel[(group_transb<<2)|group_transa]); + group_small_matrix_opt_routine=SMALL_KERNEL_ADDR(zgemm_small_kernel, ((group_transb<<2)|group_transa)); } #endif diff --git a/utest/test_extensions/test_cgemm.c b/utest/test_extensions/test_cgemm.c index ba5223f0be..055203f5a1 100644 --- a/utest/test_extensions/test_cgemm.c +++ b/utest/test_extensions/test_cgemm.c @@ -287,6 +287,16 @@ CTEST(cgemm, strided_batch_complex_scalars) ASSERT_DBL_NEAR_TOL(-14.0, c[0], SINGLE_EPS); ASSERT_DBL_NEAR_TOL(24.0, c[1], SINGLE_EPS); + + beta[0] = beta[1] = 0.0; + c[0] = 5.0; + c[1] = 6.0; + cblas_cgemm_batch_strided(CblasColMajor, CblasNoTrans, CblasNoTrans, + 1, 1, 1, alpha, a, 1, 2, b, 1, 2, + beta, c, 1, 2, 1); + + ASSERT_DBL_NEAR_TOL(-40.0, c[0], SINGLE_EPS); + ASSERT_DBL_NEAR_TOL(5.0, c[1], SINGLE_EPS); } #endif #endif diff --git a/utest/test_extensions/test_zgemm.c b/utest/test_extensions/test_zgemm.c index 241088f1ad..1238a5a437 100644 --- a/utest/test_extensions/test_zgemm.c +++ b/utest/test_extensions/test_zgemm.c @@ -287,6 +287,16 @@ CTEST(zgemm, strided_batch_complex_scalars) ASSERT_DBL_NEAR_TOL(-14.0, c[0], DOUBLE_EPS); ASSERT_DBL_NEAR_TOL(24.0, c[1], DOUBLE_EPS); + + beta[0] = beta[1] = 0.0; + c[0] = 5.0; + c[1] = 6.0; + cblas_zgemm_batch_strided(CblasColMajor, CblasNoTrans, CblasNoTrans, + 1, 1, 1, alpha, a, 1, 2, b, 1, 2, + beta, c, 1, 2, 1); + + ASSERT_DBL_NEAR_TOL(-40.0, c[0], DOUBLE_EPS); + ASSERT_DBL_NEAR_TOL(5.0, c[1], DOUBLE_EPS); } #endif #endif