From cd207549103f529b8af27162474db97703041d7f Mon Sep 17 00:00:00 2001 From: onikolskyy Date: Wed, 20 Aug 2025 12:07:29 +0200 Subject: [PATCH] refactor int -> size_t --- .../dnn/src/layers/cpu_kernels/fast_gemm.cpp | 30 +- .../cpu_kernels/fast_gemm_kernels.default.hpp | 230 ++++++------ .../cpu_kernels/fast_gemm_kernels.simd.hpp | 326 +++++++++--------- 3 files changed, 301 insertions(+), 285 deletions(-) diff --git a/modules/dnn/src/layers/cpu_kernels/fast_gemm.cpp b/modules/dnn/src/layers/cpu_kernels/fast_gemm.cpp index f8fe2bb40e..5aab38dd2e 100644 --- a/modules/dnn/src/layers/cpu_kernels/fast_gemm.cpp +++ b/modules/dnn/src/layers/cpu_kernels/fast_gemm.cpp @@ -62,60 +62,60 @@ void fastGemmPackB(const Mat &B, std::vector &packed_B, bool trans, FastG #if CV_TRY_NEON if (opt.use_neon) { - int size_packed_B = opt_NEON::fastGemmPackBSize(N, K); + size_t size_packed_B = opt_NEON::fastGemmPackBSize(N, K); packed_B.resize(size_packed_B * batch); auto *packed_b = (char*)packed_B.data(); for (int i = 0; i < batch; i++) { opt_NEON::fastGemmPackBKernel(b, packed_b, N, K, ldb0, ldb1, esz); - b += N * K * esz; - packed_b += size_packed_B * esz; + b += (size_t)N * (size_t)K * (size_t)esz; + packed_b += size_packed_B * (size_t)esz; } } else #endif #if CV_TRY_AVX2 if (opt.use_avx2) { - int size_packed_B = opt_AVX2::fastGemmPackBSize(N, K); + size_t size_packed_B = opt_AVX2::fastGemmPackBSize(N, K); packed_B.resize(size_packed_B * batch); auto *packed_b = (char*)packed_B.data(); for (int i = 0; i < batch; i++) { opt_AVX2::fastGemmPackBKernel(b, packed_b, N, K, ldb0, ldb1, esz); - b += N * K * esz; - packed_b += size_packed_B * esz; + b += (size_t)N * (size_t)K * (size_t)esz; + packed_b += size_packed_B * (size_t)esz; } } else #endif #if CV_TRY_AVX if (opt.use_avx) { - int size_packed_B = opt_AVX::fastGemmPackBSize(N, K); + size_t size_packed_B = opt_AVX::fastGemmPackBSize(N, K); packed_B.resize(size_packed_B * batch); auto *packed_b = (char*)packed_B.data(); for (int i = 0; i < batch; i++) { opt_AVX::fastGemmPackBKernel(b, packed_b, N, K, ldb0, ldb1, esz); - b += N * K * esz; - packed_b += size_packed_B * esz; + b += (size_t)N * (size_t)K * (size_t)esz; + packed_b += size_packed_B * (size_t)esz; } } else #endif #if CV_TRY_LASX if (opt.use_lasx) { - int size_packed_B = opt_LASX::fastGemmPackBSize(N, K); + size_t size_packed_B = opt_LASX::fastGemmPackBSize(N, K); packed_B.resize(size_packed_B * batch); auto *packed_b = (char*)packed_B.data(); for (int i = 0; i < batch; i++) { opt_LASX::fastGemmPackBKernel(b, packed_b, N, K, ldb0, ldb1, esz); - b += N * K * esz; - packed_b += size_packed_B * esz; + b += (size_t)N * (size_t)K * (size_t)esz; + packed_b += size_packed_B * (size_t)esz; } } else #endif { - int size_packed_B = cpu_baseline::fastGemmPackBSize(N, K); + size_t size_packed_B = cpu_baseline::fastGemmPackBSize(N, K); packed_B.resize(size_packed_B * batch); auto *packed_b = (char*)packed_B.data(); for (int i = 0; i < batch; i++) { cpu_baseline::fastGemmPackBKernel(b, packed_b, N, K, ldb0, ldb1, esz); - b += N * K * esz; - packed_b += size_packed_B * esz; + b += (size_t)N * (size_t)K * (size_t)esz; + packed_b += size_packed_B * (size_t)esz; } } } diff --git a/modules/dnn/src/layers/cpu_kernels/fast_gemm_kernels.default.hpp b/modules/dnn/src/layers/cpu_kernels/fast_gemm_kernels.default.hpp index f6bd7317a2..980796cf12 100644 --- a/modules/dnn/src/layers/cpu_kernels/fast_gemm_kernels.default.hpp +++ b/modules/dnn/src/layers/cpu_kernels/fast_gemm_kernels.default.hpp @@ -76,47 +76,47 @@ static void fast_gemm_pack##N##suffix( int m, int k, const void* A_, \ namespace cv { namespace dnn { namespace cpu_baseline { -int fastGemmPackBSize(int N, int K); +size_t fastGemmPackBSize(int N, int K); -void fastGemmPackBKernel(const char *B, char *packed_B, int N, int K, int ldb0, int ldb1, int esz); +void fastGemmPackBKernel(const char *B, char *packed_B, size_t N, size_t K, size_t ldb0, size_t ldb1, size_t esz); -void fastGemmKernel(int M, int N, int K, - float alpha, const char *A, int lda0, int lda1, - const char *B, int ldb0, int ldb1, - float beta, char *C, int ldc, int esz, bool multi_thread); -void fastGemmKernel(int M, int N, int K, - float alpha, const char *A, int lda0, int lda1, - const char *packed_B, float beta, char *C, int ldc, int esz, bool multi_thread); +void fastGemmKernel(size_t M, size_t N, size_t K, + float alpha, const char *A, size_t lda0, size_t lda1, + const char *B, size_t ldb0, size_t ldb1, + float beta, char *C, size_t ldc, size_t esz, bool multi_thread); +void fastGemmKernel(size_t M, size_t N, size_t K, + float alpha, const char *A, size_t lda0, size_t lda1, + const char *packed_B, float beta, char *C, size_t ldc, size_t esz, bool multi_thread); void fastGemmBatchKernel(size_t batch, const size_t *A_offsets, const size_t *B_offsets, const size_t *C_offsets, - int M, int N, int K, float alpha, const char *A, int lda0, int lda1, - const char *B, int ldb0, int ldb1, float beta, char *C, int ldc, int esz); + size_t M, size_t N, size_t K, float alpha, const char *A, size_t lda0, size_t lda1, + const char *B, size_t ldb0, size_t ldb1, float beta, char *C, size_t ldc, size_t esz); void fastGemmBatchKernel(size_t batch, const size_t *A_offsets, const size_t *B_offsets, const size_t *C_offsets, - int M, int N, int K, float alpha, const char *A, int lda0, int lda1, - const char *packed_B, float beta, char *C, int ldc, int esz); + size_t M, size_t N, size_t K, float alpha, const char *A, size_t lda0, size_t lda1, + const char *packed_B, float beta, char *C, size_t ldc, size_t esz); FAST_GEMM_IMPLEMENT_PACK(8, _f32, float, float) FAST_GEMM_IMPLEMENT_PACK(12, _f32, float, float) -int fastGemmPackBSize(int N, int K) { - int GEMM_NC = FAST_GEMM_F32_NC, GEMM_NR = FAST_GEMM_F32_NR; - int NC = (((GEMM_NC < N ? GEMM_NC : N) + GEMM_NR - 1) / GEMM_NR) * GEMM_NR; +size_t fastGemmPackBSize(int N, int K) { + size_t GEMM_NC = FAST_GEMM_F32_NC, GEMM_NR = FAST_GEMM_F32_NR; + size_t NC = (((GEMM_NC < N ? GEMM_NC : N) + GEMM_NR - 1) / GEMM_NR) * GEMM_NR; - return static_cast((N + NC - 1) / NC) * NC * K; + return static_cast((N + NC - 1) / NC) * NC * K; } -void fastGemmPackBKernel(const char *B, char *packed_B, int N, int K, int ldb0, int ldb1, int esz) { - int GEMM_NC = FAST_GEMM_F32_NC, GEMM_NR = FAST_GEMM_F32_NR; - int NC = (((GEMM_NC < N ? GEMM_NC : N) + GEMM_NR - 1) / GEMM_NR) * GEMM_NR; - int KC = std::min(FAST_GEMM_F32_PACKED_STRIDE_K, K); +void fastGemmPackBKernel(const char *B, char *packed_B, size_t N, size_t K, size_t ldb0, size_t ldb1, size_t esz) { + size_t GEMM_NC = FAST_GEMM_F32_NC, GEMM_NR = FAST_GEMM_F32_NR; + size_t NC = (((GEMM_NC < N ? GEMM_NC : N) + GEMM_NR - 1) / GEMM_NR) * GEMM_NR; + size_t KC = std::min(static_cast(FAST_GEMM_F32_PACKED_STRIDE_K), K); - int n_tiles = (N + NC - 1) / NC; - for (int r = 0; r < n_tiles; ++r) { - int j0 = r * NC; - int nc = N - j0 < NC ? N - j0 : NC; - int _nc = static_cast((nc + GEMM_NR - 1) / GEMM_NR) * GEMM_NR * esz; - for (int k = 0; k < K; k += KC) { - int kc = K - k < KC ? K - k : KC; + size_t n_tiles = (N + NC - 1) / NC; + for (size_t r = 0; r < n_tiles; ++r) { + size_t j0 = r * NC; + size_t nc = N - j0 < NC ? N - j0 : NC; + size_t _nc = (nc + GEMM_NR - 1) / GEMM_NR * GEMM_NR * esz; + for (size_t k = 0; k < K; k += KC) { + size_t kc = K - k < KC ? K - k : KC; fast_gemm_pack12_f32(nc, kc, B + (k * ldb0 + j0 * ldb1) * esz, ldb1, ldb0, packed_B); packed_B += _nc * kc; } @@ -176,55 +176,55 @@ static void fast_gemm_macro_kernel(int m, int n, int k, } } -void fastGemmKernel(int M, int N, int K, - float alpha, const char *A, int lda0, int lda1, - const char *B, int ldb0, int ldb1, - float beta, char *C, int ldc, int esz, bool multi_thread) { - int GEMM_MC = FAST_GEMM_F32_MC, +void fastGemmKernel(size_t M, size_t N, size_t K, + float alpha, const char *A, size_t lda0, size_t lda1, + const char *B, size_t ldb0, size_t ldb1, + float beta, char *C, size_t ldc, size_t esz, bool multi_thread) { + size_t GEMM_MC = FAST_GEMM_F32_MC, GEMM_NC = FAST_GEMM_F32_NC, GEMM_MR = FAST_GEMM_F32_MR, GEMM_NR = FAST_GEMM_F32_NR; - int MC = (((GEMM_MC < M ? GEMM_MC : M) + GEMM_MR - 1) / GEMM_MR) * GEMM_MR; - int NC = (((GEMM_NC < N ? GEMM_NC : N) + GEMM_NR - 1) / GEMM_NR) * GEMM_NR; - int KC = FAST_GEMM_STORAGE / ((MC + NC) * esz); + size_t MC = (((GEMM_MC < M ? GEMM_MC : M) + GEMM_MR - 1) / GEMM_MR) * GEMM_MR; + size_t NC = (((GEMM_NC < N ? GEMM_NC : N) + GEMM_NR - 1) / GEMM_NR) * GEMM_NR; + size_t KC = FAST_GEMM_STORAGE / ((MC + NC) * esz); KC = KC > 8 ? KC : 8; KC = KC < K ? KC : K; size_t buff_size = KC * (MC + NC) * esz; bool use_stackbuff = buff_size <= FAST_GEMM_MAX_STACKBUF; - int m_tiles = (M + MC - 1) / MC; - int n_tiles = (N + NC - 1) / NC; - int total_tiles = m_tiles * n_tiles; + size_t m_tiles = (M + MC - 1) / MC; + size_t n_tiles = (N + NC - 1) / NC; + size_t total_tiles = m_tiles * n_tiles; auto fn = [&](const Range &r) { char* packed_a = (char*)(use_stackbuff ? alloca(buff_size) : malloc(buff_size)); char* packed_b = packed_a + KC * MC * esz; - int start = r.start; - int end = r.end; + size_t start = r.start; + size_t end = r.end; - for (int tile_idx = start; tile_idx < end; tile_idx++) { - int i0 = (tile_idx / n_tiles) * MC; - int j0 = (tile_idx % n_tiles) * NC; - int mc = M - i0 < MC ? M - i0 : MC; - int nc = N - j0 < NC ? N - j0 : NC; - int ldc_block = ldc; + for (size_t tile_idx = start; tile_idx < end; tile_idx++) { + size_t i0 = (tile_idx / n_tiles) * MC; + size_t j0 = (tile_idx % n_tiles) * NC; + size_t mc = M - i0 < MC ? M - i0 : MC; + size_t nc = N - j0 < NC ? N - j0 : NC; + size_t ldc_block = ldc; char* c_block = C + (i0 * ldc + j0) * esz; if (beta == 0.f) { - for(int i = 0; i < mc; i++) + for(size_t i = 0; i < mc; i++) memset(c_block + i * ldc_block * esz, 0, nc * esz); } else if (beta != 1.f) { - for(int i = 0; i < mc; i++) { + for(size_t i = 0; i < mc; i++) { float* c_i = (float*)c_block + i * ldc_block; - for(int j = 0; j < nc; j++) + for(size_t j = 0; j < nc; j++) c_i[j] *= beta; } } - for(int k0 = 0; k0 < K; k0 += KC) + for(size_t k0 = 0; k0 < K; k0 += KC) { - int kc = K - k0 < KC ? K - k0 : KC; + size_t kc = K - k0 < KC ? K - k0 : KC; fast_gemm_pack8_f32(mc, kc, A + (i0 * lda0 + k0 * lda1) * esz, lda0, lda1, packed_a); fast_gemm_pack12_f32(nc, kc, B + (k0 * ldb0 + j0 * ldb1) * esz, ldb1, ldb0, packed_b); fast_gemm_macro_kernel(mc, nc, kc, packed_a, packed_b, alpha, c_block, ldc_block, esz); @@ -245,36 +245,36 @@ void fastGemmKernel(int M, int N, int K, } } -void fastGemmKernel(int M, int N, int K, - float alpha, const char *A, int lda0, int lda1, - const char *packed_B, float beta, char *C, int ldc, int esz, bool multi_thread) { - int GEMM_MC = FAST_GEMM_F32_MC, +void fastGemmKernel(size_t M, size_t N, size_t K, + float alpha, const char *A, size_t lda0, size_t lda1, + const char *packed_B, float beta, char *C, size_t ldc, size_t esz, bool multi_thread) { + size_t GEMM_MC = FAST_GEMM_F32_MC, GEMM_NC = FAST_GEMM_F32_NC, GEMM_MR = FAST_GEMM_F32_MR, GEMM_NR = FAST_GEMM_F32_NR; - int MC = (((GEMM_MC < M ? GEMM_MC : M) + GEMM_MR - 1) / GEMM_MR) * GEMM_MR; - int NC = (((GEMM_NC < N ? GEMM_NC : N) + GEMM_NR - 1) / GEMM_NR) * GEMM_NR; - int KC = std::min(FAST_GEMM_F32_PACKED_STRIDE_K, K); + size_t MC = (((GEMM_MC < M ? GEMM_MC : M) + GEMM_MR - 1) / GEMM_MR) * GEMM_MR; + size_t NC = (((GEMM_NC < N ? GEMM_NC : N) + GEMM_NR - 1) / GEMM_NR) * GEMM_NR; + size_t KC = std::min(static_cast(FAST_GEMM_F32_PACKED_STRIDE_K), K); size_t buff_size = KC * MC * esz; bool use_stackbuff = buff_size <= FAST_GEMM_MAX_STACKBUF; - int m_tiles = (M + MC - 1) / MC; - int n_tiles = (N + NC - 1) / NC; - int total_tiles = m_tiles * n_tiles; + size_t m_tiles = (M + MC - 1) / MC; + size_t n_tiles = (N + NC - 1) / NC; + size_t total_tiles = m_tiles * n_tiles; auto fn = [&](const Range &r) { char* packed_a = (char*)(use_stackbuff ? alloca(buff_size) : malloc(buff_size)); // TODO: use AutoBuffer const char *packed_b_ = packed_B; - int start = r.start; - int end = r.end; + size_t start = r.start; + size_t end = r.end; - for (int tile_idx = start; tile_idx < end; tile_idx++) { - int i0 = (tile_idx / n_tiles) * MC; - int j0 = (tile_idx % n_tiles) * NC; - int mc = M - i0 < MC ? M - i0 : MC; - int nc = N - j0 < NC ? N - j0 : NC; - int ldc_block = ldc; + for (size_t tile_idx = start; tile_idx < end; tile_idx++) { + size_t i0 = (tile_idx / n_tiles) * MC; + size_t j0 = (tile_idx % n_tiles) * NC; + size_t mc = M - i0 < MC ? M - i0 : MC; + size_t nc = N - j0 < NC ? N - j0 : NC; + size_t ldc_block = ldc; char* c_block = C + (i0 * ldc + j0) * esz; packed_b_ = packed_B + j0 * K * esz; @@ -289,10 +289,10 @@ void fastGemmKernel(int M, int N, int K, } } - int _nc = static_cast((nc + GEMM_NR - 1) / GEMM_NR) * GEMM_NR * esz; - for(int k0 = 0; k0 < K; k0 += KC) + size_t _nc = static_cast((nc + GEMM_NR - 1) / GEMM_NR) * GEMM_NR * esz; + for(size_t k0 = 0; k0 < K; k0 += KC) { - int kc = K - k0 < KC ? K - k0 : KC; + size_t kc = K - k0 < KC ? K - k0 : KC; fast_gemm_pack8_f32(mc, kc, A + (i0 * lda0 + k0 * lda1) * esz, lda0, lda1, packed_a); fast_gemm_macro_kernel(mc, nc, kc, packed_a, packed_b_, alpha, c_block, ldc_block, esz); packed_b_ += _nc * kc; @@ -314,39 +314,39 @@ void fastGemmKernel(int M, int N, int K, } void fastGemmBatchKernel(size_t batch, const size_t *A_offsets, const size_t *B_offsets, const size_t *C_offsets, - int M, int N, int K, float alpha, const char *A, int lda0, int lda1, - const char *B, int ldb0, int ldb1, float beta, char *C, int ldc, int esz) { - int GEMM_MC = FAST_GEMM_F32_MC, + size_t M, size_t N, size_t K, float alpha, const char *A, size_t lda0, size_t lda1, + const char *B, size_t ldb0, size_t ldb1, float beta, char *C, size_t ldc, size_t esz) { + size_t GEMM_MC = FAST_GEMM_F32_MC, GEMM_NC = FAST_GEMM_F32_NC, GEMM_MR = FAST_GEMM_F32_MR, GEMM_NR = FAST_GEMM_F32_NR; - int MC = (((GEMM_MC < M ? GEMM_MC : M) + GEMM_MR - 1) / GEMM_MR) * GEMM_MR; - int NC = (((GEMM_NC < N ? GEMM_NC : N) + GEMM_NR - 1) / GEMM_NR) * GEMM_NR; - int KC = std::min(FAST_GEMM_F32_PACKED_STRIDE_K, K); + size_t MC = (((GEMM_MC < M ? GEMM_MC : M) + GEMM_MR - 1) / GEMM_MR) * GEMM_MR; + size_t NC = (((GEMM_NC < N ? GEMM_NC : N) + GEMM_NR - 1) / GEMM_NR) * GEMM_NR; + size_t KC = std::min(static_cast(FAST_GEMM_F32_PACKED_STRIDE_K), K); size_t buff_size = KC * (MC + NC) * esz; bool use_stackbuff = buff_size <= FAST_GEMM_MAX_STACKBUF; - int m_tiles = (M + MC - 1) / MC; - int n_tiles = (N + NC - 1) / NC; - int total_tiles = m_tiles * n_tiles; + size_t m_tiles = (M + MC - 1) / MC; + size_t n_tiles = (N + NC - 1) / NC; + size_t total_tiles = m_tiles * n_tiles; auto fn = [&](const Range &r) { char* packed_a = (char*)(use_stackbuff ? alloca(buff_size) : malloc(buff_size)); char* packed_b = packed_a + KC * MC * esz; - int start = r.start; - int end = r.end; + size_t start = r.start; + size_t end = r.end; - for (int tile_idx = start; tile_idx < end; tile_idx++) { - const int batch_index = static_cast(tile_idx / total_tiles); - const int m_tiles_index = static_cast((tile_idx - batch_index * total_tiles) / n_tiles); - const int n_tiles_index = static_cast(tile_idx % n_tiles); + for (size_t tile_idx = start; tile_idx < end; tile_idx++) { + const size_t batch_index = tile_idx / total_tiles; + const size_t m_tiles_index = (tile_idx - batch_index * total_tiles) / n_tiles; + const size_t n_tiles_index = tile_idx % n_tiles; - int i0 = m_tiles_index * MC; - int j0 = n_tiles_index * NC; - int mc = M - i0 < MC ? M - i0 : MC; - int nc = N - j0 < NC ? N - j0 : NC; - int ldc_block = ldc; + size_t i0 = m_tiles_index * MC; + size_t j0 = n_tiles_index * NC; + size_t mc = M - i0 < MC ? M - i0 : MC; + size_t nc = N - j0 < NC ? N - j0 : NC; + size_t ldc_block = ldc; const char *a_block = A + A_offsets[batch_index] * esz; const char *b_block = B + B_offsets[batch_index] * esz; char* c_block = C + C_offsets[batch_index] * esz + (i0 * ldc + j0) * esz; @@ -388,39 +388,40 @@ void fastGemmBatchKernel(size_t batch, const size_t *A_offsets, const size_t *B_ } void fastGemmBatchKernel(size_t batch, const size_t *A_offsets, const size_t *B_offsets, const size_t *C_offsets, - int M, int N, int K, float alpha, const char *A, int lda0, int lda1, - const char *packed_B, float beta, char *C, int ldc, int esz) { - int GEMM_MC = FAST_GEMM_F32_MC, + size_t M, size_t N, size_t K, float alpha, const char *A, size_t lda0, size_t lda1, + const char *packed_B, float beta, char *C, size_t ldc, size_t esz) { + size_t GEMM_MC = FAST_GEMM_F32_MC, GEMM_NC = FAST_GEMM_F32_NC, GEMM_MR = FAST_GEMM_F32_MR, GEMM_NR = FAST_GEMM_F32_NR; - int MC = (((GEMM_MC < M ? GEMM_MC : M) + GEMM_MR - 1) / GEMM_MR) * GEMM_MR; - int NC = (((GEMM_NC < N ? GEMM_NC : N) + GEMM_NR - 1) / GEMM_NR) * GEMM_NR; - int KC = std::min(FAST_GEMM_F32_PACKED_STRIDE_K, K); + size_t MC = (((GEMM_MC < M ? GEMM_MC : M) + GEMM_MR - 1) / GEMM_MR) * GEMM_MR; + size_t NC = (((GEMM_NC < N ? GEMM_NC : N) + GEMM_NR - 1) / GEMM_NR) * GEMM_NR; + size_t KC = std::min(static_cast(FAST_GEMM_F32_PACKED_STRIDE_K), K); size_t buff_size = KC * MC * esz; bool use_stackbuff = buff_size <= FAST_GEMM_MAX_STACKBUF; - int m_tiles = (M + MC - 1) / MC; - int n_tiles = (N + NC - 1) / NC; - int total_tiles = m_tiles * n_tiles; + size_t m_tiles = (M + MC - 1) / MC; + size_t n_tiles = (N + NC - 1) / NC; + size_t total_tiles = m_tiles * n_tiles; auto fn = [&](const Range &r) { char* packed_a = (char*)(use_stackbuff ? alloca(buff_size) : malloc(buff_size)); const char *packed_b = packed_B; - int start = r.start; - int end = r.end; + size_t start = r.start; + size_t end = r.end; - for (int tile_idx = start; tile_idx < end; tile_idx++) { - const int batch_index = static_cast(tile_idx / total_tiles); - const int m_tiles_index = static_cast((tile_idx - batch_index * total_tiles) / n_tiles); - const int n_tiles_index = static_cast(tile_idx % n_tiles); + for (size_t tile_idx = start; tile_idx < end; tile_idx++) { + const size_t batch_index = tile_idx / total_tiles; + const size_t m_tiles_index = (tile_idx - batch_index * total_tiles) / n_tiles; + const size_t n_tiles_index = tile_idx % n_tiles; + + size_t i0 = m_tiles_index * MC; + size_t j0 = n_tiles_index * NC; + size_t mc = M - i0 < MC ? M - i0 : MC; + size_t nc = N - j0 < NC ? N - j0 : NC; + size_t ldc_block = ldc; - int i0 = m_tiles_index * MC; - int j0 = n_tiles_index * NC; - int mc = M - i0 < MC ? M - i0 : MC; - int nc = N - j0 < NC ? N - j0 : NC; - int ldc_block = ldc; const char *a_block = A + A_offsets[batch_index] * esz; packed_b = packed_B + B_offsets[batch_index] * esz + j0 * K * esz; char* c_block = C + C_offsets[batch_index] * esz + (i0 * ldc + j0) * esz; @@ -441,7 +442,8 @@ void fastGemmBatchKernel(size_t batch, const size_t *A_offsets, const size_t *B_ { int kc = K - k0 < KC ? K - k0 : KC; // pack a - fast_gemm_pack8_f32(mc, kc, a_block + (i0 * lda0 + k0 * lda1) * esz, lda0, lda1, packed_a); + size_t step_a = (i0 * lda0 + k0 * lda1) * esz; + fast_gemm_pack8_f32(mc, kc, a_block + step_a, lda0, lda1, packed_a); // run kernel fast_gemm_macro_kernel(mc, nc, kc, packed_a, packed_b, alpha, c_block, ldc_block, esz); diff --git a/modules/dnn/src/layers/cpu_kernels/fast_gemm_kernels.simd.hpp b/modules/dnn/src/layers/cpu_kernels/fast_gemm_kernels.simd.hpp index 39bf6800d7..98a3095bf6 100644 --- a/modules/dnn/src/layers/cpu_kernels/fast_gemm_kernels.simd.hpp +++ b/modules/dnn/src/layers/cpu_kernels/fast_gemm_kernels.simd.hpp @@ -115,24 +115,24 @@ namespace cv { namespace dnn { CV_CPU_OPTIMIZATION_NAMESPACE_BEGIN -int fastGemmPackBSize(int N, int K); +size_t fastGemmPackBSize(int N, int K); -void fastGemmPackBKernel(const char *B, char *packed_B, int N, int K, int ldb0, int ldb1, int esz); +void fastGemmPackBKernel(const char *B, char *packed_B, size_t N, size_t K, size_t ldb0, size_t ldb1, size_t esz); -void fastGemmKernel(int M, int N, int K, - float alpha, const char *A, int lda0, int lda1, - const char *B, int ldb0, int ldb1, - float beta, char *C, int ldc, int esz, bool multi_thread); -void fastGemmKernel(int M, int N, int K, - float alpha, const char *A, int lda0, int lda1, - const char *packed_B, float beta, char *C, int ldc, int esz, bool multi_thread); +void fastGemmKernel(size_t M, size_t N, size_t K, + float alpha, const char *A, size_t lda0, size_t lda1, + const char *B, size_t ldb0, size_t ldb1, + float beta, char *C, size_t ldc, size_t esz, bool multi_thread); +void fastGemmKernel(size_t M, size_t N, size_t K, + float alpha, const char *A, size_t lda0, size_t lda1, + const char *packed_B, float beta, char *C, size_t ldc, size_t esz, bool multi_thread); void fastGemmBatchKernel(size_t batch, const size_t *A_offsets, const size_t *B_offsets, const size_t *C_offsets, - int M, int N, int K, float alpha, const char *A, int lda0, int lda1, - const char *B, int ldb0, int ldb1, float beta, char *C, int ldc, int esz); + size_t M, size_t N, size_t K, float alpha, const char *A, size_t lda0, size_t lda1, + const char *B, size_t ldb0, size_t ldb1, float beta, char *C, size_t ldc, size_t esz); void fastGemmBatchKernel(size_t batch, const size_t *A_offsets, const size_t *B_offsets, const size_t *C_offsets, - int M, int N, int K, float alpha, const char *A, int lda0, int lda1, - const char *packed_B, float beta, char *C, int ldc, int esz); + size_t M, size_t N, size_t K, float alpha, const char *A, size_t lda0, size_t lda1, + const char *packed_B, float beta, char *C, size_t ldc, size_t esz); #ifndef CV_CPU_OPTIMIZATION_DECLARATIONS_ONLY @@ -531,88 +531,89 @@ static inline void fast_gemm_macro_kernel(int m, int n, int k, } } -int fastGemmPackBSize(int N, int K) { - int GEMM_NC = FAST_GEMM_F32_NC, GEMM_NR = FAST_GEMM_F32_NR; - int NC = (((GEMM_NC < N ? GEMM_NC : N) + GEMM_NR - 1) / GEMM_NR) * GEMM_NR; +size_t fastGemmPackBSize(int N, int K) { + size_t GEMM_NC = FAST_GEMM_F32_NC, GEMM_NR = FAST_GEMM_F32_NR; + size_t NC = (((GEMM_NC < N ? GEMM_NC : N) + GEMM_NR - 1) / GEMM_NR) * GEMM_NR; - return static_cast((N + NC - 1) / NC) * NC * K; + return static_cast((N + NC - 1) / NC) * NC * K; } -void fastGemmPackBKernel(const char *B, char *packed_B, int N, int K, int ldb0, int ldb1, int esz) { - int GEMM_NC = FAST_GEMM_F32_NC, GEMM_NR = FAST_GEMM_F32_NR; - int NC = (((GEMM_NC < N ? GEMM_NC : N) + GEMM_NR - 1) / GEMM_NR) * GEMM_NR; - int KC = std::min(FAST_GEMM_F32_PACKED_STRIDE_K, K); +void fastGemmPackBKernel(const char *B, char *packed_B, size_t N, size_t K, size_t ldb0, size_t ldb1, size_t esz) { + size_t GEMM_NC = FAST_GEMM_F32_NC, GEMM_NR = FAST_GEMM_F32_NR; + size_t NC = (((GEMM_NC < N ? GEMM_NC : N) + GEMM_NR - 1) / GEMM_NR) * GEMM_NR; + size_t KC = std::min(static_cast(FAST_GEMM_F32_PACKED_STRIDE_K), K); - int n_tiles = (N + NC - 1) / NC; - for (int r = 0; r < n_tiles; ++r) { - int j0 = r * NC; - int nc = N - j0 < NC ? N - j0 : NC; - int _nc = static_cast((nc + GEMM_NR - 1) / GEMM_NR) * GEMM_NR * esz; - for (int k = 0; k < K; k += KC) { - int kc = K - k < KC ? K - k : KC; + size_t n_tiles = (N + NC - 1) / NC; + for (size_t r = 0; r < n_tiles; ++r) { + size_t j0 = r * NC; + size_t nc = N - j0 < NC ? N - j0 : NC; + size_t _nc = (nc + GEMM_NR - 1) / GEMM_NR * GEMM_NR * esz; + for (size_t k = 0; k < K; k += KC) { + size_t kc = K - k < KC ? K - k : KC; + size_t step = (k * ldb0 + j0 * ldb1) * esz; #if CV_NEON && CV_NEON_AARCH64 - fast_gemm_pack12_f32(nc, kc, B + (k * ldb0 + j0 * ldb1) * esz, ldb1, ldb0, packed_B); + fast_gemm_pack12_f32(nc, kc, B + step, ldb1, ldb0, packed_B); #elif CV_AVX - fast_gemm_pack8_f32(nc, kc, B + (k * ldb0 + j0 * ldb1) * esz, ldb1, ldb0, packed_B); + fast_gemm_pack8_f32(nc, kc, B + step, ldb1, ldb0, packed_B); #elif CV_LASX - fast_gemm_pack16_f32(nc, kc, B + (k * ldb0 + j0 * ldb1) * esz, ldb1, ldb0, packed_B); + fast_gemm_pack16_f32(nc, kc, B + step, ldb1, ldb0, packed_B); #elif CV_SIMD128 - fast_gemm_pack12_f32(nc, kc, B + (k * ldb0 + j0 * ldb1) * esz, ldb1, ldb0, packed_B); + fast_gemm_pack12_f32(nc, kc, B + step, ldb1, ldb0, packed_B); #endif packed_B += _nc * kc; } } } -void fastGemmKernel(int M, int N, int K, - float alpha, const char *A, int lda0, int lda1, - const char *B, int ldb0, int ldb1, - float beta, char *C, int ldc, int esz, bool multi_thread) { - int GEMM_MC = FAST_GEMM_F32_MC, +void fastGemmKernel(size_t M, size_t N, size_t K, + float alpha, const char *A, size_t lda0, size_t lda1, + const char *B, size_t ldb0, size_t ldb1, + float beta, char *C, size_t ldc, size_t esz, bool multi_thread) { + size_t GEMM_MC = FAST_GEMM_F32_MC, GEMM_NC = FAST_GEMM_F32_NC, GEMM_MR = FAST_GEMM_F32_MR, GEMM_NR = FAST_GEMM_F32_NR; - int MC = (((GEMM_MC < M ? GEMM_MC : M) + GEMM_MR - 1) / GEMM_MR) * GEMM_MR; - int NC = (((GEMM_NC < N ? GEMM_NC : N) + GEMM_NR - 1) / GEMM_NR) * GEMM_NR; - int KC = FAST_GEMM_STORAGE / ((MC + NC) * esz); + size_t MC = (((GEMM_MC < M ? GEMM_MC : M) + GEMM_MR - 1) / GEMM_MR) * GEMM_MR; + size_t NC = (((GEMM_NC < N ? GEMM_NC : N) + GEMM_NR - 1) / GEMM_NR) * GEMM_NR; + size_t KC = FAST_GEMM_STORAGE / ((MC + NC) * esz); KC = KC > 8 ? KC : 8; KC = KC < K ? KC : K; size_t buff_size = KC * (MC + NC) * esz; bool use_stackbuff = buff_size <= FAST_GEMM_MAX_STACKBUF; - int m_tiles = (M + MC - 1) / MC; - int n_tiles = (N + NC - 1) / NC; - int total_tiles = m_tiles * n_tiles; + size_t m_tiles = (M + MC - 1) / MC; + size_t n_tiles = (N + NC - 1) / NC; + size_t total_tiles = m_tiles * n_tiles; auto fn = [&](const Range &r) { char* packed_a = (char*)(use_stackbuff ? alloca(buff_size) : malloc(buff_size)); char* packed_b = packed_a + KC * MC * esz; - int start = r.start; - int end = r.end; + size_t start = r.start; + size_t end = r.end; - for (int tile_idx = start; tile_idx < end; tile_idx++) { - int i0 = (tile_idx / n_tiles) * MC; - int j0 = (tile_idx % n_tiles) * NC; - int mc = M - i0 < MC ? M - i0 : MC; - int nc = N - j0 < NC ? N - j0 : NC; - int ldc_block = ldc; + for (size_t tile_idx = start; tile_idx < end; tile_idx++) { + size_t i0 = (tile_idx / n_tiles) * MC; + size_t j0 = (tile_idx % n_tiles) * NC; + size_t mc = M - i0 < MC ? M - i0 : MC; + size_t nc = N - j0 < NC ? N - j0 : NC; + size_t ldc_block = ldc; char* c_block = C + (i0 * ldc + j0) * esz; if (beta == 0.f) { - for(int i = 0; i < mc; i++) + for(size_t i = 0; i < mc; i++) memset(c_block + i * ldc_block * esz, 0, nc * esz); } else if (beta != 1.f) { - for(int i = 0; i < mc; i++) { + for(size_t i = 0; i < mc; i++) { float* c_i = (float*)c_block + i * ldc_block; - for(int j = 0; j < nc; j++) + for(size_t j = 0; j < nc; j++) c_i[j] *= beta; } } - for(int k0 = 0; k0 < K; k0 += KC) + for(size_t k0 = 0; k0 < K; k0 += KC) { - int kc = K - k0 < KC ? K - k0 : KC; + size_t kc = K - k0 < KC ? K - k0 : KC; // pack a #if CV_NEON && CV_NEON_AARCH64 fast_gemm_pack8_f32(mc, kc, A + (i0 * lda0 + k0 * lda1) * esz, lda0, lda1, packed_a); @@ -623,12 +624,11 @@ void fastGemmKernel(int M, int N, int K, #elif CV_SIMD128 fast_gemm_pack8_f32(mc, kc, A + (i0 * lda0 + k0 * lda1) * esz, lda0, lda1, packed_a); #endif - // pack b #if CV_NEON && CV_NEON_AARCH64 fast_gemm_pack12_f32(nc, kc, B + (k0 * ldb0 + j0 * ldb1) * esz, ldb1, ldb0, packed_b); #elif CV_AVX - fast_gemm_pack8_f32(nc, kc, B + (k0 * ldb0 + j0 * ldb1) * esz, ldb1, ldb0, packed_b); + fast_gemm_pack8_f32(nc, kc, B + (k0 * ldb0 + j0 * ldb1) * esz, ldb1, ldb0, packed_b); #elif CV_LASX fast_gemm_pack16_f32(nc, kc, B + (k0 * ldb0 + j0 * ldb1) * esz, ldb1, ldb0, packed_b); #elif CV_SIMD128 @@ -655,63 +655,67 @@ void fastGemmKernel(int M, int N, int K, } -void fastGemmKernel(int M, int N, int K, - float alpha, const char *A, int lda0, int lda1, - const char *packed_B, float beta, char *C, int ldc, int esz, bool multi_thread) { - int GEMM_MC = FAST_GEMM_F32_MC, +void fastGemmKernel(size_t M, size_t N, size_t K, + float alpha, const char *A, size_t lda0, size_t lda1, + const char *packed_B, float beta, char *C, size_t ldc, size_t esz, bool multi_thread) { + size_t GEMM_MC = FAST_GEMM_F32_MC, GEMM_NC = FAST_GEMM_F32_NC, GEMM_MR = FAST_GEMM_F32_MR, GEMM_NR = FAST_GEMM_F32_NR; - int MC = (((GEMM_MC < M ? GEMM_MC : M) + GEMM_MR - 1) / GEMM_MR) * GEMM_MR; - int NC = (((GEMM_NC < N ? GEMM_NC : N) + GEMM_NR - 1) / GEMM_NR) * GEMM_NR; - int KC = std::min(FAST_GEMM_F32_PACKED_STRIDE_K, K); + size_t MC = (((GEMM_MC < M ? GEMM_MC : M) + GEMM_MR - 1) / GEMM_MR) * GEMM_MR; + size_t NC = (((GEMM_NC < N ? GEMM_NC : N) + GEMM_NR - 1) / GEMM_NR) * GEMM_NR; + size_t KC = std::min(static_cast(FAST_GEMM_F32_PACKED_STRIDE_K), K); size_t buff_size = KC * MC * esz; bool use_stackbuff = buff_size <= FAST_GEMM_MAX_STACKBUF; - int m_tiles = (M + MC - 1) / MC; - int n_tiles = (N + NC - 1) / NC; - int total_tiles = m_tiles * n_tiles; + size_t m_tiles = (M + MC - 1) / MC; + size_t n_tiles = (N + NC - 1) / NC; + size_t total_tiles = m_tiles * n_tiles; auto fn = [&](const Range &r) { char* packed_a = (char*)(use_stackbuff ? alloca(buff_size) : malloc(buff_size)); // TODO: use AutoBuffer const char *packed_b_ = packed_B; - int start = r.start; - int end = r.end; + size_t start = r.start; + size_t end = r.end; - for (int tile_idx = start; tile_idx < end; tile_idx++) { - int i0 = (tile_idx / n_tiles) * MC; - int j0 = (tile_idx % n_tiles) * NC; - int mc = M - i0 < MC ? M - i0 : MC; - int nc = N - j0 < NC ? N - j0 : NC; - int ldc_block = ldc; - char* c_block = C + (i0 * ldc + j0) * esz; - packed_b_ = packed_B + j0 * K * esz; + for (size_t tile_idx = start; tile_idx < end; tile_idx++) { + size_t i0 = (tile_idx / n_tiles) * MC; + size_t j0 = (tile_idx % n_tiles) * NC; + size_t mc = M - i0 < MC ? M - i0 : MC; + size_t nc = N - j0 < NC ? N - j0 : NC; + size_t ldc_block = ldc; + size_t step_c = ((size_t)i0 * (size_t)ldc + (size_t)j0) * (size_t)esz; + char* c_block = C + step_c; + + size_t step_b = (size_t)j0 * (size_t)K * (size_t)esz; + packed_b_ = packed_B + step_b; if (beta == 0.f) { - for(int i = 0; i < mc; i++) + for(size_t i = 0; i < mc; i++) memset(c_block + i * ldc_block * esz, 0, nc * esz); } else if (beta != 1.f) { - for(int i = 0; i < mc; i++) { + for(size_t i = 0; i < mc; i++) { float* c_i = (float*)c_block + i * ldc_block; - for(int j = 0; j < nc; j++) + for(size_t j = 0; j < nc; j++) c_i[j] *= beta; } } - int _nc = static_cast((nc + GEMM_NR - 1) / GEMM_NR) * GEMM_NR * esz; - for(int k0 = 0; k0 < K; k0 += KC) + size_t _nc = static_cast((nc + GEMM_NR - 1) / GEMM_NR) * GEMM_NR * esz; + for(size_t k0 = 0; k0 < K; k0 += KC) { - int kc = K - k0 < KC ? K - k0 : KC; + size_t kc = K - k0 < KC ? K - k0 : KC; + size_t step_a = (i0 * lda0 + k0 * lda1) * esz; // pack a #if CV_NEON && CV_NEON_AARCH64 - fast_gemm_pack8_f32(mc, kc, A + (i0 * lda0 + k0 * lda1) * esz, lda0, lda1, packed_a); + fast_gemm_pack8_f32(mc, kc, A + step_a, lda0, lda1, packed_a); #elif CV_AVX - fast_gemm_pack12_f32(mc, kc, A + (i0 * lda0 + k0 * lda1) * esz, lda0, lda1, packed_a); + fast_gemm_pack12_f32(mc, kc, A + step_a, lda0, lda1, packed_a); #elif CV_LASX - fast_gemm_pack12_f32(mc, kc, A + (i0 * lda0 + k0 * lda1) * esz, lda0, lda1, packed_a); + fast_gemm_pack12_f32(mc, kc, A + step_a, lda0, lda1, packed_a); #elif CV_SIMD128 - fast_gemm_pack8_f32(mc, kc, A + (i0 * lda0 + k0 * lda1) * esz, lda0, lda1, packed_a); + fast_gemm_pack8_f32(mc, kc, A + step_a, lda0, lda1, packed_a); #endif // run kernel @@ -735,77 +739,80 @@ void fastGemmKernel(int M, int N, int K, } void fastGemmBatchKernel(size_t batch, const size_t *A_offsets, const size_t *B_offsets, const size_t *C_offsets, - int M, int N, int K, float alpha, const char *A, int lda0, int lda1, - const char *B, int ldb0, int ldb1, float beta, char *C, int ldc, int esz) { - int GEMM_MC = FAST_GEMM_F32_MC, + size_t M, size_t N, size_t K, float alpha, const char *A, size_t lda0, size_t lda1, + const char *B, size_t ldb0, size_t ldb1, float beta, char *C, size_t ldc, size_t esz) { + size_t GEMM_MC = FAST_GEMM_F32_MC, GEMM_NC = FAST_GEMM_F32_NC, GEMM_MR = FAST_GEMM_F32_MR, GEMM_NR = FAST_GEMM_F32_NR; - int MC = (((GEMM_MC < M ? GEMM_MC : M) + GEMM_MR - 1) / GEMM_MR) * GEMM_MR; - int NC = (((GEMM_NC < N ? GEMM_NC : N) + GEMM_NR - 1) / GEMM_NR) * GEMM_NR; - int KC = std::min(FAST_GEMM_F32_PACKED_STRIDE_K, K); + size_t MC = (((GEMM_MC < M ? GEMM_MC : M) + GEMM_MR - 1) / GEMM_MR) * GEMM_MR; + size_t NC = (((GEMM_NC < N ? GEMM_NC : N) + GEMM_NR - 1) / GEMM_NR) * GEMM_NR; + size_t KC = std::min(static_cast(FAST_GEMM_F32_PACKED_STRIDE_K), K); size_t buff_size = KC * (MC + NC) * esz; bool use_stackbuff = buff_size <= FAST_GEMM_MAX_STACKBUF; - int m_tiles = (M + MC - 1) / MC; - int n_tiles = (N + NC - 1) / NC; - int total_tiles = m_tiles * n_tiles; + size_t m_tiles = (M + MC - 1) / MC; + size_t n_tiles = (N + NC - 1) / NC; + size_t total_tiles = m_tiles * n_tiles; auto fn = [&](const Range &r) { char* packed_a = (char*)(use_stackbuff ? alloca(buff_size) : malloc(buff_size)); char* packed_b = packed_a + KC * MC * esz; - int start = r.start; - int end = r.end; + size_t start = r.start; + size_t end = r.end; - for (int tile_idx = start; tile_idx < end; tile_idx++) { - const int batch_index = static_cast(tile_idx / total_tiles); - const int m_tiles_index = static_cast((tile_idx - batch_index * total_tiles) / n_tiles); - const int n_tiles_index = static_cast(tile_idx % n_tiles); + for (size_t tile_idx = start; tile_idx < end; tile_idx++) { + const size_t batch_index = tile_idx / total_tiles; + const size_t m_tiles_index = (tile_idx - batch_index * total_tiles) / n_tiles; + const size_t n_tiles_index = tile_idx % n_tiles; - int i0 = m_tiles_index * MC; - int j0 = n_tiles_index * NC; - int mc = M - i0 < MC ? M - i0 : MC; - int nc = N - j0 < NC ? N - j0 : NC; - int ldc_block = ldc; + size_t i0 = m_tiles_index * MC; + size_t j0 = n_tiles_index * NC; + size_t mc = M - i0 < MC ? M - i0 : MC; + size_t nc = N - j0 < NC ? N - j0 : NC; + size_t ldc_block = ldc; const char *a_block = A + A_offsets[batch_index] * esz; const char *b_block = B + B_offsets[batch_index] * esz; + char* c_block = C + C_offsets[batch_index] * esz + (i0 * ldc + j0) * esz; if (beta == 0.f) { - for(int i = 0; i < mc; i++) + for(size_t i = 0; i < mc; i++) memset(c_block + i * ldc_block * esz, 0, nc * esz); } else if (beta != 1.f) { - for(int i = 0; i < mc; i++) { + for(size_t i = 0; i < mc; i++) { float* c_i = (float*)c_block + i * ldc_block; - for(int j = 0; j < nc; j++) + for(size_t j = 0; j < nc; j++) c_i[j] *= beta; } } - for(int k0 = 0; k0 < K; k0 += KC) + for(size_t k0 = 0; k0 < K; k0 += KC) { - int kc = K - k0 < KC ? K - k0 : KC; + size_t kc = K - k0 < KC ? K - k0 : KC; + size_t step_a = (i0 * lda0 + k0 * lda1) * esz; + // pack a #if CV_NEON && CV_NEON_AARCH64 - fast_gemm_pack8_f32(mc, kc, a_block + (i0 * lda0 + k0 * lda1) * esz, lda0, lda1, packed_a); + fast_gemm_pack8_f32(mc, kc, a_block + step_a, lda0, lda1, packed_a); #elif CV_AVX - fast_gemm_pack12_f32(mc, kc, a_block + (i0 * lda0 + k0 * lda1) * esz, lda0, lda1, packed_a); + fast_gemm_pack12_f32(mc, kc, a_block + step_a, lda0, lda1, packed_a); #elif CV_LASX - fast_gemm_pack12_f32(mc, kc, a_block + (i0 * lda0 + k0 * lda1) * esz, lda0, lda1, packed_a); + fast_gemm_pack12_f32(mc, kc, a_block + step_a, lda0, lda1, packed_a); #elif CV_SIMD128 - fast_gemm_pack8_f32(mc, kc, a_block + (i0 * lda0 + k0 * lda1) * esz, lda0, lda1, packed_a); + fast_gemm_pack8_f32(mc, kc, a_block + step_a, lda0, lda1, packed_a); #endif - + size_t step_b = (k0 * ldb0 + j0 * ldb1) * esz; // pack b #if CV_NEON && CV_NEON_AARCH64 - fast_gemm_pack12_f32(nc, kc, b_block + (k0 * ldb0 + j0 * ldb1) * esz, ldb1, ldb0, packed_b); + fast_gemm_pack12_f32(nc, kc, b_block + step_b, ldb1, ldb0, packed_b); #elif CV_AVX - fast_gemm_pack8_f32(nc, kc, b_block + (k0 * ldb0 + j0 * ldb1) * esz, ldb1, ldb0, packed_b); + fast_gemm_pack8_f32(nc, kc, b_block + step_b, ldb1, ldb0, packed_b); #elif CV_LASX - fast_gemm_pack16_f32(nc, kc, b_block + (k0 * ldb0 + j0 * ldb1) * esz, ldb1, ldb0, packed_b); + fast_gemm_pack16_f32(nc, kc, b_block + step_b, ldb1, ldb0, packed_b); #elif CV_SIMD128 - fast_gemm_pack12_f32(nc, kc, b_block + (k0 * ldb0 + j0 * ldb1) * esz, ldb1, ldb0, packed_b); + fast_gemm_pack12_f32(nc, kc, b_block + step_b, ldb1, ldb0, packed_b); #endif // run kernel @@ -825,67 +832,74 @@ void fastGemmBatchKernel(size_t batch, const size_t *A_offsets, const size_t *B_ } void fastGemmBatchKernel(size_t batch, const size_t *A_offsets, const size_t *B_offsets, const size_t *C_offsets, - int M, int N, int K, float alpha, const char *A, int lda0, int lda1, - const char *packed_B, float beta, char *C, int ldc, int esz) { - int GEMM_MC = FAST_GEMM_F32_MC, - GEMM_NC = FAST_GEMM_F32_NC, - GEMM_MR = FAST_GEMM_F32_MR, - GEMM_NR = FAST_GEMM_F32_NR; + size_t M, size_t N, size_t K, float alpha, const char *A, size_t lda0, size_t lda1, + const char *packed_B, float beta, char *C, size_t ldc, size_t esz) { + size_t GEMM_MC = static_cast(FAST_GEMM_F32_MC), + GEMM_NC = static_cast(FAST_GEMM_F32_NC), + GEMM_MR = static_cast(FAST_GEMM_F32_MR), + GEMM_NR = static_cast(FAST_GEMM_F32_NR); - int MC = (((GEMM_MC < M ? GEMM_MC : M) + GEMM_MR - 1) / GEMM_MR) * GEMM_MR; - int NC = (((GEMM_NC < N ? GEMM_NC : N) + GEMM_NR - 1) / GEMM_NR) * GEMM_NR; - int KC = std::min(FAST_GEMM_F32_PACKED_STRIDE_K, K); + size_t MC = (((GEMM_MC < M ? GEMM_MC : M) + GEMM_MR - 1) / GEMM_MR) * GEMM_MR; + size_t NC = (((GEMM_NC < N ? GEMM_NC : N) + GEMM_NR - 1) / GEMM_NR) * GEMM_NR; + size_t KC = std::min( static_cast(FAST_GEMM_F32_PACKED_STRIDE_K), K); size_t buff_size = KC * MC * esz; bool use_stackbuff = buff_size <= FAST_GEMM_MAX_STACKBUF; - int m_tiles = (M + MC - 1) / MC; - int n_tiles = (N + NC - 1) / NC; - int total_tiles = m_tiles * n_tiles; + size_t m_tiles = (M + MC - 1) / MC; + size_t n_tiles = (N + NC - 1) / NC; + size_t total_tiles = m_tiles * n_tiles; auto fn = [&](const Range &r) { char* packed_a = (char*)(use_stackbuff ? alloca(buff_size) : malloc(buff_size)); const char *packed_b = packed_B; - int start = r.start; - int end = r.end; + size_t start = r.start; + size_t end = r.end; - for (int tile_idx = start; tile_idx < end; tile_idx++) { - const int batch_index = static_cast(tile_idx / total_tiles); - const int m_tiles_index = static_cast((tile_idx - batch_index * total_tiles) / n_tiles); - const int n_tiles_index = static_cast(tile_idx % n_tiles); + for (size_t tile_idx = start; tile_idx < end; tile_idx++) { + const size_t batch_index = tile_idx / total_tiles; + const size_t m_tiles_index = (tile_idx - batch_index * total_tiles) / n_tiles; + const size_t n_tiles_index = tile_idx % n_tiles; - int i0 = m_tiles_index * MC; - int j0 = n_tiles_index * NC; - int mc = M - i0 < MC ? M - i0 : MC; - int nc = N - j0 < NC ? N - j0 : NC; - int ldc_block = ldc; + size_t i0 = m_tiles_index * MC; + size_t j0 = n_tiles_index * NC; + size_t mc = M - i0 < MC ? M - i0 : MC; + size_t nc = N - j0 < NC ? N - j0 : NC; + size_t ldc_block = ldc; const char *a_block = A + A_offsets[batch_index] * esz; - packed_b = packed_B + B_offsets[batch_index] * esz + j0 * K * esz; - char* c_block = C + C_offsets[batch_index] * esz + (i0 * ldc + j0) * esz; + + size_t step_b = B_offsets[batch_index] * esz + j0 * K * esz; + size_t step_c = C_offsets[batch_index] * esz + (i0 * ldc + j0) * esz; + + packed_b = packed_B + step_b; + char* c_block = C + step_c; if (beta == 0.f) { - for(int i = 0; i < mc; i++) + for(size_t i = 0; i < mc; i++){ memset(c_block + i * ldc_block * esz, 0, nc * esz); + } } else if (beta != 1.f) { - for(int i = 0; i < mc; i++) { + for(size_t i = 0; i < mc; i++) { float* c_i = (float*)c_block + i * ldc_block; - for(int j = 0; j < nc; j++) + for(size_t j = 0; j < nc; j++) c_i[j] *= beta; } } - int _nc = static_cast((nc + GEMM_NR - 1) / GEMM_NR) * GEMM_NR * esz; + size_t _nc = static_cast((nc + GEMM_NR - 1) / GEMM_NR) * GEMM_NR * esz; for(int k0 = 0; k0 < K; k0 += KC) { - int kc = K - k0 < KC ? K - k0 : KC; + size_t kc = K - k0 < KC ? K - k0 : KC; + // size_t step = ((size_t)k0 * (size_t)lda1 + (size_t)i0 * (size_t)lda0) * (size_t)esz; + size_t step = (k0 * lda1 + i0 * lda0) * esz; // pack a #if CV_NEON && CV_NEON_AARCH64 - fast_gemm_pack8_f32(mc, kc, a_block + (i0 * lda0 + k0 * lda1) * esz, lda0, lda1, packed_a); + fast_gemm_pack8_f32(mc, kc, a_block + step, lda0, lda1, packed_a); #elif CV_AVX - fast_gemm_pack12_f32(mc, kc, a_block + (i0 * lda0 + k0 * lda1) * esz, lda0, lda1, packed_a); + fast_gemm_pack12_f32(mc, kc, a_block + step, lda0, lda1, packed_a); #elif CV_LASX - fast_gemm_pack12_f32(mc, kc, a_block + (i0 * lda0 + k0 * lda1) * esz, lda0, lda1, packed_a); + fast_gemm_pack12_f32(mc, kc, a_block + step, lda0, lda1, packed_a); #elif CV_SIMD128 - fast_gemm_pack8_f32(mc, kc, a_block + (i0 * lda0 + k0 * lda1) * esz, lda0, lda1, packed_a); + fast_gemm_pack8_f32(mc, kc, a_block + step, lda0, lda1, packed_a); #endif // run kernel