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

Merge pull request #27693 from nklskyoy:fix-fastgemm-large-mats

DNN fix fastgemm kernel for large mats
This commit is contained in:
Alexander Smorkalov
2025-08-24 17:12:31 +03:00
committed by GitHub
3 changed files with 301 additions and 285 deletions
@@ -62,60 +62,60 @@ void fastGemmPackB(const Mat &B, std::vector<float> &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;
}
}
}
@@ -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<int>((N + NC - 1) / NC) * NC * K;
return static_cast<size_t>((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<size_t>(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<int>((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<size_t>(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<int>((nc + GEMM_NR - 1) / GEMM_NR) * GEMM_NR * esz;
for(int k0 = 0; k0 < K; k0 += KC)
size_t _nc = static_cast<size_t>((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<size_t>(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<int>(tile_idx / total_tiles);
const int m_tiles_index = static_cast<int>((tile_idx - batch_index * total_tiles) / n_tiles);
const int n_tiles_index = static_cast<int>(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<size_t>(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<int>(tile_idx / total_tiles);
const int m_tiles_index = static_cast<int>((tile_idx - batch_index * total_tiles) / n_tiles);
const int n_tiles_index = static_cast<int>(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);
@@ -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<int>((N + NC - 1) / NC) * NC * K;
return static_cast<size_t>((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<size_t>(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<int>((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<size_t>(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<int>((nc + GEMM_NR - 1) / GEMM_NR) * GEMM_NR * esz;
for(int k0 = 0; k0 < K; k0 += KC)
size_t _nc = static_cast<size_t>((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<size_t>(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<int>(tile_idx / total_tiles);
const int m_tiles_index = static_cast<int>((tile_idx - batch_index * total_tiles) / n_tiles);
const int n_tiles_index = static_cast<int>(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<size_t>(FAST_GEMM_F32_MC),
GEMM_NC = static_cast<size_t>(FAST_GEMM_F32_NC),
GEMM_MR = static_cast<size_t>(FAST_GEMM_F32_MR),
GEMM_NR = static_cast<size_t>(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<size_t>(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<int>(tile_idx / total_tiles);
const int m_tiles_index = static_cast<int>((tile_idx - batch_index * total_tiles) / n_tiles);
const int n_tiles_index = static_cast<int>(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<int>((nc + GEMM_NR - 1) / GEMM_NR) * GEMM_NR * esz;
size_t _nc = static_cast<size_t>((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