From 2c7f22d65417295d8fb69e8218f962530212a694 Mon Sep 17 00:00:00 2001 From: nihui Date: Thu, 16 Jul 2026 23:41:58 +0800 Subject: [PATCH 1/4] qwq --- src/layer/arm/gemm_arm.cpp | 258 + src/layer/arm/gemm_arm.h | 10 + src/layer/arm/gemm_arm_asimddp.cpp | 26 + src/layer/arm/gemm_arm_i8mm.cpp | 26 + src/layer/arm/gemm_arm_svei8mm.cpp | 26 + src/layer/arm/gemm_wq_int8.h | 6135 ++++++++++++ src/layer/arm/multiheadattention_arm.cpp | 608 +- src/layer/arm/multiheadattention_arm.h | 5 + src/layer/gemm.cpp | 164 +- src/layer/loongarch/gemm_loongarch.cpp | 238 + src/layer/loongarch/gemm_loongarch.h | 11 + src/layer/loongarch/gemm_wq_int8.h | 6297 ++++++++++++ .../multiheadattention_loongarch.cpp | 397 +- .../loongarch/multiheadattention_loongarch.h | 5 + src/layer/mips/gemm_mips.cpp | 235 + src/layer/mips/gemm_mips.h | 9 + src/layer/mips/gemm_mips_mmi.cpp | 25 + src/layer/mips/gemm_wq_int8.h | 5273 ++++++++++ src/layer/mips/multiheadattention_mips.cpp | 397 +- src/layer/mips/multiheadattention_mips.h | 5 + src/layer/multiheadattention.cpp | 331 +- src/layer/riscv/gemm_riscv.cpp | 261 +- src/layer/riscv/gemm_riscv.h | 8 + src/layer/riscv/gemm_wq_int8.h | 2217 +++++ src/layer/riscv/multiheadattention_riscv.cpp | 921 ++ src/layer/riscv/multiheadattention_riscv.h | 40 + src/layer/x86/gemm_wq_int8.h | 8579 +++++++++++++++++ src/layer/x86/gemm_x86.cpp | 299 +- src/layer/x86/gemm_x86.h | 8 + src/layer/x86/gemm_x86_avx2.cpp | 34 + src/layer/x86/gemm_x86_avx512vnni.cpp | 34 + src/layer/x86/gemm_x86_avxvnni.cpp | 34 + src/layer/x86/gemm_x86_avxvnniint8.cpp | 34 + src/layer/x86/gemm_x86_xop.cpp | 9 + src/layer/x86/multiheadattention_x86.cpp | 601 +- src/layer/x86/multiheadattention_x86.h | 5 + 36 files changed, 33333 insertions(+), 232 deletions(-) create mode 100644 src/layer/arm/gemm_arm_svei8mm.cpp create mode 100644 src/layer/arm/gemm_wq_int8.h create mode 100644 src/layer/loongarch/gemm_wq_int8.h create mode 100644 src/layer/mips/gemm_wq_int8.h create mode 100644 src/layer/riscv/gemm_wq_int8.h create mode 100644 src/layer/riscv/multiheadattention_riscv.cpp create mode 100644 src/layer/riscv/multiheadattention_riscv.h create mode 100644 src/layer/x86/gemm_wq_int8.h diff --git a/src/layer/arm/gemm_arm.cpp b/src/layer/arm/gemm_arm.cpp index 38528293aa7f..1822c15396a0 100644 --- a/src/layer/arm/gemm_arm.cpp +++ b/src/layer/arm/gemm_arm.cpp @@ -25,6 +25,10 @@ namespace ncnn { #endif #endif +#if NCNN_WEIGHT_QUANT +#include "gemm_wq_int8.h" +#endif + Gemm_arm::Gemm_arm() { #if __ARM_NEON @@ -4582,10 +4586,124 @@ static int gemm_AT_BT_arm(const Mat& AT, const Mat& BT, const Mat& C, Mat& top_b return 0; } +#if NCNN_WEIGHT_QUANT +static int gemm_BT_arm_wq_int8(const Mat& A, const Mat& packed_B, const Mat& packed_B_descales, const Mat& input_scales, const Mat& C, Mat& top_blob, int broadcast_type_C, int N, int K, int block_size, int transA, int output_transpose, float alpha, float beta, int constant_TILE_M, int constant_TILE_N, int constant_TILE_K, int nT, const Option& opt) +{ + const int M = transA ? A.w : (A.dims == 3 ? A.c : A.h) * A.elempack; + const int block_count = (K + block_size - 1) / block_size; + const Mat BT = packed_B.reshape(K, N); + const Mat BT_descales = packed_B_descales.reshape(block_count, N); + int TILE_M, TILE_N, TILE_K; + get_optimal_tile_mnk_wq_int8(M, N, K, constant_TILE_M, constant_TILE_N, constant_TILE_K, TILE_M, TILE_N, TILE_K, nT); + + (void)TILE_K; + const int mr = std::min(M, TILE_M); + const int nr = std::min(N, TILE_N); + const int nn_M = (M + TILE_M - 1) / TILE_M; + const int nn_N = (N + TILE_N - 1) / TILE_N; + const float* input_scale_ptr = input_scales; + + Mat topT(nr * mr, 1, nT, (size_t)4u, opt.workspace_allocator); + if (topT.empty()) + return -100; + + if (nT > nn_M) + { + Mat AT(K, mr, nn_M, (size_t)1u, opt.workspace_allocator); + Mat AT_descales(block_count, mr, nn_M, (size_t)4u, opt.workspace_allocator); + if (AT.empty() || AT_descales.empty()) + return -100; + + #pragma omp parallel for num_threads(nT) + for (int ppi = 0; ppi < nn_M; ppi++) + { + const int i = ppi * TILE_M; + const int max_ii = std::min(M - i, TILE_M); + + Mat AT_tile = AT.channel(i / TILE_M); + Mat AT_descales_tile = AT_descales.channel(i / TILE_M); + + if (transA) + transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); + else + quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); + } + + const int nn_MN = nn_M * nn_N; + #pragma omp parallel for num_threads(nT) + for (int ppij = 0; ppij < nn_MN; ppij++) + { + const int ppi = ppij / nn_N; + const int ppj = ppij % nn_N; + + const int i = ppi * TILE_M; + const int j = ppj * TILE_N; + + const int max_ii = std::min(M - i, TILE_M); + const int max_jj = std::min(N - j, TILE_N); + + Mat AT_tile = AT.channel(i / TILE_M); + Mat AT_descales_tile = AT_descales.channel(i / TILE_M); + Mat BT_tile = BT.row_range(j, max_jj); + Mat BT_descales_tile = BT_descales.row_range(j, max_jj); + Mat topT_tile = topT.channel(get_omp_thread_num()); + + gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); + if (output_transpose) + transpose_unpack_output_tile_wq_int8(topT_tile, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, N, alpha, beta); + else + unpack_output_tile_wq_int8(topT_tile, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, N, alpha, beta); + } + } + else + { + Mat ATX(K, mr, nT, (size_t)1u, opt.workspace_allocator); + Mat ATX_descales(block_count, mr, nT, (size_t)4u, opt.workspace_allocator); + if (ATX.empty() || ATX_descales.empty()) + return -100; + + #pragma omp parallel for num_threads(nT) + for (int ppi = 0; ppi < nn_M; ppi++) + { + const int i = ppi * TILE_M; + const int max_ii = std::min(M - i, TILE_M); + + Mat AT_tile = ATX.channel(get_omp_thread_num()); + Mat AT_descales_tile = ATX_descales.channel(get_omp_thread_num()); + Mat topT_tile = topT.channel(get_omp_thread_num()); + + if (transA) + transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); + else + quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); + + for (int j = 0; j < N; j += TILE_N) + { + const int max_jj = std::min(N - j, TILE_N); + + Mat BT_tile = BT.row_range(j, max_jj); + Mat BT_descales_tile = BT_descales.row_range(j, max_jj); + gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); + if (output_transpose) + transpose_unpack_output_tile_wq_int8(topT_tile, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, N, alpha, beta); + else + unpack_output_tile_wq_int8(topT_tile, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, N, alpha, beta); + } + } + } + + return 0; +} +#endif // NCNN_WEIGHT_QUANT + int Gemm_arm::create_pipeline(const Option& opt) { if (weight_block_quantize) { +#if NCNN_WEIGHT_QUANT + if (quantize_term / 100 == 8) + return create_pipeline_wq_int8(opt); +#endif return 0; } @@ -4745,10 +4863,150 @@ int Gemm_arm::create_pipeline(const Option& opt) return 0; } +int Gemm_arm::destroy_pipeline(const Option& /*opt*/) +{ +#if NCNN_WEIGHT_QUANT + BT_data_wq_int8.release(); + BT_data_wq_int8_descales.release(); +#endif + + return 0; +} + +#if NCNN_WEIGHT_QUANT +int Gemm_arm::create_pipeline_wq_int8(const Option& opt) +{ + if (!BT_data_wq_int8.empty()) + return 0; + + if (B_data.empty() || B_data_quantize_scales.empty()) + return -100; + + const int block_size_code = quantize_term % 10; + const int block_size = block_size_code == 0 ? 32 : block_size_code == 1 ? 64 : 128; + + Mat BT_data_packed; + Mat BT_data_packed_descales; + int ret = pack_B_wq_int8(B_data, B_data_quantize_scales, BT_data_packed, BT_data_packed_descales, constantN, constantK, block_size, opt.num_threads); + if (ret != 0) + return ret; + + BT_data_wq_int8 = BT_data_packed; + BT_data_wq_int8_descales = BT_data_packed_descales; + + B_data.release(); + B_data_quantize_scales.release(); + + return 0; +} + +int Gemm_arm::forward_wq_int8(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& opt) const +{ + const Mat& A = bottom_blobs[0]; + if (A.elemsize != 4u || A.elempack != 1) + { + NCNN_LOGE("Gemm unsupported input"); + return -1; + } + + if (transA && A.dims != 2) + { + NCNN_LOGE("Gemm unsupported input"); + return -1; + } + + const int K = transA ? A.h : A.w; + if (K != constantK) + { + NCNN_LOGE("Gemm weight block quantize K mismatch"); + return -1; + } + + const int M = transA ? A.w : A.dims == 3 ? A.c : A.h; + const int N = constantN; + + Mat C; + int broadcast_type_C = -1; + if (constantC) + { + C = C_data; + broadcast_type_C = constant_broadcast_type_C; + } + else + { + if (bottom_blobs.size() == 2) + C = bottom_blobs[1]; + + if (!C.empty()) + { + bool matched = false; + if (C.dims == 1 && C.w == 1) + { + broadcast_type_C = 0; + matched = true; + } + if (C.dims == 1 && C.w == M) + { + broadcast_type_C = 1; + matched = true; + } + if (C.dims == 1 && C.w == N) + { + broadcast_type_C = 4; + matched = true; + } + if (C.dims == 2 && C.w == 1 && C.h == M) + { + broadcast_type_C = 2; + matched = true; + } + if (C.dims == 2 && C.w == N && C.h == M) + { + broadcast_type_C = 3; + matched = true; + } + if (C.dims == 2 && C.w == N && C.h == 1) + { + broadcast_type_C = 4; + matched = true; + } + + if (!matched || C.elemsize != 4u || C.elempack != 1) + { + NCNN_LOGE("Gemm unsupported C"); + return -1; + } + } + } + + if (!C.empty() && (C.elemsize != 4u || C.elempack != 1)) + { + NCNN_LOGE("Gemm unsupported C"); + return -1; + } + + Mat& top_blob = top_blobs[0]; + if (output_transpose) + top_blob.create(M, N, (size_t)4u, opt.blob_allocator); + else + top_blob.create(N, M, (size_t)4u, opt.blob_allocator); + if (top_blob.empty()) + return -100; + + const int block_size_code = quantize_term % 10; + const int block_size = block_size_code == 0 ? 32 : block_size_code == 1 ? 64 : 128; + return gemm_BT_arm_wq_int8(A, BT_data_wq_int8, BT_data_wq_int8_descales, B_data_input_scales, C, top_blob, broadcast_type_C, N, K, block_size, transA, output_transpose, alpha, beta, constant_TILE_M, constant_TILE_N, constant_TILE_K, opt.num_threads, opt); +} +#endif // NCNN_WEIGHT_QUANT + int Gemm_arm::forward(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& opt) const { if (weight_block_quantize) { +#if NCNN_WEIGHT_QUANT + if (quantize_term / 100 == 8 && !BT_data_wq_int8.empty()) + return forward_wq_int8(bottom_blobs, top_blobs, opt); +#endif return Gemm::forward(bottom_blobs, top_blobs, opt); } diff --git a/src/layer/arm/gemm_arm.h b/src/layer/arm/gemm_arm.h index 7caf73e3876f..f5ab207dd465 100644 --- a/src/layer/arm/gemm_arm.h +++ b/src/layer/arm/gemm_arm.h @@ -14,6 +14,7 @@ class Gemm_arm : public Gemm Gemm_arm(); virtual int create_pipeline(const Option& opt); + virtual int destroy_pipeline(const Option& opt); virtual int forward(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& opt) const; @@ -34,6 +35,10 @@ class Gemm_arm : public Gemm int create_pipeline_int8(const Option& opt); int forward_int8(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& opt) const; #endif +#if NCNN_WEIGHT_QUANT + int create_pipeline_wq_int8(const Option& opt); + int forward_wq_int8(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& opt) const; +#endif public: int nT; @@ -42,6 +47,11 @@ class Gemm_arm : public Gemm Mat CT_data; int input_elemtype; // 0=auto 1=fp32 2=fp16 3=bf16 + +#if NCNN_WEIGHT_QUANT + Mat BT_data_wq_int8; + Mat BT_data_wq_int8_descales; +#endif }; } // namespace ncnn diff --git a/src/layer/arm/gemm_arm_asimddp.cpp b/src/layer/arm/gemm_arm_asimddp.cpp index 07b389bb4e96..240ac68d0a17 100644 --- a/src/layer/arm/gemm_arm_asimddp.cpp +++ b/src/layer/arm/gemm_arm_asimddp.cpp @@ -14,6 +14,32 @@ namespace ncnn { #include "gemm_int8_bf16s.h" #endif +#if NCNN_WEIGHT_QUANT +#include "gemm_wq_int8.h" +#endif + +#if NCNN_WEIGHT_QUANT +int pack_B_wq_int8_asimddp(const Mat& B, const Mat& B_scales, Mat& BT, Mat& BT_descales, int N, int K, int block_size, int num_threads) +{ + return pack_B_wq_int8(B, B_scales, BT, BT_descales, N, K, block_size, num_threads); +} + +void quantize_A_tile_wq_int8_asimddp(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr) +{ + quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); +} + +void transpose_quantize_A_tile_wq_int8_asimddp(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr) +{ + transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); +} + +void gemm_transB_packed_tile_wq_int8_asimddp(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int block_size) +{ + gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); +} +#endif // NCNN_WEIGHT_QUANT + void pack_A_tile_int8_asimddp(const Mat& A, Mat& AT, int i, int max_ii, int k, int max_kk) { pack_A_tile_int8(A, AT, i, max_ii, k, max_kk); diff --git a/src/layer/arm/gemm_arm_i8mm.cpp b/src/layer/arm/gemm_arm_i8mm.cpp index 596fef8a2489..882f6d57f5f1 100644 --- a/src/layer/arm/gemm_arm_i8mm.cpp +++ b/src/layer/arm/gemm_arm_i8mm.cpp @@ -14,6 +14,32 @@ namespace ncnn { #include "gemm_int8_bf16s.h" #endif +#if NCNN_WEIGHT_QUANT +#include "gemm_wq_int8.h" +#endif + +#if NCNN_WEIGHT_QUANT +int pack_B_wq_int8_i8mm(const Mat& B, const Mat& B_scales, Mat& BT, Mat& BT_descales, int N, int K, int block_size, int num_threads) +{ + return pack_B_wq_int8(B, B_scales, BT, BT_descales, N, K, block_size, num_threads); +} + +void quantize_A_tile_wq_int8_i8mm(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr) +{ + quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); +} + +void transpose_quantize_A_tile_wq_int8_i8mm(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr) +{ + transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); +} + +void gemm_transB_packed_tile_wq_int8_i8mm(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int block_size) +{ + gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); +} +#endif // NCNN_WEIGHT_QUANT + void pack_A_tile_int8_i8mm(const Mat& A, Mat& AT, int i, int max_ii, int k, int max_kk) { pack_A_tile_int8(A, AT, i, max_ii, k, max_kk); diff --git a/src/layer/arm/gemm_arm_svei8mm.cpp b/src/layer/arm/gemm_arm_svei8mm.cpp new file mode 100644 index 000000000000..73512ebca56e --- /dev/null +++ b/src/layer/arm/gemm_arm_svei8mm.cpp @@ -0,0 +1,26 @@ +// Copyright 2026 Tencent +// SPDX-License-Identifier: BSD-3-Clause + +#include + +#include "cpu.h" +#include "mat.h" +#include "arm_usability.h" + +namespace ncnn { + +#if NCNN_WEIGHT_QUANT +#include "gemm_wq_int8.h" + +int pack_B_wq_int8_svei8mm(const Mat& B, const Mat& B_scales, Mat& BT, Mat& BT_descales, int N, int K, int block_size, int num_threads) +{ + return pack_B_wq_int8(B, B_scales, BT, BT_descales, N, K, block_size, num_threads); +} + +void gemm_transB_packed_tile_wq_int8_svei8mm(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int block_size) +{ + gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); +} +#endif // NCNN_WEIGHT_QUANT + +} // namespace ncnn diff --git a/src/layer/arm/gemm_wq_int8.h b/src/layer/arm/gemm_wq_int8.h new file mode 100644 index 000000000000..3aa907278537 --- /dev/null +++ b/src/layer/arm/gemm_wq_int8.h @@ -0,0 +1,6135 @@ +// Copyright 2026 Tencent +// SPDX-License-Identifier: BSD-3-Clause + +#include + +#if NCNN_RUNTIME_CPU && NCNN_ARM86SVEI8MM && __aarch64__ && !__ARM_FEATURE_SVE_MATMUL_INT8 +int pack_B_wq_int8_svei8mm(const Mat& B, const Mat& B_scales, Mat& BT, Mat& BT_descales, int N, int K, int block_size, int num_threads); +void gemm_transB_packed_tile_wq_int8_svei8mm(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int block_size); +#endif + +#if NCNN_RUNTIME_CPU && NCNN_ARM84I8MM && __aarch64__ && !__ARM_FEATURE_MATMUL_INT8 +int pack_B_wq_int8_i8mm(const Mat& B, const Mat& B_scales, Mat& BT, Mat& BT_descales, int N, int K, int block_size, int num_threads); +void quantize_A_tile_wq_int8_i8mm(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr); +void transpose_quantize_A_tile_wq_int8_i8mm(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr); +void gemm_transB_packed_tile_wq_int8_i8mm(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int block_size); +#endif + +#if NCNN_RUNTIME_CPU && NCNN_ARM82DOT && __aarch64__ && !__ARM_FEATURE_DOTPROD && !__ARM_FEATURE_MATMUL_INT8 +int pack_B_wq_int8_asimddp(const Mat& B, const Mat& B_scales, Mat& BT, Mat& BT_descales, int N, int K, int block_size, int num_threads); +void quantize_A_tile_wq_int8_asimddp(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr); +void transpose_quantize_A_tile_wq_int8_asimddp(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr); +void gemm_transB_packed_tile_wq_int8_asimddp(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int block_size); +#endif + +static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr) +{ +#if NCNN_RUNTIME_CPU && NCNN_ARM84I8MM && __aarch64__ && !__ARM_FEATURE_MATMUL_INT8 + if (ncnn::cpu_support_arm_i8mm()) + { + quantize_A_tile_wq_int8_i8mm(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); + return; + } +#endif +#if NCNN_RUNTIME_CPU && NCNN_ARM82DOT && __aarch64__ && !__ARM_FEATURE_DOTPROD && !__ARM_FEATURE_MATMUL_INT8 + if (ncnn::cpu_support_arm_asimddp()) + { + quantize_A_tile_wq_int8_asimddp(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); + return; + } +#endif + + signed char* outptr = AT_tile; + const int out_hstep = AT_tile.w; + float* descales = AT_descales_tile; + const int descales_hstep = AT_descales_tile.w; + const int block_count = (K + block_size - 1) / block_size; + const size_t A_hstep = A.dims == 3 ? A.cstep : (size_t)A.w; + +#if __ARM_NEON && __aarch64__ + if (max_ii >= 8) + { + signed char* pp = AT_tile; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + float absmax[8]; + float scales[8]; + + for (int r = 0; r < 8; r++) + { + const float* ptrA = (const float*)A + (i + r) * A_hstep + k0; + float32x4_t _absmax = vdupq_n_f32(0.f); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + float32x4_t _v = vld1q_f32(ptrA + kk); + if (input_scale_ptr) + _v = vmulq_f32(_v, vld1q_f32(input_scale_ptr + k0 + kk)); + _absmax = vmaxq_f32(_absmax, vabsq_f32(_v)); + } +#if __aarch64__ + absmax[r] = vmaxvq_f32(_absmax); +#else + float32x2_t _max2 = vmax_f32(vget_low_f32(_absmax), vget_high_f32(_absmax)); + _max2 = vpmax_f32(_max2, _max2); + absmax[r] = vget_lane_f32(_max2, 0); +#endif + for (; kk < max_kk; kk++) + { + float v = ptrA[kk]; + if (input_scale_ptr) + v *= input_scale_ptr[k0 + kk]; + absmax[r] = std::max(absmax[r], fabsf(v)); + } + + descales[g * 8 + r] = absmax[r] / 127.f; + volatile double scale_fp64 = absmax[r] == 0.f ? 0.0 : 127.0 / (double)absmax[r]; + scales[r] = (float)scale_fp64; + } + + int kk = 0; +#if __ARM_FEATURE_MATMUL_INT8 + for (; kk + 7 < max_kk; kk += 8) + { + for (int r = 0; r < 8; r += 2) + { + const float* ptrA0 = (const float*)A + (i + r) * A_hstep + k0 + kk; + const float* ptrA1 = ptrA0 + A_hstep; + float32x4_t _v00 = vld1q_f32(ptrA0); + float32x4_t _v01 = vld1q_f32(ptrA0 + 4); + float32x4_t _v10 = vld1q_f32(ptrA1); + float32x4_t _v11 = vld1q_f32(ptrA1 + 4); + if (input_scale_ptr) + { + const float32x4_t _s0 = vld1q_f32(input_scale_ptr + k0 + kk); + const float32x4_t _s1 = vld1q_f32(input_scale_ptr + k0 + kk + 4); + _v00 = vmulq_f32(_v00, _s0); + _v01 = vmulq_f32(_v01, _s1); + _v10 = vmulq_f32(_v10, _s0); + _v11 = vmulq_f32(_v11, _s1); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+w"(_v00), "+w"(_v01), "+w"(_v10), "+w"(_v11)); +#endif + } + const float32x4_t _scale0 = vdupq_n_f32(scales[r]); + const float32x4_t _scale1 = vdupq_n_f32(scales[r + 1]); + const int8x8_t _q0 = float2int8(vmulq_f32(_v00, _scale0), vmulq_f32(_v01, _scale0)); + const int8x8_t _q1 = float2int8(vmulq_f32(_v10, _scale1), vmulq_f32(_v11, _scale1)); + vst1q_s8(pp, vcombine_s8(_q0, _q1)); + pp += 16; + } + } +#endif // __ARM_FEATURE_MATMUL_INT8 +#if __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk; kk += 4) + { + for (int r = 0; r < 8; r++) + { + const float* ptrA = (const float*)A + (i + r) * A_hstep + k0 + kk; + float32x4_t _v = vld1q_f32(ptrA); + if (input_scale_ptr) + { + _v = vmulq_f32(_v, vld1q_f32(input_scale_ptr + k0 + kk)); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+w"(_v)); +#endif + } + const float32x4_t _scale = vdupq_n_f32(scales[r]); + const int8x8_t _q = float2int8(vmulq_f32(_v, _scale), vmulq_f32(_v, _scale)); + vst1_lane_s32((int*)pp, vreinterpret_s32_s8(_q), 0); + pp += 4; + } + } +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 1 < max_kk; kk += 2) + { + for (int r = 0; r < 8; r++) + { + const float* ptrA = (const float*)A + (i + r) * A_hstep + k0 + kk; + float v0 = ptrA[0]; + float v1 = ptrA[1]; + if (input_scale_ptr) + { + v0 *= input_scale_ptr[k0 + kk]; + v1 *= input_scale_ptr[k0 + kk + 1]; +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+w"(v0), "+w"(v1)); +#endif + } + *pp++ = float2int8(v0 * scales[r]); + *pp++ = float2int8(v1 * scales[r]); + } + } + if (kk < max_kk) + { + for (int r = 0; r < 8; r++) + { + const float* ptrA = (const float*)A + (i + r) * A_hstep + k0 + kk; + float v = ptrA[0]; + if (input_scale_ptr) + { + v *= input_scale_ptr[k0 + kk]; +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+w"(v)); +#endif + } + *pp++ = float2int8(v * scales[r]); + } + } + } + + return; + } +#endif // __ARM_NEON && __aarch64__ + + int ii = 0; +#if __ARM_NEON + for (; ii + 3 < max_ii; ii += 4) + { + signed char* pp = outptr + ii * out_hstep; + float* descale_ptr = descales + ii * descales_hstep; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + float absmax[4]; + float scales[4]; + + for (int r = 0; r < 4; r++) + { + const float* ptrA = (const float*)A + (i + ii + r) * A_hstep + k0; + float32x4_t _absmax = vdupq_n_f32(0.f); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + float32x4_t _v = vld1q_f32(ptrA + kk); + if (input_scale_ptr) + _v = vmulq_f32(_v, vld1q_f32(input_scale_ptr + k0 + kk)); + _absmax = vmaxq_f32(_absmax, vabsq_f32(_v)); + } +#if __aarch64__ + absmax[r] = vmaxvq_f32(_absmax); +#else + float32x2_t _max2 = vmax_f32(vget_low_f32(_absmax), vget_high_f32(_absmax)); + _max2 = vpmax_f32(_max2, _max2); + absmax[r] = vget_lane_f32(_max2, 0); +#endif + for (; kk < max_kk; kk++) + { + float v = ptrA[kk]; + if (input_scale_ptr) + v *= input_scale_ptr[k0 + kk]; + absmax[r] = std::max(absmax[r], fabsf(v)); + } + descale_ptr[g * 4 + r] = absmax[r] / 127.f; + volatile double scale_fp64 = absmax[r] == 0.f ? 0.0 : 127.0 / (double)absmax[r]; + scales[r] = (float)scale_fp64; + } + + int kk = 0; +#if __ARM_FEATURE_MATMUL_INT8 + for (; kk + 7 < max_kk; kk += 8) + { + for (int r = 0; r < 4; r += 2) + { + const float* ptrA0 = (const float*)A + (i + ii + r) * A_hstep + k0 + kk; + const float* ptrA1 = ptrA0 + A_hstep; + float32x4_t _v00 = vld1q_f32(ptrA0); + float32x4_t _v01 = vld1q_f32(ptrA0 + 4); + float32x4_t _v10 = vld1q_f32(ptrA1); + float32x4_t _v11 = vld1q_f32(ptrA1 + 4); + if (input_scale_ptr) + { + const float32x4_t _s0 = vld1q_f32(input_scale_ptr + k0 + kk); + const float32x4_t _s1 = vld1q_f32(input_scale_ptr + k0 + kk + 4); + _v00 = vmulq_f32(_v00, _s0); + _v01 = vmulq_f32(_v01, _s1); + _v10 = vmulq_f32(_v10, _s0); + _v11 = vmulq_f32(_v11, _s1); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+w"(_v00), "+w"(_v01), "+w"(_v10), "+w"(_v11)); +#endif + } + const float32x4_t _scale0 = vdupq_n_f32(scales[r]); + const float32x4_t _scale1 = vdupq_n_f32(scales[r + 1]); + const int8x8_t _q0 = float2int8(vmulq_f32(_v00, _scale0), vmulq_f32(_v01, _scale0)); + const int8x8_t _q1 = float2int8(vmulq_f32(_v10, _scale1), vmulq_f32(_v11, _scale1)); + vst1q_s8(pp, vcombine_s8(_q0, _q1)); + pp += 16; + } + } +#endif // __ARM_FEATURE_MATMUL_INT8 +#if __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk; kk += 4) + { + for (int r = 0; r < 4; r++) + { + const float* ptrA = (const float*)A + (i + ii + r) * A_hstep + k0 + kk; + float32x4_t _v = vld1q_f32(ptrA); + if (input_scale_ptr) + { + _v = vmulq_f32(_v, vld1q_f32(input_scale_ptr + k0 + kk)); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+w"(_v)); +#endif + } + const int8x8_t _q = float2int8(vmulq_n_f32(_v, scales[r]), vmulq_n_f32(_v, scales[r])); + vst1_lane_s32((int*)pp, vreinterpret_s32_s8(_q), 0); + pp += 4; + } + } +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 1 < max_kk; kk += 2) + { + for (int r = 0; r < 4; r++) + { + const float* ptrA = (const float*)A + (i + ii + r) * A_hstep + k0 + kk; + float v0 = ptrA[0]; + float v1 = ptrA[1]; + if (input_scale_ptr) + { + v0 *= input_scale_ptr[k0 + kk]; + v1 *= input_scale_ptr[k0 + kk + 1]; +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+w"(v0), "+w"(v1)); +#endif + } + *pp++ = float2int8(v0 * scales[r]); + *pp++ = float2int8(v1 * scales[r]); + } + } + if (kk < max_kk) + { + for (int r = 0; r < 4; r++) + { + const float* ptrA = (const float*)A + (i + ii + r) * A_hstep + k0 + kk; + float v = ptrA[0]; + if (input_scale_ptr) + { + v *= input_scale_ptr[k0 + kk]; +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+w"(v)); +#endif + } + *pp++ = float2int8(v * scales[r]); + } + } + } + } + for (; ii + 1 < max_ii; ii += 2) + { + signed char* pp = outptr + ii * out_hstep; + float* descale_ptr = descales + ii * descales_hstep; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + float absmax[2]; + float scales[2]; + + for (int r = 0; r < 2; r++) + { + const float* ptrA = (const float*)A + (i + ii + r) * A_hstep + k0; + float32x4_t _absmax = vdupq_n_f32(0.f); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + float32x4_t _v = vld1q_f32(ptrA + kk); + if (input_scale_ptr) + _v = vmulq_f32(_v, vld1q_f32(input_scale_ptr + k0 + kk)); + _absmax = vmaxq_f32(_absmax, vabsq_f32(_v)); + } +#if __aarch64__ + absmax[r] = vmaxvq_f32(_absmax); +#else + float32x2_t _max2 = vmax_f32(vget_low_f32(_absmax), vget_high_f32(_absmax)); + _max2 = vpmax_f32(_max2, _max2); + absmax[r] = vget_lane_f32(_max2, 0); +#endif + for (; kk < max_kk; kk++) + { + float v = ptrA[kk]; + if (input_scale_ptr) + v *= input_scale_ptr[k0 + kk]; + absmax[r] = std::max(absmax[r], fabsf(v)); + } + descale_ptr[g * 2 + r] = absmax[r] / 127.f; + volatile double scale_fp64 = absmax[r] == 0.f ? 0.0 : 127.0 / (double)absmax[r]; + scales[r] = (float)scale_fp64; + } + + int kk = 0; +#if __ARM_FEATURE_MATMUL_INT8 + for (; kk + 7 < max_kk; kk += 8) + { + const float* ptrA0 = (const float*)A + (i + ii) * A_hstep + k0 + kk; + const float* ptrA1 = ptrA0 + A_hstep; + float32x4_t _v00 = vld1q_f32(ptrA0); + float32x4_t _v01 = vld1q_f32(ptrA0 + 4); + float32x4_t _v10 = vld1q_f32(ptrA1); + float32x4_t _v11 = vld1q_f32(ptrA1 + 4); + if (input_scale_ptr) + { + const float32x4_t _s0 = vld1q_f32(input_scale_ptr + k0 + kk); + const float32x4_t _s1 = vld1q_f32(input_scale_ptr + k0 + kk + 4); + _v00 = vmulq_f32(_v00, _s0); + _v01 = vmulq_f32(_v01, _s1); + _v10 = vmulq_f32(_v10, _s0); + _v11 = vmulq_f32(_v11, _s1); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+w"(_v00), "+w"(_v01), "+w"(_v10), "+w"(_v11)); +#endif + } + const int8x8_t _q0 = float2int8(vmulq_n_f32(_v00, scales[0]), vmulq_n_f32(_v01, scales[0])); + const int8x8_t _q1 = float2int8(vmulq_n_f32(_v10, scales[1]), vmulq_n_f32(_v11, scales[1])); + vst1q_s8(pp, vcombine_s8(_q0, _q1)); + pp += 16; + } +#endif // __ARM_FEATURE_MATMUL_INT8 +#if __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk; kk += 4) + { + int32x2_t _q01 = vdup_n_s32(0); + for (int r = 0; r < 2; r++) + { + const float* ptrA = (const float*)A + (i + ii + r) * A_hstep + k0 + kk; + float32x4_t _v = vld1q_f32(ptrA); + if (input_scale_ptr) + { + _v = vmulq_f32(_v, vld1q_f32(input_scale_ptr + k0 + kk)); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+w"(_v)); +#endif + } + const int8x8_t _q = float2int8(vmulq_n_f32(_v, scales[r]), vmulq_n_f32(_v, scales[r])); + _q01 = vset_lane_s32(vget_lane_s32(vreinterpret_s32_s8(_q), 0), _q01, r); + } + vst1_s8(pp, vreinterpret_s8_s32(_q01)); + pp += 8; + } +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 1 < max_kk; kk += 2) + { + for (int r = 0; r < 2; r++) + { + const float* ptrA = (const float*)A + (i + ii + r) * A_hstep + k0 + kk; + float v0 = ptrA[0]; + float v1 = ptrA[1]; + if (input_scale_ptr) + { + v0 *= input_scale_ptr[k0 + kk]; + v1 *= input_scale_ptr[k0 + kk + 1]; + } + *pp++ = float2int8(v0 * scales[r]); + *pp++ = float2int8(v1 * scales[r]); + } + } + if (kk < max_kk) + { + for (int r = 0; r < 2; r++) + { + const float* ptrA = (const float*)A + (i + ii + r) * A_hstep + k0 + kk; + float v = ptrA[0]; + if (input_scale_ptr) + v *= input_scale_ptr[k0 + kk]; + *pp++ = float2int8(v * scales[r]); + } + } + } + } +#elif __ARM_FEATURE_SIMD32 && NCNN_GNU_INLINE_ASM + for (; ii + 1 < max_ii; ii += 2) + { + signed char* pp = outptr + ii * out_hstep; + float* descale_ptr = descales + ii * descales_hstep; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + const float* ptrA0 = (const float*)A + (i + ii) * A_hstep + k0; + const float* ptrA1 = ptrA0 + A_hstep; + float absmax0 = 0.f; + float absmax1 = 0.f; + + for (int kk = 0; kk < max_kk; kk++) + { + float v0 = ptrA0[kk]; + float v1 = ptrA1[kk]; + if (input_scale_ptr) + { + v0 *= input_scale_ptr[k0 + kk]; + v1 *= input_scale_ptr[k0 + kk]; + } + absmax0 = std::max(absmax0, fabsf(v0)); + absmax1 = std::max(absmax1, fabsf(v1)); + } + + descale_ptr[g * 2] = absmax0 / 127.f; + descale_ptr[g * 2 + 1] = absmax1 / 127.f; + volatile double scale0_fp64 = absmax0 == 0.f ? 0.0 : 127.0 / (double)absmax0; + volatile double scale1_fp64 = absmax1 == 0.f ? 0.0 : 127.0 / (double)absmax1; + const float scale0 = (float)scale0_fp64; + const float scale1 = (float)scale1_fp64; + + for (int kk = 0; kk < max_kk; kk++) + { + float v0 = ptrA0[kk]; + float v1 = ptrA1[kk]; + if (input_scale_ptr) + { + v0 *= input_scale_ptr[k0 + kk]; + v1 *= input_scale_ptr[k0 + kk]; + asm volatile("" : "+w"(v0), "+w"(v1)); + } + *pp++ = float2int8(v0 * scale0); + *pp++ = float2int8(v1 * scale1); + } + } + } +#endif // __ARM_NEON + for (; ii < max_ii; ii++) + { + const float* ptrA = (const float*)A + (i + ii) * A_hstep; + signed char* outptr0 = outptr + ii * out_hstep; + float* descale_ptr = descales + ii * descales_hstep; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + + float absmax = 0.f; + int kk = 0; +#if __ARM_NEON + float32x4_t _absmax = vdupq_n_f32(0.f); + for (; kk + 3 < max_kk; kk += 4) + { + float32x4_t _v = vld1q_f32(ptrA + k0 + kk); + if (input_scale_ptr) + _v = vmulq_f32(_v, vld1q_f32(input_scale_ptr + k0 + kk)); + _absmax = vmaxq_f32(_absmax, vabsq_f32(_v)); + } +#if __aarch64__ + absmax = vmaxvq_f32(_absmax); +#else + float32x2_t _max2 = vmax_f32(vget_low_f32(_absmax), vget_high_f32(_absmax)); + _max2 = vpmax_f32(_max2, _max2); + absmax = vget_lane_f32(_max2, 0); +#endif +#endif // __ARM_NEON + for (; kk < max_kk; kk++) + { + const int k = k0 + kk; + float v = ptrA[k]; + if (input_scale_ptr) + v *= input_scale_ptr[k]; + absmax = std::max(absmax, fabsf(v)); + } + + if (absmax == 0.f) + { + descale_ptr[g] = 0.f; + for (int k = 0; k < max_kk; k++) + outptr0[k0 + k] = 0; + continue; + } + + volatile double scale_fp64 = 127.0 / (double)absmax; + const float scale = (float)scale_fp64; + descale_ptr[g] = absmax / 127.f; + + kk = 0; +#if __ARM_NEON + const float32x4_t _scale = vdupq_n_f32(scale); + for (; kk + 7 < max_kk; kk += 8) + { + float32x4_t _v0 = vld1q_f32(ptrA + k0 + kk); + float32x4_t _v1 = vld1q_f32(ptrA + k0 + kk + 4); + if (input_scale_ptr) + { + _v0 = vmulq_f32(_v0, vld1q_f32(input_scale_ptr + k0 + kk)); + _v1 = vmulq_f32(_v1, vld1q_f32(input_scale_ptr + k0 + kk + 4)); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+w"(_v0), "+w"(_v1)); +#else + volatile float32x4_t _v0_ordered = _v0; + volatile float32x4_t _v1_ordered = _v1; + _v0 = _v0_ordered; + _v1 = _v1_ordered; +#endif + } + vst1_s8(outptr0 + k0 + kk, float2int8(vmulq_f32(_v0, _scale), vmulq_f32(_v1, _scale))); + } + for (; kk + 3 < max_kk; kk += 4) + { + float32x4_t _v = vld1q_f32(ptrA + k0 + kk); + if (input_scale_ptr) + { + _v = vmulq_f32(_v, vld1q_f32(input_scale_ptr + k0 + kk)); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+w"(_v)); +#else + volatile float32x4_t _v_ordered = _v; + _v = _v_ordered; +#endif + } + const int8x8_t _q = float2int8(vmulq_f32(_v, _scale), vmulq_f32(_v, _scale)); + vst1_lane_s32((int*)(outptr0 + k0 + kk), vreinterpret_s32_s8(_q), 0); + } +#endif // __ARM_NEON + for (; kk < max_kk; kk++) + { + const int k = k0 + kk; + float v = ptrA[k]; + if (input_scale_ptr) + { + v *= input_scale_ptr[k]; +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+w"(v)); +#else + volatile float v_ordered = v; + v = v_ordered; +#endif + } + outptr0[k] = float2int8(v * scale); + } + } + } +} + +static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr) +{ +#if NCNN_RUNTIME_CPU && NCNN_ARM84I8MM && __aarch64__ && !__ARM_FEATURE_MATMUL_INT8 + if (ncnn::cpu_support_arm_i8mm()) + { + transpose_quantize_A_tile_wq_int8_i8mm(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); + return; + } +#endif +#if NCNN_RUNTIME_CPU && NCNN_ARM82DOT && __aarch64__ && !__ARM_FEATURE_DOTPROD && !__ARM_FEATURE_MATMUL_INT8 + if (ncnn::cpu_support_arm_asimddp()) + { + transpose_quantize_A_tile_wq_int8_asimddp(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); + return; + } +#endif + + signed char* outptr = AT_tile; + const int out_hstep = AT_tile.w; + float* descales = AT_descales_tile; + const int descales_hstep = AT_descales_tile.w; + const int block_count = (K + block_size - 1) / block_size; + const size_t A_hstep = A.dims == 3 ? A.cstep : (size_t)A.w; + + int ii = 0; +#if __ARM_NEON +#if __aarch64__ + for (; ii + 7 < max_ii; ii += 8) + { + const float* ptrA = (const float*)A + i + ii; + signed char* pp = outptr + ii * out_hstep; + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + float32x4_t _absmax0 = vdupq_n_f32(0.f); + float32x4_t _absmax1 = vdupq_n_f32(0.f); + for (int kk = 0; kk < max_kk; kk++) + { + const int k = k0 + kk; + float32x4_t _v0 = vld1q_f32(ptrA + (size_t)k * A_hstep); + float32x4_t _v1 = vld1q_f32(ptrA + (size_t)k * A_hstep + 4); + if (input_scale_ptr) + { + _v0 = vmulq_n_f32(_v0, input_scale_ptr[k]); + _v1 = vmulq_n_f32(_v1, input_scale_ptr[k]); + } + _absmax0 = vmaxq_f32(_absmax0, vabsq_f32(_v0)); + _absmax1 = vmaxq_f32(_absmax1, vabsq_f32(_v1)); + } + + float absmax[8]; + vst1q_f32(absmax, _absmax0); + vst1q_f32(absmax + 4, _absmax1); + descales[g * 8] = absmax[0] / 127.f; + descales[g * 8 + 1] = absmax[1] / 127.f; + descales[g * 8 + 2] = absmax[2] / 127.f; + descales[g * 8 + 3] = absmax[3] / 127.f; + descales[g * 8 + 4] = absmax[4] / 127.f; + descales[g * 8 + 5] = absmax[5] / 127.f; + descales[g * 8 + 6] = absmax[6] / 127.f; + descales[g * 8 + 7] = absmax[7] / 127.f; + + float scales[8]; + for (int r = 0; r < 8; r++) + { + volatile double scale_fp64 = absmax[r] == 0.f ? 0.0 : 127.0 / (double)absmax[r]; + scales[r] = (float)scale_fp64; + } + const float32x4_t _scale0 = vld1q_f32(scales); + const float32x4_t _scale1 = vld1q_f32(scales + 4); + + int kk = 0; +#if __ARM_FEATURE_MATMUL_INT8 + for (; kk + 7 < max_kk; kk += 8) + { + int8x8_t _q[8]; + for (int t = 0; t < 8; t++) + { + const int k = k0 + kk + t; + float32x4_t _v0 = vld1q_f32(ptrA + (size_t)k * A_hstep); + float32x4_t _v1 = vld1q_f32(ptrA + (size_t)k * A_hstep + 4); + if (input_scale_ptr) + { + _v0 = vmulq_n_f32(_v0, input_scale_ptr[k]); + _v1 = vmulq_n_f32(_v1, input_scale_ptr[k]); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+w"(_v0), "+w"(_v1)); +#endif + } + _q[t] = float2int8(vmulq_f32(_v0, _scale0), vmulq_f32(_v1, _scale1)); + } + int8x8x2_t _r04 = vzip_s8(_q[0], _q[4]); + int8x8x2_t _r15 = vzip_s8(_q[1], _q[5]); + int8x8x2_t _r26 = vzip_s8(_q[2], _q[6]); + int8x8x2_t _r37 = vzip_s8(_q[3], _q[7]); + int8x8x4_t _r0123; + _r0123.val[0] = _r04.val[0]; + _r0123.val[1] = _r15.val[0]; + _r0123.val[2] = _r26.val[0]; + _r0123.val[3] = _r37.val[0]; + int8x8x4_t _r4567; + _r4567.val[0] = _r04.val[1]; + _r4567.val[1] = _r15.val[1]; + _r4567.val[2] = _r26.val[1]; + _r4567.val[3] = _r37.val[1]; + vst4_s8(pp, _r0123); + vst4_s8(pp + 32, _r4567); + pp += 64; + } +#endif // __ARM_FEATURE_MATMUL_INT8 +#if __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk; kk += 4) + { + int8x8x4_t _q; + for (int t = 0; t < 4; t++) + { + const int k = k0 + kk + t; + float32x4_t _v0 = vld1q_f32(ptrA + (size_t)k * A_hstep); + float32x4_t _v1 = vld1q_f32(ptrA + (size_t)k * A_hstep + 4); + if (input_scale_ptr) + { + _v0 = vmulq_n_f32(_v0, input_scale_ptr[k]); + _v1 = vmulq_n_f32(_v1, input_scale_ptr[k]); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+w"(_v0), "+w"(_v1)); +#endif + } + _q.val[t] = float2int8(vmulq_f32(_v0, _scale0), vmulq_f32(_v1, _scale1)); + } + vst4_s8(pp, _q); + pp += 32; + } +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 1 < max_kk; kk += 2) + { + int8x8x2_t _q; + for (int t = 0; t < 2; t++) + { + const int k = k0 + kk + t; + float32x4_t _v0 = vld1q_f32(ptrA + (size_t)k * A_hstep); + float32x4_t _v1 = vld1q_f32(ptrA + (size_t)k * A_hstep + 4); + if (input_scale_ptr) + { + _v0 = vmulq_n_f32(_v0, input_scale_ptr[k]); + _v1 = vmulq_n_f32(_v1, input_scale_ptr[k]); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+w"(_v0), "+w"(_v1)); +#endif + } + _q.val[t] = float2int8(vmulq_f32(_v0, _scale0), vmulq_f32(_v1, _scale1)); + } + vst2_s8(pp, _q); + pp += 16; + } + if (kk < max_kk) + { + const int k = k0 + kk; + float32x4_t _v0 = vld1q_f32(ptrA + (size_t)k * A_hstep); + float32x4_t _v1 = vld1q_f32(ptrA + (size_t)k * A_hstep + 4); + if (input_scale_ptr) + { + _v0 = vmulq_n_f32(_v0, input_scale_ptr[k]); + _v1 = vmulq_n_f32(_v1, input_scale_ptr[k]); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+w"(_v0), "+w"(_v1)); +#endif + } + vst1_s8(pp, float2int8(vmulq_f32(_v0, _scale0), vmulq_f32(_v1, _scale1))); + pp += 8; + } + } + } +#endif // __aarch64__ + for (; ii + 3 < max_ii; ii += 4) + { + const float* ptrA = (const float*)A + i + ii; + signed char* pp = outptr + ii * out_hstep; + float* descale_ptr = descales + ii * descales_hstep; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + + float32x4_t _absmax = vdupq_n_f32(0.f); + for (int kk = 0; kk < max_kk; kk++) + { + const int k = k0 + kk; + float32x4_t _v = vld1q_f32(ptrA + (size_t)k * A_hstep); + if (input_scale_ptr) + _v = vmulq_f32(_v, vdupq_n_f32(input_scale_ptr[k])); + _absmax = vmaxq_f32(_absmax, vabsq_f32(_v)); + } + + float absmax[4]; + vst1q_f32(absmax, _absmax); + vst1q_f32(descale_ptr + g * 4, vmulq_n_f32(_absmax, 1.f / 127.f)); + + volatile double scale0_fp64 = absmax[0] == 0.f ? 0.0 : 127.0 / (double)absmax[0]; + volatile double scale1_fp64 = absmax[1] == 0.f ? 0.0 : 127.0 / (double)absmax[1]; + volatile double scale2_fp64 = absmax[2] == 0.f ? 0.0 : 127.0 / (double)absmax[2]; + volatile double scale3_fp64 = absmax[3] == 0.f ? 0.0 : 127.0 / (double)absmax[3]; + const float scales[4] = { + (float)scale0_fp64, + (float)scale1_fp64, + (float)scale2_fp64, + (float)scale3_fp64 + }; + const float32x4_t _scale = vld1q_f32(scales); + + int kk = 0; +#if __ARM_FEATURE_MATMUL_INT8 + for (; kk + 7 < max_kk; kk += 8) + { + int8x8_t _q[8]; + for (int t = 0; t < 8; t++) + { + const int k = k0 + kk + t; + float32x4_t _v = vld1q_f32(ptrA + (size_t)k * A_hstep); + if (input_scale_ptr) + { + _v = vmulq_n_f32(_v, input_scale_ptr[k]); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+w"(_v)); +#else + volatile float32x4_t _v_ordered = _v; + _v = _v_ordered; +#endif + } + _q[t] = float2int8(vmulq_f32(_v, _scale), vmulq_f32(_v, _scale)); + } + int8x8x2_t _r04 = vzip_s8(_q[0], _q[4]); + int8x8x2_t _r15 = vzip_s8(_q[1], _q[5]); + int8x8x2_t _r26 = vzip_s8(_q[2], _q[6]); + int8x8x2_t _r37 = vzip_s8(_q[3], _q[7]); + int8x8x4_t _r0123; + _r0123.val[0] = _r04.val[0]; + _r0123.val[1] = _r15.val[0]; + _r0123.val[2] = _r26.val[0]; + _r0123.val[3] = _r37.val[0]; + vst4_s8(pp, _r0123); + pp += 32; + } +#endif // __ARM_FEATURE_MATMUL_INT8 +#if __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk; kk += 4) + { + int8x8x4_t _q; + for (int t = 0; t < 4; t++) + { + const int k = k0 + kk + t; + float32x4_t _v = vld1q_f32(ptrA + (size_t)k * A_hstep); + if (input_scale_ptr) + { + _v = vmulq_n_f32(_v, input_scale_ptr[k]); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+w"(_v)); +#else + volatile float32x4_t _v_ordered = _v; + _v = _v_ordered; +#endif + } + _q.val[t] = float2int8(vmulq_f32(_v, _scale), vmulq_f32(_v, _scale)); + } + vst4_lane_s8(pp, _q, 0); + vst4_lane_s8(pp + 4, _q, 1); + vst4_lane_s8(pp + 8, _q, 2); + vst4_lane_s8(pp + 12, _q, 3); + pp += 16; + } +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 1 < max_kk; kk += 2) + { + int8x8x2_t _q; + for (int t = 0; t < 2; t++) + { + const int k = k0 + kk + t; + float32x4_t _v = vld1q_f32(ptrA + (size_t)k * A_hstep); + if (input_scale_ptr) + { + _v = vmulq_n_f32(_v, input_scale_ptr[k]); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+w"(_v)); +#else + volatile float32x4_t _v_ordered = _v; + _v = _v_ordered; +#endif + } + _q.val[t] = float2int8(vmulq_f32(_v, _scale), vmulq_f32(_v, _scale)); + } + vst2_lane_s8(pp, _q, 0); + vst2_lane_s8(pp + 2, _q, 1); + vst2_lane_s8(pp + 4, _q, 2); + vst2_lane_s8(pp + 6, _q, 3); + pp += 8; + } + if (kk < max_kk) + { + const int k = k0 + kk; + float32x4_t _v = vld1q_f32(ptrA + (size_t)k * A_hstep); + if (input_scale_ptr) + { + _v = vmulq_n_f32(_v, input_scale_ptr[k]); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+w"(_v)); +#else + volatile float32x4_t _v_ordered = _v; + _v = _v_ordered; +#endif + } + vst1_lane_s32((int*)pp, vreinterpret_s32_s8(float2int8(vmulq_f32(_v, _scale), vmulq_f32(_v, _scale))), 0); + pp += 4; + } + } + } + for (; ii + 1 < max_ii; ii += 2) + { + const float* ptrA = (const float*)A + i + ii; + signed char* pp = outptr + ii * out_hstep; + float* descale_ptr = descales + ii * descales_hstep; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + float32x2_t _absmax = vdup_n_f32(0.f); + for (int kk = 0; kk < max_kk; kk++) + { + const int k = k0 + kk; + float32x2_t _v = vld1_f32(ptrA + (size_t)k * A_hstep); + if (input_scale_ptr) + _v = vmul_n_f32(_v, input_scale_ptr[k]); + _absmax = vmax_f32(_absmax, vabs_f32(_v)); + } + + vst1_f32(descale_ptr + g * 2, vmul_n_f32(_absmax, 1.f / 127.f)); + float absmax[2]; + vst1_f32(absmax, _absmax); + float scales[2]; + for (int r = 0; r < 2; r++) + { + volatile double scale_fp64 = absmax[r] == 0.f ? 0.0 : 127.0 / (double)absmax[r]; + scales[r] = (float)scale_fp64; + } + const float32x2_t _scale = vld1_f32(scales); + + int kk = 0; +#if __ARM_FEATURE_MATMUL_INT8 + for (; kk + 7 < max_kk; kk += 8) + { + int8x8x4_t _q0; + int8x8x4_t _q1; + for (int t = 0; t < 4; t++) + { + const int k0t = k0 + kk + t; + const int k1t = k0 + kk + 4 + t; + float32x2_t _v0 = vld1_f32(ptrA + (size_t)k0t * A_hstep); + float32x2_t _v1 = vld1_f32(ptrA + (size_t)k1t * A_hstep); + if (input_scale_ptr) + { + _v0 = vmul_n_f32(_v0, input_scale_ptr[k0t]); + _v1 = vmul_n_f32(_v1, input_scale_ptr[k1t]); + } + const float32x4_t _s = vcombine_f32(_scale, _scale); + const float32x4_t _v0q = vmulq_f32(vcombine_f32(_v0, _v0), _s); + const float32x4_t _v1q = vmulq_f32(vcombine_f32(_v1, _v1), _s); + _q0.val[t] = float2int8(_v0q, _v0q); + _q1.val[t] = float2int8(_v1q, _v1q); + } + vst4_lane_s8(pp, _q0, 0); + vst4_lane_s8(pp + 4, _q1, 0); + vst4_lane_s8(pp + 8, _q0, 1); + vst4_lane_s8(pp + 12, _q1, 1); + pp += 16; + } +#endif // __ARM_FEATURE_MATMUL_INT8 +#if __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk; kk += 4) + { + int8x8x4_t _q; + for (int t = 0; t < 4; t++) + { + const int k = k0 + kk + t; + float32x2_t _v = vld1_f32(ptrA + (size_t)k * A_hstep); + if (input_scale_ptr) + _v = vmul_n_f32(_v, input_scale_ptr[k]); + const float32x4_t _vq = vmulq_f32(vcombine_f32(_v, _v), vcombine_f32(_scale, _scale)); + _q.val[t] = float2int8(_vq, _vq); + } + vst4_lane_s8(pp, _q, 0); + vst4_lane_s8(pp + 4, _q, 1); + pp += 8; + } +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 1 < max_kk; kk += 2) + { + int8x8x2_t _q; + for (int t = 0; t < 2; t++) + { + const int k = k0 + kk + t; + float32x2_t _v = vld1_f32(ptrA + (size_t)k * A_hstep); + if (input_scale_ptr) + _v = vmul_n_f32(_v, input_scale_ptr[k]); + const float32x4_t _vq = vmulq_f32(vcombine_f32(_v, _v), vcombine_f32(_scale, _scale)); + _q.val[t] = float2int8(_vq, _vq); + } + vst2_lane_s8(pp, _q, 0); + vst2_lane_s8(pp + 2, _q, 1); + pp += 4; + } + if (kk < max_kk) + { + const int k = k0 + kk; + float32x2_t _v = vld1_f32(ptrA + (size_t)k * A_hstep); + if (input_scale_ptr) + _v = vmul_n_f32(_v, input_scale_ptr[k]); + const float32x4_t _vq = vmulq_f32(vcombine_f32(_v, _v), vcombine_f32(_scale, _scale)); + vst1_lane_s16((short*)pp, vreinterpret_s16_s8(float2int8(_vq, _vq)), 0); + pp += 2; + } + } + } +#elif __ARM_FEATURE_SIMD32 && NCNN_GNU_INLINE_ASM + for (; ii + 1 < max_ii; ii += 2) + { + const float* ptrA = (const float*)A + i + ii; + signed char* pp = outptr + ii * out_hstep; + float* descale_ptr = descales + ii * descales_hstep; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + const float* ptrAk = ptrA + (size_t)k0 * A_hstep; + float absmax0 = 0.f; + float absmax1 = 0.f; + + for (int kk = 0; kk < max_kk; kk++) + { + float v0 = ptrAk[0]; + float v1 = ptrAk[1]; + if (input_scale_ptr) + { + v0 *= input_scale_ptr[k0 + kk]; + v1 *= input_scale_ptr[k0 + kk]; + } + absmax0 = std::max(absmax0, fabsf(v0)); + absmax1 = std::max(absmax1, fabsf(v1)); + ptrAk += A_hstep; + } + + descale_ptr[g * 2] = absmax0 / 127.f; + descale_ptr[g * 2 + 1] = absmax1 / 127.f; + volatile double scale0_fp64 = absmax0 == 0.f ? 0.0 : 127.0 / (double)absmax0; + volatile double scale1_fp64 = absmax1 == 0.f ? 0.0 : 127.0 / (double)absmax1; + const float scale0 = (float)scale0_fp64; + const float scale1 = (float)scale1_fp64; + + ptrAk = ptrA + (size_t)k0 * A_hstep; + for (int kk = 0; kk < max_kk; kk++) + { + float v0 = ptrAk[0]; + float v1 = ptrAk[1]; + if (input_scale_ptr) + { + v0 *= input_scale_ptr[k0 + kk]; + v1 *= input_scale_ptr[k0 + kk]; + asm volatile("" : "+w"(v0), "+w"(v1)); + } + *pp++ = float2int8(v0 * scale0); + *pp++ = float2int8(v1 * scale1); + ptrAk += A_hstep; + } + } + } +#endif // __ARM_NEON + for (; ii < max_ii; ii++) + { + const float* ptrA = (const float*)A + i + ii; + signed char* outptr0 = outptr + ii * out_hstep; + float* descale_ptr = descales + ii * descales_hstep; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + + float absmax = 0.f; + for (int kk = 0; kk < max_kk; kk++) + { + const int k = k0 + kk; + float v = ptrA[(size_t)k * A_hstep]; + if (input_scale_ptr) + v *= input_scale_ptr[k]; + absmax = std::max(absmax, fabsf(v)); + } + + if (absmax == 0.f) + { + descale_ptr[g] = 0.f; + for (int k = 0; k < max_kk; k++) + outptr0[k0 + k] = 0; + continue; + } + + volatile double scale_fp64 = 127.0 / (double)absmax; + const float scale = (float)scale_fp64; + descale_ptr[g] = absmax / 127.f; + + for (int kk = 0; kk < max_kk; kk++) + { + const int k = k0 + kk; + float v = ptrA[(size_t)k * A_hstep]; + if (input_scale_ptr) + { + v *= input_scale_ptr[k]; +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+w"(v)); +#else + volatile float v_ordered = v; + v = v_ordered; +#endif + } + outptr0[k] = float2int8(v * scale); + } + } + } +} + +// Persistent B uses the baseline gemm_int8 K2/K1 byte order inside each +// (output-column panel, quantization group). Every panel is exact and the +// address of the panel starting at output column j is always j * K. +// Two consecutive nr4 panels form the logical nr8 panel on aarch64. +static int pack_B_wq_int8(const Mat& B, const Mat& B_scales, Mat& BT, Mat& BT_descales, int N, int K, int block_size, int num_threads) +{ +#if NCNN_RUNTIME_CPU && NCNN_ARM86SVEI8MM && __aarch64__ && !__ARM_FEATURE_SVE_MATMUL_INT8 + if (ncnn::cpu_support_arm_svei8mm()) + return pack_B_wq_int8_svei8mm(B, B_scales, BT, BT_descales, N, K, block_size, num_threads); +#endif +#if NCNN_RUNTIME_CPU && NCNN_ARM84I8MM && __aarch64__ && !__ARM_FEATURE_MATMUL_INT8 + if (ncnn::cpu_support_arm_i8mm()) + return pack_B_wq_int8_i8mm(B, B_scales, BT, BT_descales, N, K, block_size, num_threads); +#endif +#if NCNN_RUNTIME_CPU && NCNN_ARM82DOT && __aarch64__ && !__ARM_FEATURE_DOTPROD && !__ARM_FEATURE_MATMUL_INT8 + if (ncnn::cpu_support_arm_asimddp()) + return pack_B_wq_int8_asimddp(B, B_scales, BT, BT_descales, N, K, block_size, num_threads); +#endif + + const int block_count = (K + block_size - 1) / block_size; + Mat BT_packed(N * K, (size_t)1u); + Mat BT_packed_descales(N * block_count, (size_t)4u); + if (BT_packed.empty() || BT_packed_descales.empty()) + return -100; + BT_packed.cstep = (size_t)N * K; + BT_packed_descales.cstep = (size_t)N * block_count; + + int panel_start = 0; +#if __ARM_NEON + const int nn4 = (N - panel_start) / 4; + const int panel_start4 = panel_start; + panel_start += nn4 * 4; +#endif + const int nn2 = (N - panel_start) / 2; + const int panel_start2 = panel_start; + panel_start += nn2 * 2; + const int nn1 = N - panel_start; + const int panel_start1 = panel_start; + + #pragma omp parallel num_threads(num_threads) + { +#if __ARM_NEON + #pragma omp for + for (int p = 0; p < nn4; p++) + { + const int j = panel_start4 + p * 4; + signed char* pp = (signed char*)BT_packed + j * K; + float* pd = (float*)BT_packed_descales + j * block_count; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + const signed char* p0 = B.row(j) + k0; + const signed char* p1 = B.row(j + 1) + k0; + const signed char* p2 = B.row(j + 2) + k0; + const signed char* p3 = B.row(j + 3) + k0; + int kk = 0; + for (; kk + 15 < max_kk; kk += 16) + { + const int8x16_t _p0 = vld1q_s8(p0); + const int8x16_t _p1 = vld1q_s8(p1); + const int8x16_t _p2 = vld1q_s8(p2); + const int8x16_t _p3 = vld1q_s8(p3); +#if __ARM_FEATURE_DOTPROD +#if __ARM_FEATURE_MATMUL_INT8 + int64x2x4_t _r0123; + _r0123.val[0] = vreinterpretq_s64_s8(_p0); + _r0123.val[1] = vreinterpretq_s64_s8(_p1); + _r0123.val[2] = vreinterpretq_s64_s8(_p2); + _r0123.val[3] = vreinterpretq_s64_s8(_p3); + vst4q_s64((int64_t*)pp, _r0123); +#else // __ARM_FEATURE_MATMUL_INT8 + int32x4x4_t _r0123; + _r0123.val[0] = vreinterpretq_s32_s8(_p0); + _r0123.val[1] = vreinterpretq_s32_s8(_p1); + _r0123.val[2] = vreinterpretq_s32_s8(_p2); + _r0123.val[3] = vreinterpretq_s32_s8(_p3); + vst4q_s32((int*)pp, _r0123); +#endif // __ARM_FEATURE_MATMUL_INT8 +#else // __ARM_FEATURE_DOTPROD + int16x8x4_t _r0123; + _r0123.val[0] = vreinterpretq_s16_s8(_p0); + _r0123.val[1] = vreinterpretq_s16_s8(_p1); + _r0123.val[2] = vreinterpretq_s16_s8(_p2); + _r0123.val[3] = vreinterpretq_s16_s8(_p3); + vst4q_s16((short*)pp, _r0123); +#endif // __ARM_FEATURE_DOTPROD + pp += 64; + p0 += 16; + p1 += 16; + p2 += 16; + p3 += 16; + } + for (; kk + 7 < max_kk; kk += 8) + { + const int8x8_t _p0 = vld1_s8(p0); + const int8x8_t _p1 = vld1_s8(p1); + const int8x8_t _p2 = vld1_s8(p2); + const int8x8_t _p3 = vld1_s8(p3); +#if __ARM_FEATURE_DOTPROD +#if __ARM_FEATURE_MATMUL_INT8 + vst1q_s8(pp, vcombine_s8(_p0, _p1)); + vst1q_s8(pp + 16, vcombine_s8(_p2, _p3)); +#else // __ARM_FEATURE_MATMUL_INT8 + int32x2x4_t _r0123; + _r0123.val[0] = vreinterpret_s32_s8(_p0); + _r0123.val[1] = vreinterpret_s32_s8(_p1); + _r0123.val[2] = vreinterpret_s32_s8(_p2); + _r0123.val[3] = vreinterpret_s32_s8(_p3); + vst4_s32((int*)pp, _r0123); +#endif // __ARM_FEATURE_MATMUL_INT8 +#else // __ARM_FEATURE_DOTPROD + int16x4x4_t _r0123; + _r0123.val[0] = vreinterpret_s16_s8(_p0); + _r0123.val[1] = vreinterpret_s16_s8(_p1); + _r0123.val[2] = vreinterpret_s16_s8(_p2); + _r0123.val[3] = vreinterpret_s16_s8(_p3); + vst4_s16((short*)pp, _r0123); +#endif // __ARM_FEATURE_DOTPROD + pp += 32; + p0 += 8; + p1 += 8; + p2 += 8; + p3 += 8; + } +#if __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk; kk += 4) + { + pp[0] = p0[0]; + pp[1] = p0[1]; + pp[2] = p0[2]; + pp[3] = p0[3]; + pp[4] = p1[0]; + pp[5] = p1[1]; + pp[6] = p1[2]; + pp[7] = p1[3]; + pp[8] = p2[0]; + pp[9] = p2[1]; + pp[10] = p2[2]; + pp[11] = p2[3]; + pp[12] = p3[0]; + pp[13] = p3[1]; + pp[14] = p3[2]; + pp[15] = p3[3]; + pp += 16; + p0 += 4; + p1 += 4; + p2 += 4; + p3 += 4; + } +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 1 < max_kk; kk += 2) + { + pp[0] = p0[0]; + pp[1] = p0[1]; + pp[2] = p1[0]; + pp[3] = p1[1]; + pp[4] = p2[0]; + pp[5] = p2[1]; + pp[6] = p3[0]; + pp[7] = p3[1]; + pp += 8; + p0 += 2; + p1 += 2; + p2 += 2; + p3 += 2; + } + if (kk < max_kk) + { + pp[0] = p0[0]; + pp[1] = p1[0]; + pp[2] = p2[0]; + pp[3] = p3[0]; + pp += 4; + } + + for (int jj = 0; jj < 4; jj++) + pd[g * 4 + jj] = 1.f / B_scales.row(j + jj)[g]; + } + } +#endif + #pragma omp for + for (int p = 0; p < nn2; p++) + { + const int j = panel_start2 + p * 2; + signed char* pp = (signed char*)BT_packed + j * K; + float* pd = (float*)BT_packed_descales + j * block_count; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + const signed char* p0 = B.row(j) + k0; + const signed char* p1 = B.row(j + 1) + k0; + int kk = 0; +#if __ARM_NEON + for (; kk + 15 < max_kk; kk += 16) + { + const int8x16_t _p0 = vld1q_s8(p0); + const int8x16_t _p1 = vld1q_s8(p1); +#if __ARM_FEATURE_DOTPROD +#if __ARM_FEATURE_MATMUL_INT8 + int64x2x2_t _r01; + _r01.val[0] = vreinterpretq_s64_s8(_p0); + _r01.val[1] = vreinterpretq_s64_s8(_p1); + vst2q_s64((int64_t*)pp, _r01); +#else // __ARM_FEATURE_MATMUL_INT8 + int32x4x2_t _r01; + _r01.val[0] = vreinterpretq_s32_s8(_p0); + _r01.val[1] = vreinterpretq_s32_s8(_p1); + vst2q_s32((int*)pp, _r01); +#endif // __ARM_FEATURE_MATMUL_INT8 +#else // __ARM_FEATURE_DOTPROD + int16x8x2_t _r01; + _r01.val[0] = vreinterpretq_s16_s8(_p0); + _r01.val[1] = vreinterpretq_s16_s8(_p1); + vst2q_s16((short*)pp, _r01); +#endif // __ARM_FEATURE_DOTPROD + pp += 32; + p0 += 16; + p1 += 16; + } + for (; kk + 7 < max_kk; kk += 8) + { + const int8x8_t _p0 = vld1_s8(p0); + const int8x8_t _p1 = vld1_s8(p1); +#if __ARM_FEATURE_DOTPROD +#if __ARM_FEATURE_MATMUL_INT8 + vst1q_s8(pp, vcombine_s8(_p0, _p1)); +#else // __ARM_FEATURE_MATMUL_INT8 + int32x2x2_t _r01; + _r01.val[0] = vreinterpret_s32_s8(_p0); + _r01.val[1] = vreinterpret_s32_s8(_p1); + vst2_s32((int*)pp, _r01); +#endif // __ARM_FEATURE_MATMUL_INT8 +#else // __ARM_FEATURE_DOTPROD + int16x4x2_t _r01; + _r01.val[0] = vreinterpret_s16_s8(_p0); + _r01.val[1] = vreinterpret_s16_s8(_p1); + vst2_s16((short*)pp, _r01); +#endif // __ARM_FEATURE_DOTPROD + pp += 16; + p0 += 8; + p1 += 8; + } +#endif // __ARM_NEON +#if __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk; kk += 4) + { + pp[0] = p0[0]; + pp[1] = p0[1]; + pp[2] = p0[2]; + pp[3] = p0[3]; + pp[4] = p1[0]; + pp[5] = p1[1]; + pp[6] = p1[2]; + pp[7] = p1[3]; + pp += 8; + p0 += 4; + p1 += 4; + } +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 1 < max_kk; kk += 2) + { +#if !__ARM_NEON && __ARM_FEATURE_SIMD32 && NCNN_GNU_INLINE_ASM + pp[0] = p0[0]; + pp[1] = p1[0]; + pp[2] = p0[1]; + pp[3] = p1[1]; + pp += 4; +#else + pp[0] = p0[0]; + pp[1] = p0[1]; + pp[2] = p1[0]; + pp[3] = p1[1]; + pp += 4; +#endif + p0 += 2; + p1 += 2; + } + if (kk < max_kk) + { + pp[0] = p0[0]; + pp[1] = p1[0]; + pp += 2; + } + + for (int jj = 0; jj < 2; jj++) + pd[g * 2 + jj] = 1.f / B_scales.row(j + jj)[g]; + } + } + #pragma omp for + for (int p = 0; p < nn1; p++) + { + const int j = panel_start1 + p * 1; + signed char* pp = (signed char*)BT_packed + j * K; + float* pd = (float*)BT_packed_descales + j * block_count; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + const signed char* p0 = B.row(j) + k0; + int kk = 0; +#if __ARM_NEON + for (; kk + 15 < max_kk; kk += 16) + { + vst1q_s8(pp, vld1q_s8(p0)); + pp += 16; + p0 += 16; + } + for (; kk + 7 < max_kk; kk += 8) + { + vst1_s8(pp, vld1_s8(p0)); + pp += 8; + p0 += 8; + } +#endif // __ARM_NEON +#if __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk; kk += 4) + { + pp[0] = p0[0]; + pp[1] = p0[1]; + pp[2] = p0[2]; + pp[3] = p0[3]; + pp += 4; + p0 += 4; + } +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 1 < max_kk; kk += 2) + { + pp[0] = p0[0]; + pp[1] = p0[1]; + pp += 2; + p0 += 2; + } + if (kk < max_kk) + *pp++ = p0[0]; + + for (int jj = 0; jj < 1; jj++) + pd[g * 1 + jj] = 1.f / B_scales.row(j + jj)[g]; + } + } + } + + BT = BT_packed; + BT_descales = BT_packed_descales; + return 0; +} + +static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int block_size) +{ +#if NCNN_RUNTIME_CPU && NCNN_ARM86SVEI8MM && __aarch64__ && !__ARM_FEATURE_SVE_MATMUL_INT8 + if (ncnn::cpu_support_arm_svei8mm()) + { + gemm_transB_packed_tile_wq_int8_svei8mm(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); + return; + } +#endif +#if NCNN_RUNTIME_CPU && NCNN_ARM84I8MM && __aarch64__ && !__ARM_FEATURE_MATMUL_INT8 + if (ncnn::cpu_support_arm_i8mm()) + { + gemm_transB_packed_tile_wq_int8_i8mm(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); + return; + } +#endif +#if NCNN_RUNTIME_CPU && NCNN_ARM82DOT && __aarch64__ && !__ARM_FEATURE_DOTPROD && !__ARM_FEATURE_MATMUL_INT8 + if (ncnn::cpu_support_arm_asimddp()) + { + gemm_transB_packed_tile_wq_int8_asimddp(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); + return; + } +#endif + + const signed char* pAT = AT_tile; + const int A_hstep = AT_tile.w; + const float* pAT_descales = AT_descales_tile; + const int A_descales_hstep = AT_descales_tile.w; + const signed char* pBT = BT_tile; + const float* pBT_descales = BT_descales_tile; + + float* outptr = topT_tile; + + int ii = 0; +#if __ARM_NEON +#if __aarch64__ + for (; ii + 7 < max_ii; ii += 8) + { + const signed char* pA_block = pAT; + const float* pA_descales_block = pAT_descales; + const signed char* pB = pBT; + const float* pB_descales = pBT_descales; + + int jj = 0; + for (; jj + 3 < max_jj; jj += 4) + { + const signed char* pA = pA_block; + const float* pA_descales = pA_descales_block; + float32x4_t _fsum0 = vdupq_n_f32(0.f); + float32x4_t _fsum1 = vdupq_n_f32(0.f); + float32x4_t _fsum2 = vdupq_n_f32(0.f); + float32x4_t _fsum3 = vdupq_n_f32(0.f); + float32x4_t _fsum4 = vdupq_n_f32(0.f); + float32x4_t _fsum5 = vdupq_n_f32(0.f); + float32x4_t _fsum6 = vdupq_n_f32(0.f); + float32x4_t _fsum7 = vdupq_n_f32(0.f); + + for (int k = 0; k < K; k += block_size) + { + int32x4_t _sum0 = vdupq_n_s32(0); + int32x4_t _sum1 = vdupq_n_s32(0); + int32x4_t _sum2 = vdupq_n_s32(0); + int32x4_t _sum3 = vdupq_n_s32(0); + int32x4_t _sum4 = vdupq_n_s32(0); + int32x4_t _sum5 = vdupq_n_s32(0); + int32x4_t _sum6 = vdupq_n_s32(0); + int32x4_t _sum7 = vdupq_n_s32(0); + const int max_kk = std::min(K - k, block_size); + int kk = 0; +#if __ARM_FEATURE_MATMUL_INT8 + int32x4_t _msum0 = vdupq_n_s32(0); + int32x4_t _msum1 = vdupq_n_s32(0); + int32x4_t _msum2 = vdupq_n_s32(0); + int32x4_t _msum3 = vdupq_n_s32(0); + int32x4_t _msum4 = vdupq_n_s32(0); + int32x4_t _msum5 = vdupq_n_s32(0); + int32x4_t _msum6 = vdupq_n_s32(0); + int32x4_t _msum7 = vdupq_n_s32(0); + for (; kk + 7 < max_kk; kk += 8) + { + const int8x16_t _b0 = vld1q_s8(pB); + const int8x16_t _b1 = vld1q_s8(pB + 16); + const int8x16_t _a01 = vld1q_s8(pA); + const int8x16_t _a23 = vld1q_s8(pA + 16); + const int8x16_t _a45 = vld1q_s8(pA + 32); + const int8x16_t _a67 = vld1q_s8(pA + 48); + _msum0 = vmmlaq_s32(_msum0, _a01, _b0); + _msum1 = vmmlaq_s32(_msum1, _a23, _b0); + _msum2 = vmmlaq_s32(_msum2, _a01, _b1); + _msum3 = vmmlaq_s32(_msum3, _a23, _b1); + _msum4 = vmmlaq_s32(_msum4, _a45, _b0); + _msum5 = vmmlaq_s32(_msum5, _a67, _b0); + _msum6 = vmmlaq_s32(_msum6, _a45, _b1); + _msum7 = vmmlaq_s32(_msum7, _a67, _b1); + pA += 64; + pB += 32; + } + _sum0 = vcombine_s32(vget_low_s32(_msum0), vget_low_s32(_msum2)); + _sum1 = vcombine_s32(vget_high_s32(_msum0), vget_high_s32(_msum2)); + _sum2 = vcombine_s32(vget_low_s32(_msum1), vget_low_s32(_msum3)); + _sum3 = vcombine_s32(vget_high_s32(_msum1), vget_high_s32(_msum3)); + _sum4 = vcombine_s32(vget_low_s32(_msum4), vget_low_s32(_msum6)); + _sum5 = vcombine_s32(vget_high_s32(_msum4), vget_high_s32(_msum6)); + _sum6 = vcombine_s32(vget_low_s32(_msum5), vget_low_s32(_msum7)); + _sum7 = vcombine_s32(vget_high_s32(_msum5), vget_high_s32(_msum7)); +#endif // __ARM_FEATURE_MATMUL_INT8 +#if __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk; kk += 4) + { + const int8x16_t _b0 = vld1q_s8(pB); + const int8x16_t _a0 = vld1q_s8(pA); + const int8x16_t _a1 = vld1q_s8(pA + 16); + _sum0 = vdotq_laneq_s32(_sum0, _b0, _a0, 0); + _sum1 = vdotq_laneq_s32(_sum1, _b0, _a0, 1); + _sum2 = vdotq_laneq_s32(_sum2, _b0, _a0, 2); + _sum3 = vdotq_laneq_s32(_sum3, _b0, _a0, 3); + _sum4 = vdotq_laneq_s32(_sum4, _b0, _a1, 0); + _sum5 = vdotq_laneq_s32(_sum5, _b0, _a1, 1); + _sum6 = vdotq_laneq_s32(_sum6, _b0, _a1, 2); + _sum7 = vdotq_laneq_s32(_sum7, _b0, _a1, 3); + pA += 32; + pB += 16; + } +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 1 < max_kk; kk += 2) + { + const int8x8_t _b0 = vld1_s8(pB); + const int16x8_t _a = vreinterpretq_s16_s8(vld1q_s8(pA)); + const int16x4_t _a0 = vget_low_s16(_a); + const int16x4_t _a1 = vget_high_s16(_a); + _sum0 = vaddq_s32(_sum0, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a0, 0))))); + _sum1 = vaddq_s32(_sum1, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a0, 1))))); + _sum2 = vaddq_s32(_sum2, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a0, 2))))); + _sum3 = vaddq_s32(_sum3, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a0, 3))))); + _sum4 = vaddq_s32(_sum4, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a1, 0))))); + _sum5 = vaddq_s32(_sum5, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a1, 1))))); + _sum6 = vaddq_s32(_sum6, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a1, 2))))); + _sum7 = vaddq_s32(_sum7, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a1, 3))))); + pA += 16; + pB += 8; + } + if (kk < max_kk) + { + const int8x8_t _b0 = vreinterpret_s8_s32(vld1_dup_s32((const int*)pB)); + const int8x8_t _a = vld1_s8(pA); + const int16x8_t _p0 = vmull_s8(_b0, vdup_lane_s8(_a, 0)); + const int16x8_t _p1 = vmull_s8(_b0, vdup_lane_s8(_a, 1)); + _sum0 = vaddq_s32(_sum0, vmovl_s16(vget_low_s16(_p0))); + _sum1 = vaddq_s32(_sum1, vmovl_s16(vget_low_s16(_p1))); + const int16x8_t _p2 = vmull_s8(_b0, vdup_lane_s8(_a, 2)); + const int16x8_t _p3 = vmull_s8(_b0, vdup_lane_s8(_a, 3)); + _sum2 = vaddq_s32(_sum2, vmovl_s16(vget_low_s16(_p2))); + _sum3 = vaddq_s32(_sum3, vmovl_s16(vget_low_s16(_p3))); + const int16x8_t _p4 = vmull_s8(_b0, vdup_lane_s8(_a, 4)); + const int16x8_t _p5 = vmull_s8(_b0, vdup_lane_s8(_a, 5)); + _sum4 = vaddq_s32(_sum4, vmovl_s16(vget_low_s16(_p4))); + _sum5 = vaddq_s32(_sum5, vmovl_s16(vget_low_s16(_p5))); + const int16x8_t _p6 = vmull_s8(_b0, vdup_lane_s8(_a, 6)); + const int16x8_t _p7 = vmull_s8(_b0, vdup_lane_s8(_a, 7)); + _sum6 = vaddq_s32(_sum6, vmovl_s16(vget_low_s16(_p6))); + _sum7 = vaddq_s32(_sum7, vmovl_s16(vget_low_s16(_p7))); + pA += 8; + pB += 4; + } + + const float32x4_t _bd0 = vld1q_f32(pB_descales); + const float32x4_t _ad0 = vld1q_f32(pA_descales); + const float32x4_t _ad1 = vld1q_f32(pA_descales + 4); + _fsum0 = vmlaq_f32(_fsum0, vcvtq_f32_s32(_sum0), vmulq_laneq_f32(_bd0, _ad0, 0)); + _fsum1 = vmlaq_f32(_fsum1, vcvtq_f32_s32(_sum1), vmulq_laneq_f32(_bd0, _ad0, 1)); + _fsum2 = vmlaq_f32(_fsum2, vcvtq_f32_s32(_sum2), vmulq_laneq_f32(_bd0, _ad0, 2)); + _fsum3 = vmlaq_f32(_fsum3, vcvtq_f32_s32(_sum3), vmulq_laneq_f32(_bd0, _ad0, 3)); + _fsum4 = vmlaq_f32(_fsum4, vcvtq_f32_s32(_sum4), vmulq_laneq_f32(_bd0, _ad1, 0)); + _fsum5 = vmlaq_f32(_fsum5, vcvtq_f32_s32(_sum5), vmulq_laneq_f32(_bd0, _ad1, 1)); + _fsum6 = vmlaq_f32(_fsum6, vcvtq_f32_s32(_sum6), vmulq_laneq_f32(_bd0, _ad1, 2)); + _fsum7 = vmlaq_f32(_fsum7, vcvtq_f32_s32(_sum7), vmulq_laneq_f32(_bd0, _ad1, 3)); + pA_descales += 8; + pB_descales += 4; + } + + vst1q_f32(outptr, _fsum0); + vst1q_f32(outptr + 4, _fsum1); + vst1q_f32(outptr + 8, _fsum2); + vst1q_f32(outptr + 12, _fsum3); + vst1q_f32(outptr + 16, _fsum4); + vst1q_f32(outptr + 20, _fsum5); + vst1q_f32(outptr + 24, _fsum6); + vst1q_f32(outptr + 28, _fsum7); + outptr += 32; + } + for (; jj + 1 < max_jj; jj += 2) + { + const signed char* pA = pA_block; + const float* pA_descales = pA_descales_block; + float32x4_t _fsum0 = vdupq_n_f32(0.f); + float32x4_t _fsum1 = vdupq_n_f32(0.f); + float32x4_t _fsum2 = vdupq_n_f32(0.f); + float32x4_t _fsum3 = vdupq_n_f32(0.f); + float32x4_t _fsum4 = vdupq_n_f32(0.f); + float32x4_t _fsum5 = vdupq_n_f32(0.f); + float32x4_t _fsum6 = vdupq_n_f32(0.f); + float32x4_t _fsum7 = vdupq_n_f32(0.f); + + for (int k = 0; k < K; k += block_size) + { + int32x4_t _sum0 = vdupq_n_s32(0); + int32x4_t _sum1 = vdupq_n_s32(0); + int32x4_t _sum2 = vdupq_n_s32(0); + int32x4_t _sum3 = vdupq_n_s32(0); + int32x4_t _sum4 = vdupq_n_s32(0); + int32x4_t _sum5 = vdupq_n_s32(0); + int32x4_t _sum6 = vdupq_n_s32(0); + int32x4_t _sum7 = vdupq_n_s32(0); + const int max_kk = std::min(K - k, block_size); + int kk = 0; +#if __ARM_FEATURE_MATMUL_INT8 + int32x4_t _msum0 = vdupq_n_s32(0); + int32x4_t _msum1 = vdupq_n_s32(0); + int32x4_t _msum2 = vdupq_n_s32(0); + int32x4_t _msum3 = vdupq_n_s32(0); + for (; kk + 7 < max_kk; kk += 8) + { + const int8x16_t _b0 = vld1q_s8(pB); + const int8x16_t _a01 = vld1q_s8(pA); + const int8x16_t _a23 = vld1q_s8(pA + 16); + const int8x16_t _a45 = vld1q_s8(pA + 32); + const int8x16_t _a67 = vld1q_s8(pA + 48); + _msum0 = vmmlaq_s32(_msum0, _a01, _b0); + _msum1 = vmmlaq_s32(_msum1, _a23, _b0); + _msum2 = vmmlaq_s32(_msum2, _a45, _b0); + _msum3 = vmmlaq_s32(_msum3, _a67, _b0); + pA += 64; + pB += 16; + } + _sum0 = vcombine_s32(vget_low_s32(_msum0), vdup_n_s32(0)); + _sum1 = vcombine_s32(vget_high_s32(_msum0), vdup_n_s32(0)); + _sum2 = vcombine_s32(vget_low_s32(_msum1), vdup_n_s32(0)); + _sum3 = vcombine_s32(vget_high_s32(_msum1), vdup_n_s32(0)); + _sum4 = vcombine_s32(vget_low_s32(_msum2), vdup_n_s32(0)); + _sum5 = vcombine_s32(vget_high_s32(_msum2), vdup_n_s32(0)); + _sum6 = vcombine_s32(vget_low_s32(_msum3), vdup_n_s32(0)); + _sum7 = vcombine_s32(vget_high_s32(_msum3), vdup_n_s32(0)); +#endif // __ARM_FEATURE_MATMUL_INT8 +#if __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk; kk += 4) + { + const int8x16_t _b0 = vcombine_s8(vld1_s8(pB), vdup_n_s8(0)); + const int8x16_t _a0 = vld1q_s8(pA); + const int8x16_t _a1 = vld1q_s8(pA + 16); + _sum0 = vdotq_laneq_s32(_sum0, _b0, _a0, 0); + _sum1 = vdotq_laneq_s32(_sum1, _b0, _a0, 1); + _sum2 = vdotq_laneq_s32(_sum2, _b0, _a0, 2); + _sum3 = vdotq_laneq_s32(_sum3, _b0, _a0, 3); + _sum4 = vdotq_laneq_s32(_sum4, _b0, _a1, 0); + _sum5 = vdotq_laneq_s32(_sum5, _b0, _a1, 1); + _sum6 = vdotq_laneq_s32(_sum6, _b0, _a1, 2); + _sum7 = vdotq_laneq_s32(_sum7, _b0, _a1, 3); + pA += 32; + pB += 8; + } +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 1 < max_kk; kk += 2) + { + const int8x8_t _b0 = vreinterpret_s8_s32(vld1_dup_s32((const int*)pB)); + const int16x8_t _a = vreinterpretq_s16_s8(vld1q_s8(pA)); + const int16x4_t _a0 = vget_low_s16(_a); + const int16x4_t _a1 = vget_high_s16(_a); + _sum0 = vaddq_s32(_sum0, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a0, 0))))); + _sum1 = vaddq_s32(_sum1, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a0, 1))))); + _sum2 = vaddq_s32(_sum2, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a0, 2))))); + _sum3 = vaddq_s32(_sum3, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a0, 3))))); + _sum4 = vaddq_s32(_sum4, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a1, 0))))); + _sum5 = vaddq_s32(_sum5, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a1, 1))))); + _sum6 = vaddq_s32(_sum6, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a1, 2))))); + _sum7 = vaddq_s32(_sum7, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a1, 3))))); + pA += 16; + pB += 4; + } + if (kk < max_kk) + { + const int8x8_t _b0 = vreinterpret_s8_s16(vld1_dup_s16((const short*)pB)); + const int8x8_t _a = vld1_s8(pA); + const int16x8_t _p0 = vmull_s8(_b0, vdup_lane_s8(_a, 0)); + const int16x8_t _p1 = vmull_s8(_b0, vdup_lane_s8(_a, 1)); + const int16x8_t _p2 = vmull_s8(_b0, vdup_lane_s8(_a, 2)); + const int16x8_t _p3 = vmull_s8(_b0, vdup_lane_s8(_a, 3)); + const int16x8_t _p4 = vmull_s8(_b0, vdup_lane_s8(_a, 4)); + const int16x8_t _p5 = vmull_s8(_b0, vdup_lane_s8(_a, 5)); + const int16x8_t _p6 = vmull_s8(_b0, vdup_lane_s8(_a, 6)); + const int16x8_t _p7 = vmull_s8(_b0, vdup_lane_s8(_a, 7)); + _sum0 = vaddq_s32(_sum0, vmovl_s16(vget_low_s16(_p0))); + _sum1 = vaddq_s32(_sum1, vmovl_s16(vget_low_s16(_p1))); + _sum2 = vaddq_s32(_sum2, vmovl_s16(vget_low_s16(_p2))); + _sum3 = vaddq_s32(_sum3, vmovl_s16(vget_low_s16(_p3))); + _sum4 = vaddq_s32(_sum4, vmovl_s16(vget_low_s16(_p4))); + _sum5 = vaddq_s32(_sum5, vmovl_s16(vget_low_s16(_p5))); + _sum6 = vaddq_s32(_sum6, vmovl_s16(vget_low_s16(_p6))); + _sum7 = vaddq_s32(_sum7, vmovl_s16(vget_low_s16(_p7))); + pA += 8; + pB += 2; + } + + const float32x4_t _bd = vcombine_f32(vld1_f32(pB_descales), vdup_n_f32(0.f)); + const float32x4_t _ad0 = vld1q_f32(pA_descales); + const float32x4_t _ad1 = vld1q_f32(pA_descales + 4); + _fsum0 = vmlaq_f32(_fsum0, vcvtq_f32_s32(_sum0), vmulq_laneq_f32(_bd, _ad0, 0)); + _fsum1 = vmlaq_f32(_fsum1, vcvtq_f32_s32(_sum1), vmulq_laneq_f32(_bd, _ad0, 1)); + _fsum2 = vmlaq_f32(_fsum2, vcvtq_f32_s32(_sum2), vmulq_laneq_f32(_bd, _ad0, 2)); + _fsum3 = vmlaq_f32(_fsum3, vcvtq_f32_s32(_sum3), vmulq_laneq_f32(_bd, _ad0, 3)); + _fsum4 = vmlaq_f32(_fsum4, vcvtq_f32_s32(_sum4), vmulq_laneq_f32(_bd, _ad1, 0)); + _fsum5 = vmlaq_f32(_fsum5, vcvtq_f32_s32(_sum5), vmulq_laneq_f32(_bd, _ad1, 1)); + _fsum6 = vmlaq_f32(_fsum6, vcvtq_f32_s32(_sum6), vmulq_laneq_f32(_bd, _ad1, 2)); + _fsum7 = vmlaq_f32(_fsum7, vcvtq_f32_s32(_sum7), vmulq_laneq_f32(_bd, _ad1, 3)); + pA_descales += 8; + pB_descales += 2; + } + + vst1_f32(outptr, vget_low_f32(_fsum0)); + vst1_f32(outptr + 2, vget_low_f32(_fsum1)); + vst1_f32(outptr + 4, vget_low_f32(_fsum2)); + vst1_f32(outptr + 6, vget_low_f32(_fsum3)); + vst1_f32(outptr + 8, vget_low_f32(_fsum4)); + vst1_f32(outptr + 10, vget_low_f32(_fsum5)); + vst1_f32(outptr + 12, vget_low_f32(_fsum6)); + vst1_f32(outptr + 14, vget_low_f32(_fsum7)); + outptr += 16; + } + for (; jj < max_jj; jj++) + { + const signed char* pA = pA_block; + const float* pA_descales = pA_descales_block; + float32x4_t _fsum0 = vdupq_n_f32(0.f); + float32x4_t _fsum1 = vdupq_n_f32(0.f); + float32x4_t _fsum2 = vdupq_n_f32(0.f); + float32x4_t _fsum3 = vdupq_n_f32(0.f); + float32x4_t _fsum4 = vdupq_n_f32(0.f); + float32x4_t _fsum5 = vdupq_n_f32(0.f); + float32x4_t _fsum6 = vdupq_n_f32(0.f); + float32x4_t _fsum7 = vdupq_n_f32(0.f); + + for (int k = 0; k < K; k += block_size) + { + int32x4_t _sum0 = vdupq_n_s32(0); + int32x4_t _sum1 = vdupq_n_s32(0); + int32x4_t _sum2 = vdupq_n_s32(0); + int32x4_t _sum3 = vdupq_n_s32(0); + int32x4_t _sum4 = vdupq_n_s32(0); + int32x4_t _sum5 = vdupq_n_s32(0); + int32x4_t _sum6 = vdupq_n_s32(0); + int32x4_t _sum7 = vdupq_n_s32(0); + const int max_kk = std::min(K - k, block_size); + int kk = 0; +#if __ARM_FEATURE_MATMUL_INT8 + for (; kk + 7 < max_kk; kk += 8) + { + const int8x16_t _a01 = vld1q_s8(pA); + const int8x16_t _a23 = vld1q_s8(pA + 16); + const int8x16_t _a45 = vld1q_s8(pA + 32); + const int8x16_t _a67 = vld1q_s8(pA + 48); + const int8x16_t _b0 = vcombine_s8(vreinterpret_s8_s32(vld1_dup_s32((const int*)pB)), vdup_n_s8(0)); + const int8x16_t _b1 = vcombine_s8(vreinterpret_s8_s32(vld1_dup_s32((const int*)(pB + 4))), vdup_n_s8(0)); + _sum0 = vdotq_laneq_s32(_sum0, _b0, _a01, 0); _sum0 = vdotq_laneq_s32(_sum0, _b1, _a01, 1); + _sum1 = vdotq_laneq_s32(_sum1, _b0, _a01, 2); _sum1 = vdotq_laneq_s32(_sum1, _b1, _a01, 3); + _sum2 = vdotq_laneq_s32(_sum2, _b0, _a23, 0); _sum2 = vdotq_laneq_s32(_sum2, _b1, _a23, 1); + _sum3 = vdotq_laneq_s32(_sum3, _b0, _a23, 2); _sum3 = vdotq_laneq_s32(_sum3, _b1, _a23, 3); + _sum4 = vdotq_laneq_s32(_sum4, _b0, _a45, 0); _sum4 = vdotq_laneq_s32(_sum4, _b1, _a45, 1); + _sum5 = vdotq_laneq_s32(_sum5, _b0, _a45, 2); _sum5 = vdotq_laneq_s32(_sum5, _b1, _a45, 3); + _sum6 = vdotq_laneq_s32(_sum6, _b0, _a67, 0); _sum6 = vdotq_laneq_s32(_sum6, _b1, _a67, 1); + _sum7 = vdotq_laneq_s32(_sum7, _b0, _a67, 2); _sum7 = vdotq_laneq_s32(_sum7, _b1, _a67, 3); + pA += 64; + pB += 8; + } +#endif // __ARM_FEATURE_MATMUL_INT8 +#if __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk; kk += 4) + { + const int8x16_t _b = vcombine_s8(vreinterpret_s8_s32(vld1_dup_s32((const int*)pB)), vdup_n_s8(0)); + const int8x16_t _a0 = vld1q_s8(pA); + const int8x16_t _a1 = vld1q_s8(pA + 16); + _sum0 = vdotq_laneq_s32(_sum0, _b, _a0, 0); _sum1 = vdotq_laneq_s32(_sum1, _b, _a0, 1); + _sum2 = vdotq_laneq_s32(_sum2, _b, _a0, 2); _sum3 = vdotq_laneq_s32(_sum3, _b, _a0, 3); + _sum4 = vdotq_laneq_s32(_sum4, _b, _a1, 0); _sum5 = vdotq_laneq_s32(_sum5, _b, _a1, 1); + _sum6 = vdotq_laneq_s32(_sum6, _b, _a1, 2); _sum7 = vdotq_laneq_s32(_sum7, _b, _a1, 3); + pA += 32; + pB += 4; + } +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 1 < max_kk; kk += 2) + { + const int8x8_t _b = vreinterpret_s8_s16(vld1_dup_s16((const short*)pB)); + const int16x8_t _a = vreinterpretq_s16_s8(vld1q_s8(pA)); + const int16x4_t _a0 = vget_low_s16(_a); + const int16x4_t _a1 = vget_high_s16(_a); + _sum0 = vaddq_s32(_sum0, vpaddlq_s16(vmull_s8(_b, vreinterpret_s8_s16(vdup_lane_s16(_a0, 0))))); + _sum1 = vaddq_s32(_sum1, vpaddlq_s16(vmull_s8(_b, vreinterpret_s8_s16(vdup_lane_s16(_a0, 1))))); + _sum2 = vaddq_s32(_sum2, vpaddlq_s16(vmull_s8(_b, vreinterpret_s8_s16(vdup_lane_s16(_a0, 2))))); + _sum3 = vaddq_s32(_sum3, vpaddlq_s16(vmull_s8(_b, vreinterpret_s8_s16(vdup_lane_s16(_a0, 3))))); + _sum4 = vaddq_s32(_sum4, vpaddlq_s16(vmull_s8(_b, vreinterpret_s8_s16(vdup_lane_s16(_a1, 0))))); + _sum5 = vaddq_s32(_sum5, vpaddlq_s16(vmull_s8(_b, vreinterpret_s8_s16(vdup_lane_s16(_a1, 1))))); + _sum6 = vaddq_s32(_sum6, vpaddlq_s16(vmull_s8(_b, vreinterpret_s8_s16(vdup_lane_s16(_a1, 2))))); + _sum7 = vaddq_s32(_sum7, vpaddlq_s16(vmull_s8(_b, vreinterpret_s8_s16(vdup_lane_s16(_a1, 3))))); + pA += 16; + pB += 2; + } + if (kk < max_kk) + { + const int8x8_t _b = vset_lane_s8(pB[0], vdup_n_s8(0), 0); + const int8x8_t _a = vld1_s8(pA); + const int16x8_t _p0 = vmull_s8(_b, vdup_lane_s8(_a, 0)); const int16x8_t _p1 = vmull_s8(_b, vdup_lane_s8(_a, 1)); + const int16x8_t _p2 = vmull_s8(_b, vdup_lane_s8(_a, 2)); const int16x8_t _p3 = vmull_s8(_b, vdup_lane_s8(_a, 3)); + const int16x8_t _p4 = vmull_s8(_b, vdup_lane_s8(_a, 4)); const int16x8_t _p5 = vmull_s8(_b, vdup_lane_s8(_a, 5)); + const int16x8_t _p6 = vmull_s8(_b, vdup_lane_s8(_a, 6)); const int16x8_t _p7 = vmull_s8(_b, vdup_lane_s8(_a, 7)); + _sum0 = vaddq_s32(_sum0, vmovl_s16(vget_low_s16(_p0))); _sum1 = vaddq_s32(_sum1, vmovl_s16(vget_low_s16(_p1))); + _sum2 = vaddq_s32(_sum2, vmovl_s16(vget_low_s16(_p2))); _sum3 = vaddq_s32(_sum3, vmovl_s16(vget_low_s16(_p3))); + _sum4 = vaddq_s32(_sum4, vmovl_s16(vget_low_s16(_p4))); _sum5 = vaddq_s32(_sum5, vmovl_s16(vget_low_s16(_p5))); + _sum6 = vaddq_s32(_sum6, vmovl_s16(vget_low_s16(_p6))); _sum7 = vaddq_s32(_sum7, vmovl_s16(vget_low_s16(_p7))); + pA += 8; + pB++; + } + + const float32x4_t _bd = vdupq_n_f32(pB_descales[0]); + const float32x4_t _ad0 = vld1q_f32(pA_descales); + const float32x4_t _ad1 = vld1q_f32(pA_descales + 4); + _fsum0 = vmlaq_f32(_fsum0, vcvtq_f32_s32(_sum0), vmulq_laneq_f32(_bd, _ad0, 0)); + _fsum1 = vmlaq_f32(_fsum1, vcvtq_f32_s32(_sum1), vmulq_laneq_f32(_bd, _ad0, 1)); + _fsum2 = vmlaq_f32(_fsum2, vcvtq_f32_s32(_sum2), vmulq_laneq_f32(_bd, _ad0, 2)); + _fsum3 = vmlaq_f32(_fsum3, vcvtq_f32_s32(_sum3), vmulq_laneq_f32(_bd, _ad0, 3)); + _fsum4 = vmlaq_f32(_fsum4, vcvtq_f32_s32(_sum4), vmulq_laneq_f32(_bd, _ad1, 0)); + _fsum5 = vmlaq_f32(_fsum5, vcvtq_f32_s32(_sum5), vmulq_laneq_f32(_bd, _ad1, 1)); + _fsum6 = vmlaq_f32(_fsum6, vcvtq_f32_s32(_sum6), vmulq_laneq_f32(_bd, _ad1, 2)); + _fsum7 = vmlaq_f32(_fsum7, vcvtq_f32_s32(_sum7), vmulq_laneq_f32(_bd, _ad1, 3)); + pA_descales += 8; + pB_descales++; + } + + vst1q_lane_f32(outptr, _fsum0, 0); outptr++; + vst1q_lane_f32(outptr, _fsum1, 0); outptr++; + vst1q_lane_f32(outptr, _fsum2, 0); outptr++; + vst1q_lane_f32(outptr, _fsum3, 0); outptr++; + vst1q_lane_f32(outptr, _fsum4, 0); outptr++; + vst1q_lane_f32(outptr, _fsum5, 0); outptr++; + vst1q_lane_f32(outptr, _fsum6, 0); outptr++; + vst1q_lane_f32(outptr, _fsum7, 0); outptr++; + } + } +#endif // __aarch64__ + for (; ii + 3 < max_ii; ii += 4) + { + const signed char* pA0_block = pAT + ii * A_hstep; + const float* pA_descales0_block = pAT_descales + ii * A_descales_hstep; + const signed char* pB = pBT; + const float* pB_descales = pBT_descales; + + int jj = 0; +#if __aarch64__ + for (; jj + 7 < max_jj; jj += 8) + { + const signed char* pB0 = pB; + const signed char* pB1 = pB + 4 * K; + const float* pB_descales0 = pB_descales; + const float* pB_descales1 = pB_descales + 4 * ((K + block_size - 1) / block_size); + float32x4_t _fsum0 = vdupq_n_f32(0.f); + float32x4_t _fsum1 = vdupq_n_f32(0.f); + float32x4_t _fsum2 = vdupq_n_f32(0.f); + float32x4_t _fsum3 = vdupq_n_f32(0.f); + float32x4_t _fsum4 = vdupq_n_f32(0.f); + float32x4_t _fsum5 = vdupq_n_f32(0.f); + float32x4_t _fsum6 = vdupq_n_f32(0.f); + float32x4_t _fsum7 = vdupq_n_f32(0.f); + + const signed char* pA = pA0_block; + const float* pA_descales = pA_descales0_block; + for (int k = 0; k < K; k += block_size) + { + int32x4_t _sum0 = vdupq_n_s32(0); + int32x4_t _sum1 = vdupq_n_s32(0); + int32x4_t _sum2 = vdupq_n_s32(0); + int32x4_t _sum3 = vdupq_n_s32(0); + int32x4_t _sum4 = vdupq_n_s32(0); + int32x4_t _sum5 = vdupq_n_s32(0); + int32x4_t _sum6 = vdupq_n_s32(0); + int32x4_t _sum7 = vdupq_n_s32(0); + const int max_kk = std::min(K - k, block_size); + int kk = 0; +#if __ARM_FEATURE_MATMUL_INT8 + int32x4_t _msum0 = vdupq_n_s32(0); + int32x4_t _msum1 = vdupq_n_s32(0); + int32x4_t _msum2 = vdupq_n_s32(0); + int32x4_t _msum3 = vdupq_n_s32(0); + int32x4_t _msum4 = vdupq_n_s32(0); + int32x4_t _msum5 = vdupq_n_s32(0); + int32x4_t _msum6 = vdupq_n_s32(0); + int32x4_t _msum7 = vdupq_n_s32(0); + for (; kk + 7 < max_kk; kk += 8) + { + const int8x16_t _a01 = vld1q_s8(pA); + const int8x16_t _a23 = vld1q_s8(pA + 16); + const int8x16_t _b00 = vld1q_s8(pB0); + const int8x16_t _b01 = vld1q_s8(pB0 + 16); + const int8x16_t _b10 = vld1q_s8(pB1); + const int8x16_t _b11 = vld1q_s8(pB1 + 16); + _msum0 = vmmlaq_s32(_msum0, _a01, _b00); + _msum1 = vmmlaq_s32(_msum1, _a01, _b01); + _msum2 = vmmlaq_s32(_msum2, _a23, _b00); + _msum3 = vmmlaq_s32(_msum3, _a23, _b01); + _msum4 = vmmlaq_s32(_msum4, _a01, _b10); + _msum5 = vmmlaq_s32(_msum5, _a01, _b11); + _msum6 = vmmlaq_s32(_msum6, _a23, _b10); + _msum7 = vmmlaq_s32(_msum7, _a23, _b11); + pA += 32; + pB0 += 32; + pB1 += 32; + } + _sum0 = vcombine_s32(vget_low_s32(_msum0), vget_low_s32(_msum1)); + _sum1 = vcombine_s32(vget_low_s32(_msum4), vget_low_s32(_msum5)); + _sum2 = vcombine_s32(vget_high_s32(_msum0), vget_high_s32(_msum1)); + _sum3 = vcombine_s32(vget_high_s32(_msum4), vget_high_s32(_msum5)); + _sum4 = vcombine_s32(vget_low_s32(_msum2), vget_low_s32(_msum3)); + _sum5 = vcombine_s32(vget_low_s32(_msum6), vget_low_s32(_msum7)); + _sum6 = vcombine_s32(vget_high_s32(_msum2), vget_high_s32(_msum3)); + _sum7 = vcombine_s32(vget_high_s32(_msum6), vget_high_s32(_msum7)); +#endif // __ARM_FEATURE_MATMUL_INT8 +#if __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk; kk += 4) + { + const int8x16_t _a = vld1q_s8(pA); + const int8x16_t _b0 = vld1q_s8(pB0); + const int8x16_t _b1 = vld1q_s8(pB1); + _sum0 = vdotq_laneq_s32(_sum0, _b0, _a, 0); + _sum1 = vdotq_laneq_s32(_sum1, _b1, _a, 0); + _sum2 = vdotq_laneq_s32(_sum2, _b0, _a, 1); + _sum3 = vdotq_laneq_s32(_sum3, _b1, _a, 1); + _sum4 = vdotq_laneq_s32(_sum4, _b0, _a, 2); + _sum5 = vdotq_laneq_s32(_sum5, _b1, _a, 2); + _sum6 = vdotq_laneq_s32(_sum6, _b0, _a, 3); + _sum7 = vdotq_laneq_s32(_sum7, _b1, _a, 3); + pA += 16; + pB0 += 16; + pB1 += 16; + } +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 1 < max_kk; kk += 2) + { + const int16x4_t _a = vreinterpret_s16_s8(vld1_s8(pA)); + const int8x8_t _b0 = vld1_s8(pB0); + const int8x8_t _b1 = vld1_s8(pB1); + _sum0 = vaddq_s32(_sum0, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 0))))); + _sum1 = vaddq_s32(_sum1, vpaddlq_s16(vmull_s8(_b1, vreinterpret_s8_s16(vdup_lane_s16(_a, 0))))); + _sum2 = vaddq_s32(_sum2, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 1))))); + _sum3 = vaddq_s32(_sum3, vpaddlq_s16(vmull_s8(_b1, vreinterpret_s8_s16(vdup_lane_s16(_a, 1))))); + _sum4 = vaddq_s32(_sum4, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 2))))); + _sum5 = vaddq_s32(_sum5, vpaddlq_s16(vmull_s8(_b1, vreinterpret_s8_s16(vdup_lane_s16(_a, 2))))); + _sum6 = vaddq_s32(_sum6, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 3))))); + _sum7 = vaddq_s32(_sum7, vpaddlq_s16(vmull_s8(_b1, vreinterpret_s8_s16(vdup_lane_s16(_a, 3))))); + pA += 8; + pB0 += 8; + pB1 += 8; + } + if (kk < max_kk) + { + const int8x8_t _a = vreinterpret_s8_s32(vld1_lane_s32((const int*)pA, vdup_n_s32(0), 0)); + const int8x8_t _b0 = vreinterpret_s8_s32(vld1_dup_s32((const int*)pB0)); + const int8x8_t _b1 = vreinterpret_s8_s32(vld1_dup_s32((const int*)pB1)); + const int16x8_t _p00 = vmull_s8(_b0, vdup_lane_s8(_a, 0)); + const int16x8_t _p01 = vmull_s8(_b1, vdup_lane_s8(_a, 0)); + const int16x8_t _p10 = vmull_s8(_b0, vdup_lane_s8(_a, 1)); + const int16x8_t _p11 = vmull_s8(_b1, vdup_lane_s8(_a, 1)); + const int16x8_t _p20 = vmull_s8(_b0, vdup_lane_s8(_a, 2)); + const int16x8_t _p21 = vmull_s8(_b1, vdup_lane_s8(_a, 2)); + const int16x8_t _p30 = vmull_s8(_b0, vdup_lane_s8(_a, 3)); + const int16x8_t _p31 = vmull_s8(_b1, vdup_lane_s8(_a, 3)); + _sum0 = vaddq_s32(_sum0, vmovl_s16(vget_low_s16(_p00))); + _sum1 = vaddq_s32(_sum1, vmovl_s16(vget_low_s16(_p01))); + _sum2 = vaddq_s32(_sum2, vmovl_s16(vget_low_s16(_p10))); + _sum3 = vaddq_s32(_sum3, vmovl_s16(vget_low_s16(_p11))); + _sum4 = vaddq_s32(_sum4, vmovl_s16(vget_low_s16(_p20))); + _sum5 = vaddq_s32(_sum5, vmovl_s16(vget_low_s16(_p21))); + _sum6 = vaddq_s32(_sum6, vmovl_s16(vget_low_s16(_p30))); + _sum7 = vaddq_s32(_sum7, vmovl_s16(vget_low_s16(_p31))); + pA += 4; + pB0 += 4; + pB1 += 4; + } + + const float32x4_t _bd0 = vld1q_f32(pB_descales0); + const float32x4_t _bd1 = vld1q_f32(pB_descales1); + const float32x4_t _ad = vld1q_f32(pA_descales); + _fsum0 = vmlaq_f32(_fsum0, vcvtq_f32_s32(_sum0), vmulq_laneq_f32(_bd0, _ad, 0)); + _fsum1 = vmlaq_f32(_fsum1, vcvtq_f32_s32(_sum1), vmulq_laneq_f32(_bd1, _ad, 0)); + _fsum2 = vmlaq_f32(_fsum2, vcvtq_f32_s32(_sum2), vmulq_laneq_f32(_bd0, _ad, 1)); + _fsum3 = vmlaq_f32(_fsum3, vcvtq_f32_s32(_sum3), vmulq_laneq_f32(_bd1, _ad, 1)); + _fsum4 = vmlaq_f32(_fsum4, vcvtq_f32_s32(_sum4), vmulq_laneq_f32(_bd0, _ad, 2)); + _fsum5 = vmlaq_f32(_fsum5, vcvtq_f32_s32(_sum5), vmulq_laneq_f32(_bd1, _ad, 2)); + _fsum6 = vmlaq_f32(_fsum6, vcvtq_f32_s32(_sum6), vmulq_laneq_f32(_bd0, _ad, 3)); + _fsum7 = vmlaq_f32(_fsum7, vcvtq_f32_s32(_sum7), vmulq_laneq_f32(_bd1, _ad, 3)); + + pA_descales += 4; + pB_descales0 += 4; + pB_descales1 += 4; + } + + pB = pB1; + pB_descales = pB_descales1; + + vst1q_f32(outptr, _fsum0); + outptr += 4; + vst1q_f32(outptr, _fsum1); + outptr += 4; + vst1q_f32(outptr, _fsum2); + outptr += 4; + vst1q_f32(outptr, _fsum3); + outptr += 4; + vst1q_f32(outptr, _fsum4); + outptr += 4; + vst1q_f32(outptr, _fsum5); + outptr += 4; + vst1q_f32(outptr, _fsum6); + outptr += 4; + vst1q_f32(outptr, _fsum7); + outptr += 4; + } +#endif // __aarch64__ + for (; jj + 3 < max_jj; jj += 4) + { + float32x4_t _fsum0 = vdupq_n_f32(0.f); + float32x4_t _fsum1 = vdupq_n_f32(0.f); + float32x4_t _fsum2 = vdupq_n_f32(0.f); + float32x4_t _fsum3 = vdupq_n_f32(0.f); + + const signed char* pA = pA0_block; + const float* pA_descales = pA_descales0_block; + for (int k = 0; k < K; k += block_size) + { + int32x4_t _sum0 = vdupq_n_s32(0); + int32x4_t _sum1 = vdupq_n_s32(0); + int32x4_t _sum2 = vdupq_n_s32(0); + int32x4_t _sum3 = vdupq_n_s32(0); + const int max_kk = std::min(K - k, block_size); + int kk = 0; +#if __ARM_FEATURE_MATMUL_INT8 + int32x4_t _msum0 = vdupq_n_s32(0); + int32x4_t _msum1 = vdupq_n_s32(0); + int32x4_t _msum2 = vdupq_n_s32(0); + int32x4_t _msum3 = vdupq_n_s32(0); + for (; kk + 7 < max_kk; kk += 8) + { + const int8x16_t _b0 = vld1q_s8(pB); + const int8x16_t _b1 = vld1q_s8(pB + 16); + const int8x16_t _a01 = vld1q_s8(pA); + const int8x16_t _a23 = vld1q_s8(pA + 16); + _msum0 = vmmlaq_s32(_msum0, _a01, _b0); + _msum1 = vmmlaq_s32(_msum1, _a23, _b0); + _msum2 = vmmlaq_s32(_msum2, _a01, _b1); + _msum3 = vmmlaq_s32(_msum3, _a23, _b1); + pA += 32; + pB += 32; + } + _sum0 = vcombine_s32(vget_low_s32(_msum0), vget_low_s32(_msum2)); + _sum1 = vcombine_s32(vget_high_s32(_msum0), vget_high_s32(_msum2)); + _sum2 = vcombine_s32(vget_low_s32(_msum1), vget_low_s32(_msum3)); + _sum3 = vcombine_s32(vget_high_s32(_msum1), vget_high_s32(_msum3)); +#endif // __ARM_FEATURE_MATMUL_INT8 +#if __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk; kk += 4) + { + const int8x16_t _b0 = vld1q_s8(pB); + const int8x16_t _a = vld1q_s8(pA); + _sum0 = vdotq_laneq_s32(_sum0, _b0, _a, 0); + _sum1 = vdotq_laneq_s32(_sum1, _b0, _a, 1); + _sum2 = vdotq_laneq_s32(_sum2, _b0, _a, 2); + _sum3 = vdotq_laneq_s32(_sum3, _b0, _a, 3); + pA += 16; + pB += 16; + } +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 1 < max_kk; kk += 2) + { + const int8x8_t _b0 = vld1_s8(pB); + const int16x4_t _a = vreinterpret_s16_s8(vld1_s8(pA)); + _sum0 = vaddq_s32(_sum0, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 0))))); + _sum1 = vaddq_s32(_sum1, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 1))))); + _sum2 = vaddq_s32(_sum2, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 2))))); + _sum3 = vaddq_s32(_sum3, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 3))))); + pA += 8; + pB += 8; + } + if (kk < max_kk) + { + const int8x8_t _b = vreinterpret_s8_s32(vld1_dup_s32((const int*)(pB))); + const int8x8_t _a = vreinterpret_s8_s32(vld1_lane_s32((const int*)pA, vdup_n_s32(0), 0)); + const int16x8_t _p0 = vmull_s8(_b, vdup_lane_s8(_a, 0)); + const int16x8_t _p1 = vmull_s8(_b, vdup_lane_s8(_a, 1)); + const int16x8_t _p2 = vmull_s8(_b, vdup_lane_s8(_a, 2)); + const int16x8_t _p3 = vmull_s8(_b, vdup_lane_s8(_a, 3)); + _sum0 = vaddq_s32(_sum0, vmovl_s16(vget_low_s16(_p0))); + _sum1 = vaddq_s32(_sum1, vmovl_s16(vget_low_s16(_p1))); + _sum2 = vaddq_s32(_sum2, vmovl_s16(vget_low_s16(_p2))); + _sum3 = vaddq_s32(_sum3, vmovl_s16(vget_low_s16(_p3))); + pA += 4; + pB += 4; + } + + const float32x4_t _bd0 = vld1q_f32(pB_descales); + const float32x4_t _ad = vld1q_f32(pA_descales); + _fsum0 = vmlaq_f32(_fsum0, vcvtq_f32_s32(_sum0), vmulq_n_f32(_bd0, vgetq_lane_f32(_ad, 0))); + _fsum1 = vmlaq_f32(_fsum1, vcvtq_f32_s32(_sum1), vmulq_n_f32(_bd0, vgetq_lane_f32(_ad, 1))); + _fsum2 = vmlaq_f32(_fsum2, vcvtq_f32_s32(_sum2), vmulq_n_f32(_bd0, vgetq_lane_f32(_ad, 2))); + _fsum3 = vmlaq_f32(_fsum3, vcvtq_f32_s32(_sum3), vmulq_n_f32(_bd0, vgetq_lane_f32(_ad, 3))); + + pA_descales += 4; + pB_descales += 4; + } + + vst1q_f32(outptr, _fsum0); + outptr += 4; + vst1q_f32(outptr, _fsum1); + outptr += 4; + vst1q_f32(outptr, _fsum2); + outptr += 4; + vst1q_f32(outptr, _fsum3); + outptr += 4; + } + for (; jj + 1 < max_jj; jj += 2) + { + float32x4_t _fsum0 = vdupq_n_f32(0.f); + float32x4_t _fsum1 = vdupq_n_f32(0.f); + float32x4_t _fsum2 = vdupq_n_f32(0.f); + float32x4_t _fsum3 = vdupq_n_f32(0.f); + + const signed char* pA = pA0_block; + const float* pA_descales = pA_descales0_block; + for (int k = 0; k < K; k += block_size) + { + int32x4_t _sum0 = vdupq_n_s32(0); + int32x4_t _sum1 = vdupq_n_s32(0); + int32x4_t _sum2 = vdupq_n_s32(0); + int32x4_t _sum3 = vdupq_n_s32(0); + const int max_kk = std::min(K - k, block_size); + int kk = 0; +#if __ARM_FEATURE_MATMUL_INT8 + int32x4_t _msum0 = vdupq_n_s32(0); + int32x4_t _msum1 = vdupq_n_s32(0); + for (; kk + 7 < max_kk; kk += 8) + { + const int8x16_t _b0 = vld1q_s8(pB); + const int8x16_t _a01 = vld1q_s8(pA); + const int8x16_t _a23 = vld1q_s8(pA + 16); + _msum0 = vmmlaq_s32(_msum0, _a01, _b0); + _msum1 = vmmlaq_s32(_msum1, _a23, _b0); + pA += 32; + pB += 16; + } + _sum0 = vcombine_s32(vget_low_s32(_msum0), vdup_n_s32(0)); + _sum1 = vcombine_s32(vget_high_s32(_msum0), vdup_n_s32(0)); + _sum2 = vcombine_s32(vget_low_s32(_msum1), vdup_n_s32(0)); + _sum3 = vcombine_s32(vget_high_s32(_msum1), vdup_n_s32(0)); +#endif // __ARM_FEATURE_MATMUL_INT8 +#if __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk; kk += 4) + { + const int8x16_t _b0 = vcombine_s8(vld1_s8(pB), vdup_n_s8(0)); + const int8x16_t _a = vld1q_s8(pA); + _sum0 = vdotq_laneq_s32(_sum0, _b0, _a, 0); + _sum1 = vdotq_laneq_s32(_sum1, _b0, _a, 1); + _sum2 = vdotq_laneq_s32(_sum2, _b0, _a, 2); + _sum3 = vdotq_laneq_s32(_sum3, _b0, _a, 3); + pA += 16; + pB += 8; + } +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 1 < max_kk; kk += 2) + { + const int8x8_t _b0 = vreinterpret_s8_s32(vld1_dup_s32((const int*)(pB))); + const int16x4_t _a = vreinterpret_s16_s8(vld1_s8(pA)); + _sum0 = vaddq_s32(_sum0, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 0))))); + _sum1 = vaddq_s32(_sum1, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 1))))); + _sum2 = vaddq_s32(_sum2, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 2))))); + _sum3 = vaddq_s32(_sum3, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 3))))); + pA += 8; + pB += 4; + } + if (kk < max_kk) + { + const int8x8_t _b = vreinterpret_s8_s16(vld1_dup_s16((const short*)(pB))); + const int8x8_t _a = vreinterpret_s8_s32(vld1_lane_s32((const int*)pA, vdup_n_s32(0), 0)); + const int16x8_t _p0 = vmull_s8(_b, vdup_lane_s8(_a, 0)); + const int16x8_t _p1 = vmull_s8(_b, vdup_lane_s8(_a, 1)); + const int16x8_t _p2 = vmull_s8(_b, vdup_lane_s8(_a, 2)); + const int16x8_t _p3 = vmull_s8(_b, vdup_lane_s8(_a, 3)); + _sum0 = vaddq_s32(_sum0, vmovl_s16(vget_low_s16(_p0))); + _sum1 = vaddq_s32(_sum1, vmovl_s16(vget_low_s16(_p1))); + _sum2 = vaddq_s32(_sum2, vmovl_s16(vget_low_s16(_p2))); + _sum3 = vaddq_s32(_sum3, vmovl_s16(vget_low_s16(_p3))); + pA += 4; + pB += 2; + } + + const float32x4_t _bd0 = vcombine_f32(vld1_f32(pB_descales), vdup_n_f32(0.f)); + const float32x4_t _ad = vld1q_f32(pA_descales); + _fsum0 = vmlaq_f32(_fsum0, vcvtq_f32_s32(_sum0), vmulq_n_f32(_bd0, vgetq_lane_f32(_ad, 0))); + _fsum1 = vmlaq_f32(_fsum1, vcvtq_f32_s32(_sum1), vmulq_n_f32(_bd0, vgetq_lane_f32(_ad, 1))); + _fsum2 = vmlaq_f32(_fsum2, vcvtq_f32_s32(_sum2), vmulq_n_f32(_bd0, vgetq_lane_f32(_ad, 2))); + _fsum3 = vmlaq_f32(_fsum3, vcvtq_f32_s32(_sum3), vmulq_n_f32(_bd0, vgetq_lane_f32(_ad, 3))); + + pA_descales += 4; + pB_descales += 2; + } + + vst1_f32(outptr, vget_low_f32(_fsum0)); + outptr += 2; + vst1_f32(outptr, vget_low_f32(_fsum1)); + outptr += 2; + vst1_f32(outptr, vget_low_f32(_fsum2)); + outptr += 2; + vst1_f32(outptr, vget_low_f32(_fsum3)); + outptr += 2; + } + for (; jj < max_jj; jj++) + { + float32x4_t _fsum0 = vdupq_n_f32(0.f); + float32x4_t _fsum1 = vdupq_n_f32(0.f); + float32x4_t _fsum2 = vdupq_n_f32(0.f); + float32x4_t _fsum3 = vdupq_n_f32(0.f); + + const signed char* pA = pA0_block; + const float* pA_descales = pA_descales0_block; + for (int k = 0; k < K; k += block_size) + { + int32x4_t _sum0 = vdupq_n_s32(0); + int32x4_t _sum1 = vdupq_n_s32(0); + int32x4_t _sum2 = vdupq_n_s32(0); + int32x4_t _sum3 = vdupq_n_s32(0); + const int max_kk = std::min(K - k, block_size); + int kk = 0; +#if __ARM_FEATURE_MATMUL_INT8 + for (; kk + 7 < max_kk; kk += 8) + { + const int8x16_t _a01 = vld1q_s8(pA); + const int8x16_t _a23 = vld1q_s8(pA + 16); + const int8x16_t _b0 = vcombine_s8(vreinterpret_s8_s32(vld1_dup_s32((const int*)pB)), vdup_n_s8(0)); + const int8x16_t _b1 = vcombine_s8(vreinterpret_s8_s32(vld1_dup_s32((const int*)(pB + 4))), vdup_n_s8(0)); + _sum0 = vdotq_laneq_s32(_sum0, _b0, _a01, 0); + _sum0 = vdotq_laneq_s32(_sum0, _b1, _a01, 1); + _sum1 = vdotq_laneq_s32(_sum1, _b0, _a01, 2); + _sum1 = vdotq_laneq_s32(_sum1, _b1, _a01, 3); + _sum2 = vdotq_laneq_s32(_sum2, _b0, _a23, 0); + _sum2 = vdotq_laneq_s32(_sum2, _b1, _a23, 1); + _sum3 = vdotq_laneq_s32(_sum3, _b0, _a23, 2); + _sum3 = vdotq_laneq_s32(_sum3, _b1, _a23, 3); + pA += 32; + pB += 8; + } +#endif // __ARM_FEATURE_MATMUL_INT8 +#if __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk; kk += 4) + { + const int8x16_t _b0 = vcombine_s8(vreinterpret_s8_s32(vld1_dup_s32((const int*)pB)), vdup_n_s8(0)); + const int8x16_t _a = vld1q_s8(pA); + _sum0 = vdotq_laneq_s32(_sum0, _b0, _a, 0); + _sum1 = vdotq_laneq_s32(_sum1, _b0, _a, 1); + _sum2 = vdotq_laneq_s32(_sum2, _b0, _a, 2); + _sum3 = vdotq_laneq_s32(_sum3, _b0, _a, 3); + pA += 16; + pB += 4; + } +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 1 < max_kk; kk += 2) + { + const int8x8_t _b0 = vreinterpret_s8_s16(vld1_dup_s16((const short*)(pB))); + const int16x4_t _a = vreinterpret_s16_s8(vld1_s8(pA)); + _sum0 = vaddq_s32(_sum0, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 0))))); + _sum1 = vaddq_s32(_sum1, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 1))))); + _sum2 = vaddq_s32(_sum2, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 2))))); + _sum3 = vaddq_s32(_sum3, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 3))))); + pA += 8; + pB += 2; + } + if (kk < max_kk) + { + const int8x8_t _b = vset_lane_s8(pB[0], vdup_n_s8(0), 0); + const int8x8_t _a = vreinterpret_s8_s32(vld1_lane_s32((const int*)pA, vdup_n_s32(0), 0)); + const int16x8_t _p0 = vmull_s8(_b, vdup_lane_s8(_a, 0)); + const int16x8_t _p1 = vmull_s8(_b, vdup_lane_s8(_a, 1)); + const int16x8_t _p2 = vmull_s8(_b, vdup_lane_s8(_a, 2)); + const int16x8_t _p3 = vmull_s8(_b, vdup_lane_s8(_a, 3)); + _sum0 = vaddq_s32(_sum0, vmovl_s16(vget_low_s16(_p0))); + _sum1 = vaddq_s32(_sum1, vmovl_s16(vget_low_s16(_p1))); + _sum2 = vaddq_s32(_sum2, vmovl_s16(vget_low_s16(_p2))); + _sum3 = vaddq_s32(_sum3, vmovl_s16(vget_low_s16(_p3))); + pA += 4; + pB += 1; + } + + const float32x4_t _bd0 = vdupq_n_f32(pB_descales[0]); + const float32x4_t _ad = vld1q_f32(pA_descales); + _fsum0 = vmlaq_f32(_fsum0, vcvtq_f32_s32(_sum0), vmulq_n_f32(_bd0, vgetq_lane_f32(_ad, 0))); + _fsum1 = vmlaq_f32(_fsum1, vcvtq_f32_s32(_sum1), vmulq_n_f32(_bd0, vgetq_lane_f32(_ad, 1))); + _fsum2 = vmlaq_f32(_fsum2, vcvtq_f32_s32(_sum2), vmulq_n_f32(_bd0, vgetq_lane_f32(_ad, 2))); + _fsum3 = vmlaq_f32(_fsum3, vcvtq_f32_s32(_sum3), vmulq_n_f32(_bd0, vgetq_lane_f32(_ad, 3))); + + pA_descales += 4; + pB_descales += 1; + } + + vst1q_lane_f32(outptr, _fsum0, 0); + outptr++; + vst1q_lane_f32(outptr, _fsum1, 0); + outptr++; + vst1q_lane_f32(outptr, _fsum2, 0); + outptr++; + vst1q_lane_f32(outptr, _fsum3, 0); + outptr++; + } + } + for (; ii + 1 < max_ii; ii += 2) + { + const signed char* pA0_block = pAT + ii * A_hstep; + const float* pA_descales0_block = pAT_descales + ii * A_descales_hstep; + const signed char* pB = pBT; + const float* pB_descales = pBT_descales; + + int jj = 0; +#if __aarch64__ + for (; jj + 7 < max_jj; jj += 8) + { + const signed char* pB0 = pB; + const signed char* pB1 = pB + 4 * K; + const float* pB_descales0 = pB_descales; + const float* pB_descales1 = pB_descales + 4 * ((K + block_size - 1) / block_size); + float32x4_t _fsum0 = vdupq_n_f32(0.f); + float32x4_t _fsum1 = vdupq_n_f32(0.f); + float32x4_t _fsum2 = vdupq_n_f32(0.f); + float32x4_t _fsum3 = vdupq_n_f32(0.f); + + const signed char* pA = pA0_block; + const float* pA_descales = pA_descales0_block; + for (int k = 0; k < K; k += block_size) + { + int32x4_t _sum0 = vdupq_n_s32(0); + int32x4_t _sum1 = vdupq_n_s32(0); + int32x4_t _sum2 = vdupq_n_s32(0); + int32x4_t _sum3 = vdupq_n_s32(0); + const int max_kk = std::min(K - k, block_size); + int kk = 0; +#if __ARM_FEATURE_MATMUL_INT8 + int32x4_t _msum0 = vdupq_n_s32(0); + int32x4_t _msum1 = vdupq_n_s32(0); + int32x4_t _msum2 = vdupq_n_s32(0); + int32x4_t _msum3 = vdupq_n_s32(0); + for (; kk + 7 < max_kk; kk += 8) + { + const int8x16_t _b0 = vld1q_s8(pB0); + const int8x16_t _b1 = vld1q_s8(pB0 + 16); + const int8x16_t _b2 = vld1q_s8(pB1); + const int8x16_t _b3 = vld1q_s8(pB1 + 16); + const int8x16_t _a0 = vld1q_s8(pA); + _msum0 = vmmlaq_s32(_msum0, _a0, _b0); + _msum1 = vmmlaq_s32(_msum1, _a0, _b1); + _msum2 = vmmlaq_s32(_msum2, _a0, _b2); + _msum3 = vmmlaq_s32(_msum3, _a0, _b3); + pA += 16; + pB0 += 32; + pB1 += 32; + } + _sum0 = vcombine_s32(vget_low_s32(_msum0), vget_low_s32(_msum1)); + _sum1 = vcombine_s32(vget_low_s32(_msum2), vget_low_s32(_msum3)); + _sum2 = vcombine_s32(vget_high_s32(_msum0), vget_high_s32(_msum1)); + _sum3 = vcombine_s32(vget_high_s32(_msum2), vget_high_s32(_msum3)); +#endif // __ARM_FEATURE_MATMUL_INT8 +#if __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk; kk += 4) + { + const int8x16_t _b0 = vld1q_s8(pB0); + const int8x16_t _b1 = vld1q_s8(pB1); + const int8x16_t _a = vcombine_s8(vld1_s8(pA), vdup_n_s8(0)); + _sum0 = vdotq_laneq_s32(_sum0, _b0, _a, 0); + _sum1 = vdotq_laneq_s32(_sum1, _b1, _a, 0); + _sum2 = vdotq_laneq_s32(_sum2, _b0, _a, 1); + _sum3 = vdotq_laneq_s32(_sum3, _b1, _a, 1); + pA += 8; + pB0 += 16; + pB1 += 16; + } +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 1 < max_kk; kk += 2) + { + const int8x16_t _b = vcombine_s8(vld1_s8(pB0), vld1_s8(pB1)); + const int8x8_t _b0 = vget_low_s8(_b); + const int8x8_t _b1 = vget_high_s8(_b); + const int16x4_t _a = vreinterpret_s16_s32(vld1_lane_s32((const int*)pA, vdup_n_s32(0), 0)); + _sum0 = vaddq_s32(_sum0, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 0))))); + _sum1 = vaddq_s32(_sum1, vpaddlq_s16(vmull_s8(_b1, vreinterpret_s8_s16(vdup_lane_s16(_a, 0))))); + _sum2 = vaddq_s32(_sum2, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 1))))); + _sum3 = vaddq_s32(_sum3, vpaddlq_s16(vmull_s8(_b1, vreinterpret_s8_s16(vdup_lane_s16(_a, 1))))); + pA += 4; + pB0 += 8; + pB1 += 8; + } + if (kk < max_kk) + { + const int8x8_t _b = vreinterpret_s8_s32(vld1_lane_s32((const int*)pB1, vld1_dup_s32((const int*)pB0), 1)); + const int8x8_t _a = vreinterpret_s8_s16(vld1_lane_s16((const short*)pA, vdup_n_s16(0), 0)); + const int16x8_t _p0 = vmull_s8(_b, vdup_lane_s8(_a, 0)); + const int16x8_t _p1 = vmull_s8(_b, vdup_lane_s8(_a, 1)); + _sum0 = vaddq_s32(_sum0, vmovl_s16(vget_low_s16(_p0))); + _sum1 = vaddq_s32(_sum1, vmovl_s16(vget_high_s16(_p0))); + _sum2 = vaddq_s32(_sum2, vmovl_s16(vget_low_s16(_p1))); + _sum3 = vaddq_s32(_sum3, vmovl_s16(vget_high_s16(_p1))); + pA += 2; + pB0 += 4; + pB1 += 4; + } + + const float32x4_t _bd0 = vld1q_f32(pB_descales0); + const float32x4_t _bd1 = vld1q_f32(pB_descales1); + const float32x2_t _ad = vld1_f32(pA_descales); + _fsum0 = vmlaq_f32(_fsum0, vcvtq_f32_s32(_sum0), vmulq_lane_f32(_bd0, _ad, 0)); + _fsum1 = vmlaq_f32(_fsum1, vcvtq_f32_s32(_sum1), vmulq_lane_f32(_bd1, _ad, 0)); + _fsum2 = vmlaq_f32(_fsum2, vcvtq_f32_s32(_sum2), vmulq_lane_f32(_bd0, _ad, 1)); + _fsum3 = vmlaq_f32(_fsum3, vcvtq_f32_s32(_sum3), vmulq_lane_f32(_bd1, _ad, 1)); + + pA_descales += 2; + pB_descales0 += 4; + pB_descales1 += 4; + } + + pB = pB1; + pB_descales = pB_descales1; + + vst1q_f32(outptr, _fsum0); + outptr += 4; + vst1q_f32(outptr, _fsum1); + outptr += 4; + vst1q_f32(outptr, _fsum2); + outptr += 4; + vst1q_f32(outptr, _fsum3); + outptr += 4; + } +#endif // __aarch64__ + for (; jj + 3 < max_jj; jj += 4) + { + float32x4_t _fsum0 = vdupq_n_f32(0.f); + float32x4_t _fsum1 = vdupq_n_f32(0.f); + + const signed char* pA = pA0_block; + const float* pA_descales = pA_descales0_block; + for (int k = 0; k < K; k += block_size) + { + int32x4_t _sum0 = vdupq_n_s32(0); + int32x4_t _sum1 = vdupq_n_s32(0); + const int max_kk = std::min(K - k, block_size); + int kk = 0; +#if __ARM_FEATURE_MATMUL_INT8 + int32x4_t _msum0 = vdupq_n_s32(0); + int32x4_t _msum1 = vdupq_n_s32(0); + for (; kk + 7 < max_kk; kk += 8) + { + const int8x16_t _b0 = vld1q_s8(pB + 0); + const int8x16_t _b1 = vld1q_s8(pB + 16); + const int8x16_t _a0 = vld1q_s8(pA); + _msum0 = vmmlaq_s32(_msum0, _a0, _b0); + _msum1 = vmmlaq_s32(_msum1, _a0, _b1); + pA += 16; + pB += 32; + } + _sum0 = vcombine_s32(vget_low_s32(_msum0), vget_low_s32(_msum1)); + _sum1 = vcombine_s32(vget_high_s32(_msum0), vget_high_s32(_msum1)); +#endif // __ARM_FEATURE_MATMUL_INT8 +#if __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk; kk += 4) + { + const int8x16_t _b0 = vld1q_s8(pB); + const int8x16_t _a = vcombine_s8(vld1_s8(pA), vdup_n_s8(0)); + _sum0 = vdotq_laneq_s32(_sum0, _b0, _a, 0); + _sum1 = vdotq_laneq_s32(_sum1, _b0, _a, 1); + pA += 8; + pB += 16; + } +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 1 < max_kk; kk += 2) + { + const int8x8_t _b0 = vld1_s8(pB); + const int16x4_t _a = vreinterpret_s16_s32(vld1_lane_s32((const int*)pA, vdup_n_s32(0), 0)); + _sum0 = vaddq_s32(_sum0, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 0))))); + _sum1 = vaddq_s32(_sum1, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 1))))); + pA += 4; + pB += 8; + } + if (kk < max_kk) + { + const int8x8_t _b = vreinterpret_s8_s32(vld1_dup_s32((const int*)(pB))); + const int8x8_t _a = vreinterpret_s8_s16(vld1_lane_s16((const short*)pA, vdup_n_s16(0), 0)); + const int16x8_t _p0 = vmull_s8(_b, vdup_lane_s8(_a, 0)); + const int16x8_t _p1 = vmull_s8(_b, vdup_lane_s8(_a, 1)); + _sum0 = vaddq_s32(_sum0, vmovl_s16(vget_low_s16(_p0))); + _sum1 = vaddq_s32(_sum1, vmovl_s16(vget_low_s16(_p1))); + pA += 2; + pB += 4; + } + + const float32x4_t _bd0 = vld1q_f32(pB_descales); + const float32x2_t _ad = vld1_f32(pA_descales); + _fsum0 = vmlaq_f32(_fsum0, vcvtq_f32_s32(_sum0), vmulq_lane_f32(_bd0, _ad, 0)); + _fsum1 = vmlaq_f32(_fsum1, vcvtq_f32_s32(_sum1), vmulq_lane_f32(_bd0, _ad, 1)); + + pA_descales += 2; + pB_descales += 4; + } + + vst1q_f32(outptr, _fsum0); + outptr += 4; + vst1q_f32(outptr, _fsum1); + outptr += 4; + } + for (; jj + 1 < max_jj; jj += 2) + { + float32x4_t _fsum0 = vdupq_n_f32(0.f); + float32x4_t _fsum1 = vdupq_n_f32(0.f); + + const signed char* pA = pA0_block; + const float* pA_descales = pA_descales0_block; + for (int k = 0; k < K; k += block_size) + { + int32x4_t _sum0 = vdupq_n_s32(0); + int32x4_t _sum1 = vdupq_n_s32(0); + const int max_kk = std::min(K - k, block_size); + int kk = 0; +#if __ARM_FEATURE_MATMUL_INT8 + int32x4_t _msum0 = vdupq_n_s32(0); + for (; kk + 7 < max_kk; kk += 8) + { + const int8x16_t _b0 = vld1q_s8(pB + 0); + const int8x16_t _a0 = vld1q_s8(pA); + _msum0 = vmmlaq_s32(_msum0, _a0, _b0); + pA += 16; + pB += 16; + } + _sum0 = vcombine_s32(vget_low_s32(_msum0), vdup_n_s32(0)); + _sum1 = vcombine_s32(vget_high_s32(_msum0), vdup_n_s32(0)); +#endif // __ARM_FEATURE_MATMUL_INT8 +#if __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk; kk += 4) + { + const int8x16_t _b0 = vcombine_s8(vld1_s8(pB), vdup_n_s8(0)); + const int8x16_t _a = vcombine_s8(vld1_s8(pA), vdup_n_s8(0)); + _sum0 = vdotq_laneq_s32(_sum0, _b0, _a, 0); + _sum1 = vdotq_laneq_s32(_sum1, _b0, _a, 1); + pA += 8; + pB += 8; + } +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 1 < max_kk; kk += 2) + { + const int8x8_t _b0 = vreinterpret_s8_s32(vld1_dup_s32((const int*)(pB))); + const int16x4_t _a = vreinterpret_s16_s32(vld1_lane_s32((const int*)pA, vdup_n_s32(0), 0)); + _sum0 = vaddq_s32(_sum0, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 0))))); + _sum1 = vaddq_s32(_sum1, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 1))))); + pA += 4; + pB += 4; + } + if (kk < max_kk) + { + const int8x8_t _b = vreinterpret_s8_s16(vld1_dup_s16((const short*)(pB))); + const int8x8_t _a = vreinterpret_s8_s16(vld1_lane_s16((const short*)pA, vdup_n_s16(0), 0)); + const int16x8_t _p0 = vmull_s8(_b, vdup_lane_s8(_a, 0)); + const int16x8_t _p1 = vmull_s8(_b, vdup_lane_s8(_a, 1)); + _sum0 = vaddq_s32(_sum0, vmovl_s16(vget_low_s16(_p0))); + _sum1 = vaddq_s32(_sum1, vmovl_s16(vget_low_s16(_p1))); + pA += 2; + pB += 2; + } + + const float32x4_t _bd0 = vcombine_f32(vld1_f32(pB_descales), vdup_n_f32(0.f)); + const float32x2_t _ad = vld1_f32(pA_descales); + _fsum0 = vmlaq_f32(_fsum0, vcvtq_f32_s32(_sum0), vmulq_lane_f32(_bd0, _ad, 0)); + _fsum1 = vmlaq_f32(_fsum1, vcvtq_f32_s32(_sum1), vmulq_lane_f32(_bd0, _ad, 1)); + + pA_descales += 2; + pB_descales += 2; + } + + vst1_f32(outptr, vget_low_f32(_fsum0)); + outptr += 2; + vst1_f32(outptr, vget_low_f32(_fsum1)); + outptr += 2; + } + for (; jj < max_jj; jj++) + { + float32x4_t _fsum0 = vdupq_n_f32(0.f); + float32x4_t _fsum1 = vdupq_n_f32(0.f); + + const signed char* pA = pA0_block; + const float* pA_descales = pA_descales0_block; + for (int k = 0; k < K; k += block_size) + { + int32x4_t _sum0 = vdupq_n_s32(0); + int32x4_t _sum1 = vdupq_n_s32(0); + const int max_kk = std::min(K - k, block_size); + int kk = 0; +#if __ARM_FEATURE_MATMUL_INT8 + for (; kk + 7 < max_kk; kk += 8) + { + const int8x16_t _a = vld1q_s8(pA); + const int8x16_t _b0 = vcombine_s8(vreinterpret_s8_s32(vld1_dup_s32((const int*)pB)), vdup_n_s8(0)); + const int8x16_t _b1 = vcombine_s8(vreinterpret_s8_s32(vld1_dup_s32((const int*)(pB + 4))), vdup_n_s8(0)); + _sum0 = vdotq_laneq_s32(_sum0, _b0, _a, 0); + _sum0 = vdotq_laneq_s32(_sum0, _b1, _a, 1); + _sum1 = vdotq_laneq_s32(_sum1, _b0, _a, 2); + _sum1 = vdotq_laneq_s32(_sum1, _b1, _a, 3); + pA += 16; + pB += 8; + } +#endif // __ARM_FEATURE_MATMUL_INT8 +#if __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk; kk += 4) + { + const int8x16_t _b0 = vcombine_s8(vreinterpret_s8_s32(vld1_dup_s32((const int*)pB)), vdup_n_s8(0)); + const int8x16_t _a = vcombine_s8(vld1_s8(pA), vdup_n_s8(0)); + _sum0 = vdotq_laneq_s32(_sum0, _b0, _a, 0); + _sum1 = vdotq_laneq_s32(_sum1, _b0, _a, 1); + pA += 8; + pB += 4; + } +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 1 < max_kk; kk += 2) + { + const int8x8_t _b0 = vreinterpret_s8_s16(vld1_dup_s16((const short*)(pB))); + const int16x4_t _a = vreinterpret_s16_s32(vld1_lane_s32((const int*)pA, vdup_n_s32(0), 0)); + _sum0 = vaddq_s32(_sum0, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 0))))); + _sum1 = vaddq_s32(_sum1, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 1))))); + pA += 4; + pB += 2; + } + if (kk < max_kk) + { + const int8x8_t _b = vset_lane_s8(pB[0], vdup_n_s8(0), 0); + const int8x8_t _a = vreinterpret_s8_s16(vld1_lane_s16((const short*)pA, vdup_n_s16(0), 0)); + const int16x8_t _p0 = vmull_s8(_b, vdup_lane_s8(_a, 0)); + const int16x8_t _p1 = vmull_s8(_b, vdup_lane_s8(_a, 1)); + _sum0 = vaddq_s32(_sum0, vmovl_s16(vget_low_s16(_p0))); + _sum1 = vaddq_s32(_sum1, vmovl_s16(vget_low_s16(_p1))); + pA += 2; + pB += 1; + } + + const float32x4_t _bd0 = vdupq_n_f32(pB_descales[0]); + const float32x2_t _ad = vld1_f32(pA_descales); + _fsum0 = vmlaq_f32(_fsum0, vcvtq_f32_s32(_sum0), vmulq_lane_f32(_bd0, _ad, 0)); + _fsum1 = vmlaq_f32(_fsum1, vcvtq_f32_s32(_sum1), vmulq_lane_f32(_bd0, _ad, 1)); + + pA_descales += 2; + pB_descales += 1; + } + + vst1q_lane_f32(outptr, _fsum0, 0); + outptr++; + vst1q_lane_f32(outptr, _fsum1, 0); + outptr++; + } + } + for (; ii < max_ii; ii++) + { + const signed char* pA0_block = pAT + ii * A_hstep; + const float* pA_descales0_block = pAT_descales + ii * A_descales_hstep; + const signed char* pB = pBT; + const float* pB_descales = pBT_descales; + + int jj = 0; +#if __aarch64__ + for (; jj + 7 < max_jj; jj += 8) + { + const signed char* pB0 = pB; + const signed char* pB1 = pB + 4 * K; + const float* pB_descales0 = pB_descales; + const float* pB_descales1 = pB_descales + 4 * ((K + block_size - 1) / block_size); + float32x4_t _fsum0 = vdupq_n_f32(0.f); + float32x4_t _fsum1 = vdupq_n_f32(0.f); + + const signed char* pA0 = pA0_block; + const float* pA_descales0 = pA_descales0_block; + for (int k = 0; k < K; k += block_size) + { + int32x4_t _sum0 = vdupq_n_s32(0); + int32x4_t _sum1 = vdupq_n_s32(0); + const int max_kk = std::min(K - k, block_size); + int kk = 0; +#if __ARM_FEATURE_MATMUL_INT8 + int32x4_t _msum0 = vdupq_n_s32(0); + int32x4_t _msum1 = vdupq_n_s32(0); + int32x4_t _msum2 = vdupq_n_s32(0); + int32x4_t _msum3 = vdupq_n_s32(0); + for (; kk + 7 < max_kk; kk += 8) + { + const int8x16_t _b0 = vld1q_s8(pB0); + const int8x16_t _b1 = vld1q_s8(pB0 + 16); + const int8x16_t _b2 = vld1q_s8(pB1); + const int8x16_t _b3 = vld1q_s8(pB1 + 16); + const int8x16_t _a0 = vcombine_s8(vld1_s8(pA0 + kk), vdup_n_s8(0)); + _msum0 = vmmlaq_s32(_msum0, _a0, _b0); + _msum1 = vmmlaq_s32(_msum1, _a0, _b1); + _msum2 = vmmlaq_s32(_msum2, _a0, _b2); + _msum3 = vmmlaq_s32(_msum3, _a0, _b3); + pB0 += 32; + pB1 += 32; + } + _sum0 = vcombine_s32(vget_low_s32(_msum0), vget_low_s32(_msum1)); + _sum1 = vcombine_s32(vget_low_s32(_msum2), vget_low_s32(_msum3)); +#endif // __ARM_FEATURE_MATMUL_INT8 +#if __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk; kk += 4) + { + const int8x16_t _b0 = vld1q_s8(pB0); + const int8x16_t _b1 = vld1q_s8(pB1); + const int8x16_t _a0 = vreinterpretq_s8_s32(vdupq_lane_s32(vld1_lane_s32((const int*)(pA0 + kk), vdup_n_s32(0), 0), 0)); + _sum0 = vdotq_s32(_sum0, _b0, _a0); + _sum1 = vdotq_s32(_sum1, _b1, _a0); + pB0 += 16; + pB1 += 16; + } +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 1 < max_kk; kk += 2) + { + const int8x16_t _b = vcombine_s8(vld1_s8(pB0), vld1_s8(pB1)); + const int8x8_t _b0 = vget_low_s8(_b); + const int8x8_t _b1 = vget_high_s8(_b); + const int8x8_t _a0 = vreinterpret_s8_s16(vdup_lane_s16(vld1_lane_s16((const short*)(pA0 + kk), vdup_n_s16(0), 0), 0)); + _sum0 = vaddq_s32(_sum0, vpaddlq_s16(vmull_s8(_b0, _a0))); + _sum1 = vaddq_s32(_sum1, vpaddlq_s16(vmull_s8(_b1, _a0))); + pB0 += 8; + pB1 += 8; + } + if (kk < max_kk) + { + const int8x8_t _b = vreinterpret_s8_s32(vld1_lane_s32((const int*)pB1, vld1_dup_s32((const int*)pB0), 1)); + const int8x8_t _a0 = vld1_lane_s8(pA0 + kk, vdup_n_s8(0), 0); + const int16x8_t _p0 = vmull_s8(_b, vdup_lane_s8(_a0, 0)); + _sum0 = vaddq_s32(_sum0, vmovl_s16(vget_low_s16(_p0))); + _sum1 = vaddq_s32(_sum1, vmovl_s16(vget_high_s16(_p0))); + pB0 += 4; + pB1 += 4; + } + + const float32x4_t _bd0 = vld1q_f32(pB_descales0); + const float32x4_t _bd1 = vld1q_f32(pB_descales1); + const float _ad0 = pA_descales0[0]; + _fsum0 = vmlaq_f32(_fsum0, vcvtq_f32_s32(_sum0), vmulq_n_f32(_bd0, _ad0)); + _fsum1 = vmlaq_f32(_fsum1, vcvtq_f32_s32(_sum1), vmulq_n_f32(_bd1, _ad0)); + + pA0 += max_kk; + pA_descales0++; + pB_descales0 += 4; + pB_descales1 += 4; + } + + pB = pB1; + pB_descales = pB_descales1; + + vst1q_f32(outptr, _fsum0); + outptr += 4; + vst1q_f32(outptr, _fsum1); + outptr += 4; + } +#endif // __aarch64__ + for (; jj + 3 < max_jj; jj += 4) + { + float32x4_t _fsum0 = vdupq_n_f32(0.f); + + const signed char* pA0 = pA0_block; + const float* pA_descales0 = pA_descales0_block; + for (int k = 0; k < K; k += block_size) + { + int32x4_t _sum0 = vdupq_n_s32(0); + const int max_kk = std::min(K - k, block_size); + int kk = 0; +#if __ARM_FEATURE_MATMUL_INT8 + int32x4_t _msum0 = vdupq_n_s32(0); + int32x4_t _msum1 = vdupq_n_s32(0); + for (; kk + 7 < max_kk; kk += 8) + { + const int8x16_t _b0 = vld1q_s8(pB + 0); + const int8x16_t _b1 = vld1q_s8(pB + 16); + const int8x16_t _a0 = vcombine_s8(vld1_s8(pA0 + kk), vdup_n_s8(0)); + _msum0 = vmmlaq_s32(_msum0, _a0, _b0); + _msum1 = vmmlaq_s32(_msum1, _a0, _b1); + pB += 32; + } + _sum0 = vcombine_s32(vget_low_s32(_msum0), vget_low_s32(_msum1)); +#endif // __ARM_FEATURE_MATMUL_INT8 +#if __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk; kk += 4) + { + const int8x16_t _b0 = vld1q_s8(pB); + const int8x16_t _a0 = vreinterpretq_s8_s32(vdupq_lane_s32(vld1_lane_s32((const int*)(pA0 + kk), vdup_n_s32(0), 0), 0)); + _sum0 = vdotq_s32(_sum0, _b0, _a0); + pB += 16; + } +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 1 < max_kk; kk += 2) + { + const int8x8_t _b0 = vld1_s8(pB); + const int8x8_t _a0 = vreinterpret_s8_s16(vdup_lane_s16(vld1_lane_s16((const short*)(pA0 + kk), vdup_n_s16(0), 0), 0)); + _sum0 = vaddq_s32(_sum0, vpaddlq_s16(vmull_s8(_b0, _a0))); + pB += 8; + } + if (kk < max_kk) + { + const int8x8_t _b = vreinterpret_s8_s32(vld1_dup_s32((const int*)(pB))); + const int8x8_t _a0 = vld1_lane_s8(pA0 + kk, vdup_n_s8(0), 0); + const int16x8_t _p0 = vmull_s8(_b, vdup_lane_s8(_a0, 0)); + _sum0 = vaddq_s32(_sum0, vmovl_s16(vget_low_s16(_p0))); + pB += 4; + } + + const float32x4_t _bd0 = vld1q_f32(pB_descales); + const float _ad0 = pA_descales0[0]; + _fsum0 = vmlaq_f32(_fsum0, vcvtq_f32_s32(_sum0), vmulq_n_f32(_bd0, _ad0)); + + pA0 += max_kk; + pA_descales0++; + pB_descales += 4; + } + + vst1q_f32(outptr, _fsum0); + outptr += 4; + } + for (; jj + 1 < max_jj; jj += 2) + { + float32x4_t _fsum0 = vdupq_n_f32(0.f); + + const signed char* pA0 = pA0_block; + const float* pA_descales0 = pA_descales0_block; + for (int k = 0; k < K; k += block_size) + { + int32x4_t _sum0 = vdupq_n_s32(0); + const int max_kk = std::min(K - k, block_size); + int kk = 0; +#if __ARM_FEATURE_MATMUL_INT8 + int32x4_t _msum0 = vdupq_n_s32(0); + for (; kk + 7 < max_kk; kk += 8) + { + const int8x16_t _b0 = vld1q_s8(pB + 0); + const int8x16_t _a0 = vcombine_s8(vld1_s8(pA0 + kk), vdup_n_s8(0)); + _msum0 = vmmlaq_s32(_msum0, _a0, _b0); + pB += 16; + } + _sum0 = vcombine_s32(vget_low_s32(_msum0), vdup_n_s32(0)); +#endif // __ARM_FEATURE_MATMUL_INT8 +#if __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk; kk += 4) + { + const int8x16_t _b0 = vcombine_s8(vld1_s8(pB), vdup_n_s8(0)); + const int8x16_t _a0 = vreinterpretq_s8_s32(vdupq_lane_s32(vld1_lane_s32((const int*)(pA0 + kk), vdup_n_s32(0), 0), 0)); + _sum0 = vdotq_s32(_sum0, _b0, _a0); + pB += 8; + } +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 1 < max_kk; kk += 2) + { + const int8x8_t _b0 = vreinterpret_s8_s32(vld1_dup_s32((const int*)(pB))); + const int8x8_t _a0 = vreinterpret_s8_s16(vdup_lane_s16(vld1_lane_s16((const short*)(pA0 + kk), vdup_n_s16(0), 0), 0)); + _sum0 = vaddq_s32(_sum0, vpaddlq_s16(vmull_s8(_b0, _a0))); + pB += 4; + } + if (kk < max_kk) + { + const int8x8_t _b = vreinterpret_s8_s16(vld1_dup_s16((const short*)(pB))); + const int8x8_t _a0 = vld1_lane_s8(pA0 + kk, vdup_n_s8(0), 0); + const int16x8_t _p0 = vmull_s8(_b, vdup_lane_s8(_a0, 0)); + _sum0 = vaddq_s32(_sum0, vmovl_s16(vget_low_s16(_p0))); + pB += 2; + } + + const float32x4_t _bd0 = vcombine_f32(vld1_f32(pB_descales), vdup_n_f32(0.f)); + const float _ad0 = pA_descales0[0]; + _fsum0 = vmlaq_f32(_fsum0, vcvtq_f32_s32(_sum0), vmulq_n_f32(_bd0, _ad0)); + + pA0 += max_kk; + pA_descales0++; + pB_descales += 2; + } + + vst1_f32(outptr, vget_low_f32(_fsum0)); + outptr += 2; + } + for (; jj < max_jj; jj++) + { + float32x4_t _fsum0 = vdupq_n_f32(0.f); + + const signed char* pA0 = pA0_block; + const float* pA_descales0 = pA_descales0_block; + for (int k = 0; k < K; k += block_size) + { + int32x4_t _sum0 = vdupq_n_s32(0); + const int max_kk = std::min(K - k, block_size); + int kk = 0; +#if __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk; kk += 4) + { + const int8x16_t _b0 = vcombine_s8(vreinterpret_s8_s32(vld1_dup_s32((const int*)pB)), vdup_n_s8(0)); + const int8x16_t _a0 = vreinterpretq_s8_s32(vdupq_lane_s32(vld1_lane_s32((const int*)(pA0 + kk), vdup_n_s32(0), 0), 0)); + _sum0 = vdotq_s32(_sum0, _b0, _a0); + pB += 4; + } +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 1 < max_kk; kk += 2) + { + const int8x8_t _b0 = vreinterpret_s8_s16(vld1_dup_s16((const short*)(pB))); + const int8x8_t _a0 = vreinterpret_s8_s16(vdup_lane_s16(vld1_lane_s16((const short*)(pA0 + kk), vdup_n_s16(0), 0), 0)); + _sum0 = vaddq_s32(_sum0, vpaddlq_s16(vmull_s8(_b0, _a0))); + pB += 2; + } + if (kk < max_kk) + { + const int8x8_t _b = vset_lane_s8(pB[0], vdup_n_s8(0), 0); + const int8x8_t _a0 = vld1_lane_s8(pA0 + kk, vdup_n_s8(0), 0); + const int16x8_t _p0 = vmull_s8(_b, vdup_lane_s8(_a0, 0)); + _sum0 = vaddq_s32(_sum0, vmovl_s16(vget_low_s16(_p0))); + pB += 1; + } + + const float32x4_t _bd0 = vdupq_n_f32(pB_descales[0]); + const float _ad0 = pA_descales0[0]; + _fsum0 = vmlaq_f32(_fsum0, vcvtq_f32_s32(_sum0), vmulq_n_f32(_bd0, _ad0)); + + pA0 += max_kk; + pA_descales0++; + pB_descales += 1; + } + + vst1q_lane_f32(outptr, _fsum0, 0); + outptr++; + } + } +#else +#if __ARM_FEATURE_SIMD32 && NCNN_GNU_INLINE_ASM + for (; ii + 1 < max_ii; ii += 2) + { + const signed char* pA_block = pAT + ii * A_hstep; + const float* pA_descales_block = pAT_descales + ii * A_descales_hstep; + const signed char* pB = pBT; + const float* pB_descales = pBT_descales; + + int jj = 0; + for (; jj + 1 < max_jj; jj += 2) + { + float fsum00 = 0.f; + float fsum01 = 0.f; + float fsum10 = 0.f; + float fsum11 = 0.f; + + const signed char* pA = pA_block; + const float* pA_descales = pA_descales_block; + for (int k = 0; k < K; k += block_size) + { + int sum00 = 0; + int sum01 = 0; + int sum10 = 0; + int sum11 = 0; + const int max_kk = std::min(K - k, block_size); + int kk = 0; + for (; kk + 1 < max_kk; kk += 2) + { +#if __OPTIMIZE__ + asm volatile( + "ldr r2, [%0], #4 \n" + "ldr r4, [%1], #4 \n" + "ror r3, r2, #8 \n" + "ror r5, r4, #8 \n" + "sxtb16 r2, r2 \n" + "sxtb16 r4, r4 \n" + "sxtb16 r3, r3 \n" + "sxtb16 r5, r5 \n" + "smlad %2, r2, r4, %2 \n" + "smlad %3, r3, r4, %3 \n" + "smlad %4, r2, r5, %4 \n" + "smlad %5, r3, r5, %5 \n" + : "=r"(pA), + "=r"(pB), + "=r"(sum00), + "=r"(sum10), + "=r"(sum01), + "=r"(sum11) + : "0"(pA), + "1"(pB), + "2"(sum00), + "3"(sum10), + "4"(sum01), + "5"(sum11) + : "memory", "r2", "r3", "r4", "r5"); +#else + int _pA0 = *((int*)pA); + int _pB0 = *((int*)pB); + int _pA1; + int _pB1; + asm volatile("ror %0, %1, #8" + : "=r"(_pA1) + : "r"(_pA0) + :); + asm volatile("ror %0, %1, #8" + : "=r"(_pB1) + : "r"(_pB0) + :); + asm volatile("sxtb16 %0, %0" + : "=r"(_pA0) + : "0"(_pA0) + :); + asm volatile("sxtb16 %0, %0" + : "=r"(_pA1) + : "0"(_pA1) + :); + asm volatile("sxtb16 %0, %0" + : "=r"(_pB0) + : "0"(_pB0) + :); + asm volatile("sxtb16 %0, %0" + : "=r"(_pB1) + : "0"(_pB1) + :); + asm volatile("smlad %0, %2, %3, %0" + : "=r"(sum00) + : "0"(sum00), "r"(_pA0), "r"(_pB0) + :); + asm volatile("smlad %0, %2, %3, %0" + : "=r"(sum10) + : "0"(sum10), "r"(_pA1), "r"(_pB0) + :); + asm volatile("smlad %0, %2, %3, %0" + : "=r"(sum01) + : "0"(sum01), "r"(_pA0), "r"(_pB1) + :); + asm volatile("smlad %0, %2, %3, %0" + : "=r"(sum11) + : "0"(sum11), "r"(_pA1), "r"(_pB1) + :); + pA += 4; + pB += 4; +#endif + } + if (kk < max_kk) + { + sum00 += pA[0] * pB[0]; + sum10 += pA[1] * pB[0]; + sum01 += pA[0] * pB[1]; + sum11 += pA[1] * pB[1]; + pA += 2; + pB += 2; + } + + const float bd0 = pB_descales[0]; + const float bd1 = pB_descales[1]; + const float ad0 = pA_descales[0]; + const float ad1 = pA_descales[1]; + fsum00 += sum00 * ad0 * bd0; + fsum01 += sum01 * ad0 * bd1; + fsum10 += sum10 * ad1 * bd0; + fsum11 += sum11 * ad1 * bd1; + + pA_descales += 2; + pB_descales += 2; + } + + *outptr++ = fsum00; + *outptr++ = fsum01; + *outptr++ = fsum10; + *outptr++ = fsum11; + } + for (; jj < max_jj; jj++) + { + float fsum00 = 0.f; + float fsum10 = 0.f; + + const signed char* pA = pA_block; + const float* pA_descales = pA_descales_block; + for (int k = 0; k < K; k += block_size) + { + int sum00 = 0; + int sum10 = 0; + const int max_kk = std::min(K - k, block_size); + for (int kk = 0; kk < max_kk; kk++) + { + sum00 += pA[0] * pB[0]; + sum10 += pA[1] * pB[0]; + pA += 2; + pB++; + } + + const float bd0 = pB_descales[0]; + const float ad0 = pA_descales[0]; + const float ad1 = pA_descales[1]; + fsum00 += sum00 * ad0 * bd0; + fsum10 += sum10 * ad1 * bd0; + + pA_descales += 2; + pB_descales++; + } + + *outptr++ = fsum00; + *outptr++ = fsum10; + } + } +#else + for (; ii + 1 < max_ii; ii += 2) + { + const signed char* pA0_block = pAT + ii * A_hstep; + const signed char* pA1_block = pA0_block + A_hstep; + const float* pA_descales0_block = pAT_descales + ii * A_descales_hstep; + const float* pA_descales1_block = pA_descales0_block + A_descales_hstep; + const signed char* pB = pBT; + const float* pB_descales = pBT_descales; + + int jj = 0; + for (; jj + 1 < max_jj; jj += 2) + { + float fsum00 = 0.f; + float fsum01 = 0.f; + float fsum10 = 0.f; + float fsum11 = 0.f; + + const signed char* pA0 = pA0_block; + const signed char* pA1 = pA1_block; + const float* pA_descales0 = pA_descales0_block; + const float* pA_descales1 = pA_descales1_block; + for (int k = 0; k < K; k += block_size) + { + int sum00 = 0; + int sum01 = 0; + int sum10 = 0; + int sum11 = 0; + const int max_kk = std::min(K - k, block_size); + int kk = 0; + for (; kk + 1 < max_kk; kk += 2) + { + const int b00 = pB[0]; + const int b01 = pB[1]; + const int b10 = pB[2]; + const int b11 = pB[3]; + sum00 += pA0[kk] * b00 + pA0[kk + 1] * b01; + sum01 += pA0[kk] * b10 + pA0[kk + 1] * b11; + sum10 += pA1[kk] * b00 + pA1[kk + 1] * b01; + sum11 += pA1[kk] * b10 + pA1[kk + 1] * b11; + pB += 4; + } + if (kk < max_kk) + { + const int b0 = pB[0]; + const int b1 = pB[1]; + sum00 += pA0[kk] * b0; + sum01 += pA0[kk] * b1; + sum10 += pA1[kk] * b0; + sum11 += pA1[kk] * b1; + pB += 2; + } + + const float bd0 = pB_descales[0]; + const float bd1 = pB_descales[1]; + const float ad0 = pA_descales0[0]; + fsum00 += sum00 * ad0 * bd0; + fsum01 += sum01 * ad0 * bd1; + const float ad1 = pA_descales1[0]; + fsum10 += sum10 * ad1 * bd0; + fsum11 += sum11 * ad1 * bd1; + + pA0 += max_kk; + pA1 += max_kk; + pA_descales0++; + pA_descales1++; + pB_descales += 2; + } + + outptr[0] = fsum00; + outptr++; + outptr[0] = fsum01; + outptr++; + outptr[0] = fsum10; + outptr++; + outptr[0] = fsum11; + outptr++; + } + for (; jj < max_jj; jj++) + { + float fsum00 = 0.f; + float fsum10 = 0.f; + + const signed char* pA0 = pA0_block; + const signed char* pA1 = pA1_block; + const float* pA_descales0 = pA_descales0_block; + const float* pA_descales1 = pA_descales1_block; + for (int k = 0; k < K; k += block_size) + { + int sum00 = 0; + int sum10 = 0; + const int max_kk = std::min(K - k, block_size); + int kk = 0; + for (; kk + 1 < max_kk; kk += 2) + { + const int b00 = pB[0]; + const int b01 = pB[1]; + sum00 += pA0[kk] * b00 + pA0[kk + 1] * b01; + sum10 += pA1[kk] * b00 + pA1[kk + 1] * b01; + pB += 2; + } + if (kk < max_kk) + { + const int b0 = pB[0]; + sum00 += pA0[kk] * b0; + sum10 += pA1[kk] * b0; + pB += 1; + } + + const float bd0 = pB_descales[0]; + const float ad0 = pA_descales0[0]; + fsum00 += sum00 * ad0 * bd0; + const float ad1 = pA_descales1[0]; + fsum10 += sum10 * ad1 * bd0; + + pA0 += max_kk; + pA1 += max_kk; + pA_descales0++; + pA_descales1++; + pB_descales += 1; + } + + outptr[0] = fsum00; + outptr++; + outptr[0] = fsum10; + outptr++; + } + } +#endif // __ARM_FEATURE_SIMD32 && NCNN_GNU_INLINE_ASM + for (; ii < max_ii; ii++) + { + const signed char* pA0_block = pAT + ii * A_hstep; + const float* pA_descales0_block = pAT_descales + ii * A_descales_hstep; + const signed char* pB = pBT; + const float* pB_descales = pBT_descales; + + int jj = 0; + for (; jj + 1 < max_jj; jj += 2) + { + float fsum00 = 0.f; + float fsum01 = 0.f; + + const signed char* pA0 = pA0_block; + const float* pA_descales0 = pA_descales0_block; + for (int k = 0; k < K; k += block_size) + { + int sum00 = 0; + int sum01 = 0; + const int max_kk = std::min(K - k, block_size); + int kk = 0; + for (; kk + 1 < max_kk; kk += 2) + { + const int b00 = pB[0]; +#if __ARM_FEATURE_SIMD32 && NCNN_GNU_INLINE_ASM + const int b01 = pB[2]; + const int b10 = pB[1]; +#else + const int b01 = pB[1]; + const int b10 = pB[2]; +#endif + const int b11 = pB[3]; + sum00 += pA0[kk] * b00 + pA0[kk + 1] * b01; + sum01 += pA0[kk] * b10 + pA0[kk + 1] * b11; + pB += 4; + } + if (kk < max_kk) + { + const int b0 = pB[0]; + const int b1 = pB[1]; + sum00 += pA0[kk] * b0; + sum01 += pA0[kk] * b1; + pB += 2; + } + + const float bd0 = pB_descales[0]; + const float bd1 = pB_descales[1]; + const float ad0 = pA_descales0[0]; + fsum00 += sum00 * ad0 * bd0; + fsum01 += sum01 * ad0 * bd1; + + pA0 += max_kk; + pA_descales0++; + pB_descales += 2; + } + + outptr[0] = fsum00; + outptr++; + outptr[0] = fsum01; + outptr++; + } + for (; jj < max_jj; jj++) + { + float fsum00 = 0.f; + + const signed char* pA0 = pA0_block; + const float* pA_descales0 = pA_descales0_block; + for (int k = 0; k < K; k += block_size) + { + int sum00 = 0; + const int max_kk = std::min(K - k, block_size); + int kk = 0; + for (; kk + 1 < max_kk; kk += 2) + { + const int b00 = pB[0]; + const int b01 = pB[1]; + sum00 += pA0[kk] * b00 + pA0[kk + 1] * b01; + pB += 2; + } + if (kk < max_kk) + { + const int b0 = pB[0]; + sum00 += pA0[kk] * b0; + pB += 1; + } + + const float bd0 = pB_descales[0]; + const float ad0 = pA_descales0[0]; + fsum00 += sum00 * ad0 * bd0; + + pA0 += max_kk; + pA_descales0++; + pB_descales += 1; + } + + outptr[0] = fsum00; + outptr++; + } + } +#endif // __ARM_NEON +} + +static void get_optimal_tile_mnk_wq_int8(int M, int N, int K, int constant_TILE_M, int constant_TILE_N, int constant_TILE_K, int& TILE_M, int& TILE_N, int& TILE_K, int nT) +{ + // resolve optimal tile size from cache size + const size_t l2_cache_size = get_cpu_level2_cache_size(); + + if (nT == 0) + nT = get_physical_big_cpu_count(); + + const int tile_size = std::max(1, (int)((float)l2_cache_size / 2 / sizeof(signed char) / std::max(1, K))); + +#if __aarch64__ + TILE_M = M >= nT * 8 ? 8 : M >= nT * 4 ? 4 : M >= nT * 2 ? 2 : 1; + TILE_N = std::max(8, tile_size / 8 * 8); +#elif __ARM_NEON + TILE_M = M >= nT * 4 ? 4 : M >= nT * 2 ? 2 : 1; + TILE_N = std::max(4, tile_size / 4 * 4); +#else + TILE_M = M >= nT * 2 ? 2 : 1; + TILE_N = std::max(2, tile_size / 2 * 2); +#endif + TILE_K = K; + + if (N > 0) + { + const int nn_N = (N + TILE_N - 1) / TILE_N; +#if __aarch64__ + TILE_N = std::min(TILE_N, ((N + nn_N - 1) / nn_N + 7) / 8 * 8); +#elif __ARM_NEON + TILE_N = std::min(TILE_N, ((N + nn_N - 1) / nn_N + 3) / 4 * 4); +#else + TILE_N = std::min(TILE_N, ((N + nn_N - 1) / nn_N + 1) / 2 * 2); +#endif + } + + // always take constant TILE_M/N value when provided + if (constant_TILE_M > 0) + { +#if __aarch64__ + TILE_M = (constant_TILE_M + 7) / 8 * 8; +#elif __ARM_NEON + TILE_M = (constant_TILE_M + 3) / 4 * 4; +#else + TILE_M = (constant_TILE_M + 1) / 2 * 2; +#endif + } + + if (constant_TILE_N > 0) + { +#if __aarch64__ + TILE_N = (constant_TILE_N + 7) / 8 * 8; +#elif __ARM_NEON + TILE_N = (constant_TILE_N + 3) / 4 * 4; +#else + TILE_N = (constant_TILE_N + 1) / 2 * 2; +#endif + } + + // one driver M tile follows the natural producer slab +#if __aarch64__ + TILE_M = std::min(TILE_M, 8); +#elif __ARM_NEON + TILE_M = std::min(TILE_M, 4); +#else + TILE_M = std::min(TILE_M, 2); +#endif + + (void)constant_TILE_K; +} + +static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta) +{ + const float* pC = C; + const float* pp = topT; + const size_t out_hstep = top_blob.dims == 3 ? top_blob.cstep : (size_t)N; + const size_t c_hstep = C.dims == 3 ? C.cstep : (size_t)C.w; + + int ii = 0; +#if __ARM_NEON +#if __aarch64__ + for (; ii + 7 < max_ii; ii += 8) + { + float* outptr0 = (float*)top_blob + (size_t)(i + ii) * out_hstep + j; + float* outptr1 = outptr0 + out_hstep; + float* outptr2 = outptr1 + out_hstep; + float* outptr3 = outptr2 + out_hstep; + float* outptr4 = outptr3 + out_hstep; + float* outptr5 = outptr4 + out_hstep; + float* outptr6 = outptr5 + out_hstep; + float* outptr7 = outptr6 + out_hstep; + pC = (const float*)C; + float c0 = 0.f; + float c1 = 0.f; + float c2 = 0.f; + float c3 = 0.f; + float c4 = 0.f; + float c5 = 0.f; + float c6 = 0.f; + float c7 = 0.f; + if (pC) + { + if (broadcast_type_C == 0) + c0 = c1 = c2 = c3 = c4 = c5 = c6 = c7 = pC[0] * beta; + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + pC = (const float*)C + i + ii; + c0 = pC[0] * beta; + c1 = pC[1] * beta; + c2 = pC[2] * beta; + c3 = pC[3] * beta; + c4 = pC[4] * beta; + c5 = pC[5] * beta; + c6 = pC[6] * beta; + c7 = pC[7] * beta; + } + if (broadcast_type_C == 3) + pC = (const float*)C + (size_t)(i + ii) * c_hstep + j; + if (broadcast_type_C == 4) + pC = (const float*)C + j; + } + + int jj = 0; + for (; jj + 3 < max_jj; jj += 4) + { + float32x4_t _out0 = vld1q_f32(pp); + float32x4_t _out1 = vld1q_f32(pp + 4); + float32x4_t _out2 = vld1q_f32(pp + 8); + float32x4_t _out3 = vld1q_f32(pp + 12); + float32x4_t _out4 = vld1q_f32(pp + 16); + float32x4_t _out5 = vld1q_f32(pp + 20); + float32x4_t _out6 = vld1q_f32(pp + 24); + float32x4_t _out7 = vld1q_f32(pp + 28); + if (pC) + { + if (broadcast_type_C <= 2) + { + _out0 = vaddq_f32(_out0, vdupq_n_f32(c0)); + _out1 = vaddq_f32(_out1, vdupq_n_f32(c1)); + _out2 = vaddq_f32(_out2, vdupq_n_f32(c2)); + _out3 = vaddq_f32(_out3, vdupq_n_f32(c3)); + _out4 = vaddq_f32(_out4, vdupq_n_f32(c4)); + _out5 = vaddq_f32(_out5, vdupq_n_f32(c5)); + _out6 = vaddq_f32(_out6, vdupq_n_f32(c6)); + _out7 = vaddq_f32(_out7, vdupq_n_f32(c7)); + } + if (broadcast_type_C == 3) + { + _out0 = beta == 1.f ? vaddq_f32(_out0, vld1q_f32(pC)) : vmlaq_n_f32(_out0, vld1q_f32(pC), beta); + _out1 = beta == 1.f ? vaddq_f32(_out1, vld1q_f32(pC + c_hstep)) : vmlaq_n_f32(_out1, vld1q_f32(pC + c_hstep), beta); + _out2 = beta == 1.f ? vaddq_f32(_out2, vld1q_f32(pC + c_hstep * 2)) : vmlaq_n_f32(_out2, vld1q_f32(pC + c_hstep * 2), beta); + _out3 = beta == 1.f ? vaddq_f32(_out3, vld1q_f32(pC + c_hstep * 3)) : vmlaq_n_f32(_out3, vld1q_f32(pC + c_hstep * 3), beta); + _out4 = beta == 1.f ? vaddq_f32(_out4, vld1q_f32(pC + c_hstep * 4)) : vmlaq_n_f32(_out4, vld1q_f32(pC + c_hstep * 4), beta); + _out5 = beta == 1.f ? vaddq_f32(_out5, vld1q_f32(pC + c_hstep * 5)) : vmlaq_n_f32(_out5, vld1q_f32(pC + c_hstep * 5), beta); + _out6 = beta == 1.f ? vaddq_f32(_out6, vld1q_f32(pC + c_hstep * 6)) : vmlaq_n_f32(_out6, vld1q_f32(pC + c_hstep * 6), beta); + _out7 = beta == 1.f ? vaddq_f32(_out7, vld1q_f32(pC + c_hstep * 7)) : vmlaq_n_f32(_out7, vld1q_f32(pC + c_hstep * 7), beta); + pC += 4; + } + if (broadcast_type_C == 4) + { + const float32x4_t _c = beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta); + _out0 = vaddq_f32(_out0, _c); + _out1 = vaddq_f32(_out1, _c); + _out2 = vaddq_f32(_out2, _c); + _out3 = vaddq_f32(_out3, _c); + _out4 = vaddq_f32(_out4, _c); + _out5 = vaddq_f32(_out5, _c); + _out6 = vaddq_f32(_out6, _c); + _out7 = vaddq_f32(_out7, _c); + pC += 4; + } + } + if (alpha != 1.f) + { + _out0 = vmulq_n_f32(_out0, alpha); + _out1 = vmulq_n_f32(_out1, alpha); + _out2 = vmulq_n_f32(_out2, alpha); + _out3 = vmulq_n_f32(_out3, alpha); + _out4 = vmulq_n_f32(_out4, alpha); + _out5 = vmulq_n_f32(_out5, alpha); + _out6 = vmulq_n_f32(_out6, alpha); + _out7 = vmulq_n_f32(_out7, alpha); + } + vst1q_f32(outptr0, _out0); + vst1q_f32(outptr1, _out1); + vst1q_f32(outptr2, _out2); + vst1q_f32(outptr3, _out3); + vst1q_f32(outptr4, _out4); + vst1q_f32(outptr5, _out5); + vst1q_f32(outptr6, _out6); + vst1q_f32(outptr7, _out7); + outptr0 += 4; + outptr1 += 4; + outptr2 += 4; + outptr3 += 4; + outptr4 += 4; + outptr5 += 4; + outptr6 += 4; + outptr7 += 4; + pp += 32; + } + for (; jj + 1 < max_jj; jj += 2) + { + float32x2_t _out0 = vld1_f32(pp); + float32x2_t _out1 = vld1_f32(pp + 2); + float32x2_t _out2 = vld1_f32(pp + 4); + float32x2_t _out3 = vld1_f32(pp + 6); + float32x2_t _out4 = vld1_f32(pp + 8); + float32x2_t _out5 = vld1_f32(pp + 10); + float32x2_t _out6 = vld1_f32(pp + 12); + float32x2_t _out7 = vld1_f32(pp + 14); + if (pC) + { + if (broadcast_type_C <= 2) + { + _out0 = vadd_f32(_out0, vdup_n_f32(c0)); + _out1 = vadd_f32(_out1, vdup_n_f32(c1)); + _out2 = vadd_f32(_out2, vdup_n_f32(c2)); + _out3 = vadd_f32(_out3, vdup_n_f32(c3)); + _out4 = vadd_f32(_out4, vdup_n_f32(c4)); + _out5 = vadd_f32(_out5, vdup_n_f32(c5)); + _out6 = vadd_f32(_out6, vdup_n_f32(c6)); + _out7 = vadd_f32(_out7, vdup_n_f32(c7)); + } + if (broadcast_type_C == 3) + { + _out0 = beta == 1.f ? vadd_f32(_out0, vld1_f32(pC)) : vmla_n_f32(_out0, vld1_f32(pC), beta); + _out1 = beta == 1.f ? vadd_f32(_out1, vld1_f32(pC + c_hstep)) : vmla_n_f32(_out1, vld1_f32(pC + c_hstep), beta); + _out2 = beta == 1.f ? vadd_f32(_out2, vld1_f32(pC + c_hstep * 2)) : vmla_n_f32(_out2, vld1_f32(pC + c_hstep * 2), beta); + _out3 = beta == 1.f ? vadd_f32(_out3, vld1_f32(pC + c_hstep * 3)) : vmla_n_f32(_out3, vld1_f32(pC + c_hstep * 3), beta); + _out4 = beta == 1.f ? vadd_f32(_out4, vld1_f32(pC + c_hstep * 4)) : vmla_n_f32(_out4, vld1_f32(pC + c_hstep * 4), beta); + _out5 = beta == 1.f ? vadd_f32(_out5, vld1_f32(pC + c_hstep * 5)) : vmla_n_f32(_out5, vld1_f32(pC + c_hstep * 5), beta); + _out6 = beta == 1.f ? vadd_f32(_out6, vld1_f32(pC + c_hstep * 6)) : vmla_n_f32(_out6, vld1_f32(pC + c_hstep * 6), beta); + _out7 = beta == 1.f ? vadd_f32(_out7, vld1_f32(pC + c_hstep * 7)) : vmla_n_f32(_out7, vld1_f32(pC + c_hstep * 7), beta); + pC += 2; + } + if (broadcast_type_C == 4) + { + const float32x2_t _c = beta == 1.f ? vld1_f32(pC) : vmul_n_f32(vld1_f32(pC), beta); + _out0 = vadd_f32(_out0, _c); _out1 = vadd_f32(_out1, _c); + _out2 = vadd_f32(_out2, _c); _out3 = vadd_f32(_out3, _c); + _out4 = vadd_f32(_out4, _c); _out5 = vadd_f32(_out5, _c); + _out6 = vadd_f32(_out6, _c); _out7 = vadd_f32(_out7, _c); + pC += 2; + } + } + if (alpha != 1.f) + { + _out0 = vmul_n_f32(_out0, alpha); _out1 = vmul_n_f32(_out1, alpha); + _out2 = vmul_n_f32(_out2, alpha); _out3 = vmul_n_f32(_out3, alpha); + _out4 = vmul_n_f32(_out4, alpha); _out5 = vmul_n_f32(_out5, alpha); + _out6 = vmul_n_f32(_out6, alpha); _out7 = vmul_n_f32(_out7, alpha); + } + vst1_f32(outptr0, _out0); vst1_f32(outptr1, _out1); + vst1_f32(outptr2, _out2); vst1_f32(outptr3, _out3); + vst1_f32(outptr4, _out4); vst1_f32(outptr5, _out5); + vst1_f32(outptr6, _out6); vst1_f32(outptr7, _out7); + outptr0 += 2; outptr1 += 2; outptr2 += 2; outptr3 += 2; + outptr4 += 2; outptr5 += 2; outptr6 += 2; outptr7 += 2; + pp += 16; + } + for (; jj < max_jj; jj++) + { + float f0 = pp[0]; float f1 = pp[1]; float f2 = pp[2]; float f3 = pp[3]; + float f4 = pp[4]; float f5 = pp[5]; float f6 = pp[6]; float f7 = pp[7]; + if (pC) + { + if (broadcast_type_C <= 2) + { + f0 += c0; f1 += c1; f2 += c2; f3 += c3; + f4 += c4; f5 += c5; f6 += c6; f7 += c7; + } + if (broadcast_type_C == 3) + { + f0 += pC[0] * beta; f1 += pC[c_hstep] * beta; + f2 += pC[c_hstep * 2] * beta; f3 += pC[c_hstep * 3] * beta; + f4 += pC[c_hstep * 4] * beta; f5 += pC[c_hstep * 5] * beta; + f6 += pC[c_hstep * 6] * beta; f7 += pC[c_hstep * 7] * beta; + pC++; + } + if (broadcast_type_C == 4) + { + const float c = pC[0] * beta; + f0 += c; f1 += c; f2 += c; f3 += c; + f4 += c; f5 += c; f6 += c; f7 += c; + pC++; + } + } + outptr0[0] = f0 * alpha; outptr1[0] = f1 * alpha; + outptr2[0] = f2 * alpha; outptr3[0] = f3 * alpha; + outptr4[0] = f4 * alpha; outptr5[0] = f5 * alpha; + outptr6[0] = f6 * alpha; outptr7[0] = f7 * alpha; + outptr0++; outptr1++; outptr2++; outptr3++; + outptr4++; outptr5++; outptr6++; outptr7++; + pp += 8; + } + } +#endif // __aarch64__ + for (; ii + 3 < max_ii; ii += 4) + { + pC = (const float*)C; + float* outptr0 = (float*)top_blob + (size_t)(i + ii) * out_hstep + j; + float* outptr1 = outptr0 + out_hstep; + float* outptr2 = outptr1 + out_hstep; + float* outptr3 = outptr2 + out_hstep; + float c0 = 0.f; + float c1 = 0.f; + float c2 = 0.f; + float c3 = 0.f; + if (pC) + { + if (broadcast_type_C == 0) + { + float c = pC[0]; + if (beta != 1.f) + c *= beta; + c0 = c; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + c0 = pC[i + ii]; + if (beta != 1.f) + c0 *= beta; + c1 = pC[i + ii + 1]; + if (beta != 1.f) + c1 *= beta; + c2 = pC[i + ii + 2]; + if (beta != 1.f) + c2 *= beta; + c3 = pC[i + ii + 3]; + if (beta != 1.f) + c3 *= beta; + } + if (broadcast_type_C == 3) + pC = (const float*)C + (i + ii) * c_hstep + j; + if (broadcast_type_C == 4) + pC = (const float*)C + j; + } + + int jj = 0; +#if __aarch64__ + for (; jj + 7 < max_jj; jj += 8) + { + float32x4_t _out0 = vld1q_f32(pp + 0); + float32x4_t _out1 = vld1q_f32(pp + 4); + float32x4_t _out2 = vld1q_f32(pp + 8); + float32x4_t _out3 = vld1q_f32(pp + 12); + float32x4_t _out4 = vld1q_f32(pp + 16); + float32x4_t _out5 = vld1q_f32(pp + 20); + float32x4_t _out6 = vld1q_f32(pp + 24); + float32x4_t _out7 = vld1q_f32(pp + 28); + + if (pC) + { + if (broadcast_type_C == 0) + { + const float32x4_t _c = vdupq_n_f32(c0); + _out0 = vaddq_f32(_out0, _c); + _out1 = vaddq_f32(_out1, _c); + _out2 = vaddq_f32(_out2, _c); + _out3 = vaddq_f32(_out3, _c); + _out4 = vaddq_f32(_out4, _c); + _out5 = vaddq_f32(_out5, _c); + _out6 = vaddq_f32(_out6, _c); + _out7 = vaddq_f32(_out7, _c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _out0 = vaddq_f32(_out0, vdupq_n_f32(c0)); + _out1 = vaddq_f32(_out1, vdupq_n_f32(c0)); + _out2 = vaddq_f32(_out2, vdupq_n_f32(c1)); + _out3 = vaddq_f32(_out3, vdupq_n_f32(c1)); + _out4 = vaddq_f32(_out4, vdupq_n_f32(c2)); + _out5 = vaddq_f32(_out5, vdupq_n_f32(c2)); + _out6 = vaddq_f32(_out6, vdupq_n_f32(c3)); + _out7 = vaddq_f32(_out7, vdupq_n_f32(c3)); + } + if (broadcast_type_C == 3) + { + _out0 = vaddq_f32(_out0, (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta))); + _out1 = vaddq_f32(_out1, (beta == 1.f ? vld1q_f32(pC + 4) : vmulq_n_f32(vld1q_f32(pC + 4), beta))); + _out2 = vaddq_f32(_out2, (beta == 1.f ? vld1q_f32(pC + c_hstep) : vmulq_n_f32(vld1q_f32(pC + c_hstep), beta))); + _out3 = vaddq_f32(_out3, (beta == 1.f ? vld1q_f32(pC + c_hstep + 4) : vmulq_n_f32(vld1q_f32(pC + c_hstep + 4), beta))); + _out4 = vaddq_f32(_out4, (beta == 1.f ? vld1q_f32(pC + c_hstep * 2) : vmulq_n_f32(vld1q_f32(pC + c_hstep * 2), beta))); + _out5 = vaddq_f32(_out5, (beta == 1.f ? vld1q_f32(pC + c_hstep * 2 + 4) : vmulq_n_f32(vld1q_f32(pC + c_hstep * 2 + 4), beta))); + _out6 = vaddq_f32(_out6, (beta == 1.f ? vld1q_f32(pC + c_hstep * 3) : vmulq_n_f32(vld1q_f32(pC + c_hstep * 3), beta))); + _out7 = vaddq_f32(_out7, (beta == 1.f ? vld1q_f32(pC + c_hstep * 3 + 4) : vmulq_n_f32(vld1q_f32(pC + c_hstep * 3 + 4), beta))); + } + if (broadcast_type_C == 4) + { + const float32x4_t _cc0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); + _out0 = vaddq_f32(_out0, _cc0); + _out2 = vaddq_f32(_out2, _cc0); + _out4 = vaddq_f32(_out4, _cc0); + _out6 = vaddq_f32(_out6, _cc0); + const float32x4_t _cc1 = (beta == 1.f ? vld1q_f32(pC + 4) : vmulq_n_f32(vld1q_f32(pC + 4), beta)); + _out1 = vaddq_f32(_out1, _cc1); + _out3 = vaddq_f32(_out3, _cc1); + _out5 = vaddq_f32(_out5, _cc1); + _out7 = vaddq_f32(_out7, _cc1); + } + } + + if (alpha != 1.f) + { + _out0 = vmulq_n_f32(_out0, alpha); + _out1 = vmulq_n_f32(_out1, alpha); + _out2 = vmulq_n_f32(_out2, alpha); + _out3 = vmulq_n_f32(_out3, alpha); + _out4 = vmulq_n_f32(_out4, alpha); + _out5 = vmulq_n_f32(_out5, alpha); + _out6 = vmulq_n_f32(_out6, alpha); + _out7 = vmulq_n_f32(_out7, alpha); + } + + vst1q_f32(outptr0, _out0); + vst1q_f32(outptr0 + 4, _out1); + vst1q_f32(outptr1, _out2); + vst1q_f32(outptr1 + 4, _out3); + vst1q_f32(outptr2, _out4); + vst1q_f32(outptr2 + 4, _out5); + vst1q_f32(outptr3, _out6); + vst1q_f32(outptr3 + 4, _out7); + + outptr0 += 8; + outptr1 += 8; + outptr2 += 8; + outptr3 += 8; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 8; + pp += 32; + } +#endif // __aarch64__ + for (; jj + 3 < max_jj; jj += 4) + { + float32x4_t _out0 = vld1q_f32(pp + 0); + float32x4_t _out1 = vld1q_f32(pp + 4); + float32x4_t _out2 = vld1q_f32(pp + 8); + float32x4_t _out3 = vld1q_f32(pp + 12); + + if (pC) + { + if (broadcast_type_C == 0) + { + const float32x4_t _c = vdupq_n_f32(c0); + _out0 = vaddq_f32(_out0, _c); + _out1 = vaddq_f32(_out1, _c); + _out2 = vaddq_f32(_out2, _c); + _out3 = vaddq_f32(_out3, _c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _out0 = vaddq_f32(_out0, vdupq_n_f32(c0)); + _out1 = vaddq_f32(_out1, vdupq_n_f32(c1)); + _out2 = vaddq_f32(_out2, vdupq_n_f32(c2)); + _out3 = vaddq_f32(_out3, vdupq_n_f32(c3)); + } + if (broadcast_type_C == 3) + { + _out0 = vaddq_f32(_out0, (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta))); + _out1 = vaddq_f32(_out1, (beta == 1.f ? vld1q_f32(pC + c_hstep) : vmulq_n_f32(vld1q_f32(pC + c_hstep), beta))); + _out2 = vaddq_f32(_out2, (beta == 1.f ? vld1q_f32(pC + c_hstep * 2) : vmulq_n_f32(vld1q_f32(pC + c_hstep * 2), beta))); + _out3 = vaddq_f32(_out3, (beta == 1.f ? vld1q_f32(pC + c_hstep * 3) : vmulq_n_f32(vld1q_f32(pC + c_hstep * 3), beta))); + } + if (broadcast_type_C == 4) + { + const float32x4_t _cc0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); + _out0 = vaddq_f32(_out0, _cc0); + _out1 = vaddq_f32(_out1, _cc0); + _out2 = vaddq_f32(_out2, _cc0); + _out3 = vaddq_f32(_out3, _cc0); + } + } + + if (alpha != 1.f) + { + _out0 = vmulq_n_f32(_out0, alpha); + _out1 = vmulq_n_f32(_out1, alpha); + _out2 = vmulq_n_f32(_out2, alpha); + _out3 = vmulq_n_f32(_out3, alpha); + } + + vst1q_f32(outptr0, _out0); + vst1q_f32(outptr1, _out1); + vst1q_f32(outptr2, _out2); + vst1q_f32(outptr3, _out3); + + outptr0 += 4; + outptr1 += 4; + outptr2 += 4; + outptr3 += 4; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 4; + pp += 16; + } + for (; jj + 1 < max_jj; jj += 2) + { + float32x4_t _out0 = vld1q_f32(pp); + float32x4_t _out1 = vld1q_f32(pp + 4); + + if (pC) + { + if (broadcast_type_C == 0) + { + const float32x4_t _c = vdupq_n_f32(c0); + _out0 = vaddq_f32(_out0, _c); + _out1 = vaddq_f32(_out1, _c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + const float32x4_t _c01 = vcombine_f32(vdup_n_f32(c0), vdup_n_f32(c1)); + const float32x4_t _c23 = vcombine_f32(vdup_n_f32(c2), vdup_n_f32(c3)); + _out0 = vaddq_f32(_out0, _c01); + _out1 = vaddq_f32(_out1, _c23); + } + if (broadcast_type_C == 3) + { + const float32x4_t _c01 = vcombine_f32(vld1_f32(pC), vld1_f32(pC + c_hstep)); + const float32x4_t _c23 = vcombine_f32(vld1_f32(pC + c_hstep * 2), vld1_f32(pC + c_hstep * 3)); + _out0 = beta == 1.f ? vaddq_f32(_out0, _c01) : vmlaq_n_f32(_out0, _c01, beta); + _out1 = beta == 1.f ? vaddq_f32(_out1, _c23) : vmlaq_n_f32(_out1, _c23, beta); + } + if (broadcast_type_C == 4) + { + float32x2_t _c = vld1_f32(pC); + if (beta != 1.f) + _c = vmul_n_f32(_c, beta); + const float32x4_t _cc0 = vcombine_f32(_c, _c); + _out0 = vaddq_f32(_out0, _cc0); + _out1 = vaddq_f32(_out1, _cc0); + } + } + + if (alpha != 1.f) + { + _out0 = vmulq_n_f32(_out0, alpha); + _out1 = vmulq_n_f32(_out1, alpha); + } + + vst1_f32(outptr0, vget_low_f32(_out0)); + vst1_f32(outptr1, vget_high_f32(_out0)); + vst1_f32(outptr2, vget_low_f32(_out1)); + vst1_f32(outptr3, vget_high_f32(_out1)); + + outptr0 += 2; + outptr1 += 2; + outptr2 += 2; + outptr3 += 2; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 2; + pp += 8; + } + for (; jj < max_jj; jj += 1) + { + float32x4_t _out0 = vld1q_f32(pp); + + if (pC) + { + if (broadcast_type_C == 0) + { + _out0 = vaddq_f32(_out0, vdupq_n_f32(c0)); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + float32x4_t _c = vdupq_n_f32(c0); + _c = vsetq_lane_f32(c1, _c, 1); + _c = vsetq_lane_f32(c2, _c, 2); + _c = vsetq_lane_f32(c3, _c, 3); + _out0 = vaddq_f32(_out0, _c); + } + if (broadcast_type_C == 3) + { + float32x4_t _c = vdupq_n_f32(pC[0]); + _c = vsetq_lane_f32(pC[c_hstep], _c, 1); + _c = vsetq_lane_f32(pC[c_hstep * 2], _c, 2); + _c = vsetq_lane_f32(pC[c_hstep * 3], _c, 3); + _out0 = beta == 1.f ? vaddq_f32(_out0, _c) : vmlaq_n_f32(_out0, _c, beta); + } + if (broadcast_type_C == 4) + { + const float32x4_t _cc0 = vdupq_n_f32(beta == 1.f ? pC[0] : pC[0] * beta); + _out0 = vaddq_f32(_out0, _cc0); + } + } + + if (alpha != 1.f) + { + _out0 = vmulq_n_f32(_out0, alpha); + } + + vst1q_lane_f32(outptr0, _out0, 0); + vst1q_lane_f32(outptr1, _out0, 1); + vst1q_lane_f32(outptr2, _out0, 2); + vst1q_lane_f32(outptr3, _out0, 3); + + outptr0++; + outptr1++; + outptr2++; + outptr3++; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC++; + pp += 4; + } + } + for (; ii + 1 < max_ii; ii += 2) + { + pC = (const float*)C; + float* outptr0 = (float*)top_blob + (size_t)(i + ii) * out_hstep + j; + float* outptr1 = outptr0 + out_hstep; + float32x4_t _c0; + float32x4_t _c1; + float c0 = 0.f; + float c1 = 0.f; + if (pC) + { + if (broadcast_type_C == 0) + { + c0 = pC[0]; + if (beta != 1.f) + c0 *= beta; + _c0 = vdupq_n_f32(c0); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + c0 = pC[i + ii]; + if (beta != 1.f) + c0 *= beta; + _c0 = vdupq_n_f32(c0); + c1 = pC[i + ii + 1]; + if (beta != 1.f) + c1 *= beta; + _c1 = vdupq_n_f32(c1); + } + if (broadcast_type_C == 3) + pC = (const float*)C + (i + ii) * c_hstep + j; + if (broadcast_type_C == 4) + pC = (const float*)C + j; + } + + int jj = 0; +#if __aarch64__ + for (; jj + 7 < max_jj; jj += 8) + { + float32x4_t _out0 = vld1q_f32(pp + 0); + float32x4_t _out1 = vld1q_f32(pp + 4); + float32x4_t _out2 = vld1q_f32(pp + 8); + float32x4_t _out3 = vld1q_f32(pp + 12); + + if (pC) + { + if (broadcast_type_C == 0) + { + _out0 = vaddq_f32(_out0, _c0); + _out1 = vaddq_f32(_out1, _c0); + _out2 = vaddq_f32(_out2, _c0); + _out3 = vaddq_f32(_out3, _c0); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _out0 = vaddq_f32(_out0, _c0); + _out1 = vaddq_f32(_out1, _c0); + _out2 = vaddq_f32(_out2, _c1); + _out3 = vaddq_f32(_out3, _c1); + } + if (broadcast_type_C == 3) + { + _out0 = vaddq_f32(_out0, (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta))); + _out1 = vaddq_f32(_out1, (beta == 1.f ? vld1q_f32(pC + 4) : vmulq_n_f32(vld1q_f32(pC + 4), beta))); + _out2 = vaddq_f32(_out2, (beta == 1.f ? vld1q_f32(pC + c_hstep) : vmulq_n_f32(vld1q_f32(pC + c_hstep), beta))); + _out3 = vaddq_f32(_out3, (beta == 1.f ? vld1q_f32(pC + c_hstep + 4) : vmulq_n_f32(vld1q_f32(pC + c_hstep + 4), beta))); + } + if (broadcast_type_C == 4) + { + const float32x4_t _c0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); + _out0 = vaddq_f32(_out0, _c0); + _out2 = vaddq_f32(_out2, _c0); + const float32x4_t _c1 = (beta == 1.f ? vld1q_f32(pC + 4) : vmulq_n_f32(vld1q_f32(pC + 4), beta)); + _out1 = vaddq_f32(_out1, _c1); + _out3 = vaddq_f32(_out3, _c1); + } + } + + if (alpha != 1.f) + { + _out0 = vmulq_n_f32(_out0, alpha); + _out1 = vmulq_n_f32(_out1, alpha); + _out2 = vmulq_n_f32(_out2, alpha); + _out3 = vmulq_n_f32(_out3, alpha); + } + + vst1q_f32(outptr0, _out0); + vst1q_f32(outptr0 + 4, _out1); + vst1q_f32(outptr1, _out2); + vst1q_f32(outptr1 + 4, _out3); + + outptr0 += 8; + outptr1 += 8; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 8; + pp += 16; + } +#endif // __aarch64__ + for (; jj + 3 < max_jj; jj += 4) + { + float32x4_t _out0 = vld1q_f32(pp + 0); + float32x4_t _out1 = vld1q_f32(pp + 4); + + if (pC) + { + if (broadcast_type_C == 0) + { + _out0 = vaddq_f32(_out0, _c0); + _out1 = vaddq_f32(_out1, _c0); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _out0 = vaddq_f32(_out0, _c0); + _out1 = vaddq_f32(_out1, _c1); + } + if (broadcast_type_C == 3) + { + _out0 = vaddq_f32(_out0, (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta))); + _out1 = vaddq_f32(_out1, (beta == 1.f ? vld1q_f32(pC + c_hstep) : vmulq_n_f32(vld1q_f32(pC + c_hstep), beta))); + } + if (broadcast_type_C == 4) + { + const float32x4_t _c0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); + _out0 = vaddq_f32(_out0, _c0); + _out1 = vaddq_f32(_out1, _c0); + } + } + + if (alpha != 1.f) + { + _out0 = vmulq_n_f32(_out0, alpha); + _out1 = vmulq_n_f32(_out1, alpha); + } + + vst1q_f32(outptr0, _out0); + vst1q_f32(outptr1, _out1); + + outptr0 += 4; + outptr1 += 4; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 4; + pp += 8; + } + for (; jj + 1 < max_jj; jj += 2) + { + float32x2_t _out0 = vld1_f32(pp); + float32x2_t _out1 = vld1_f32(pp + 2); + + if (pC) + { + if (broadcast_type_C == 0) + { + const float32x2_t _c = vdup_n_f32(c0); + _out0 = vadd_f32(_out0, _c); + _out1 = vadd_f32(_out1, _c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _out0 = vadd_f32(_out0, vdup_n_f32(c0)); + _out1 = vadd_f32(_out1, vdup_n_f32(c1)); + } + if (broadcast_type_C == 3) + { + const float32x2_t _c0 = vld1_f32(pC); + const float32x2_t _c1 = vld1_f32(pC + c_hstep); + _out0 = beta == 1.f ? vadd_f32(_out0, _c0) : vmla_n_f32(_out0, _c0, beta); + _out1 = beta == 1.f ? vadd_f32(_out1, _c1) : vmla_n_f32(_out1, _c1, beta); + } + if (broadcast_type_C == 4) + { + const float32x2_t _c0 = beta == 1.f ? vld1_f32(pC) : vmul_n_f32(vld1_f32(pC), beta); + _out0 = vadd_f32(_out0, _c0); + _out1 = vadd_f32(_out1, _c0); + } + } + + if (alpha != 1.f) + { + _out0 = vmul_n_f32(_out0, alpha); + _out1 = vmul_n_f32(_out1, alpha); + } + + vst1_f32(outptr0, _out0); + vst1_f32(outptr1, _out1); + + outptr0 += 2; + outptr1 += 2; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 2; + pp += 4; + } + for (; jj < max_jj; jj += 1) + { + float32x2_t _out0 = vld1_f32(pp); + + if (pC) + { + if (broadcast_type_C == 0) + { + _out0 = vadd_f32(_out0, vdup_n_f32(c0)); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + float32x2_t _c = vdup_n_f32(c0); + _c = vset_lane_f32(c1, _c, 1); + _out0 = vadd_f32(_out0, _c); + } + if (broadcast_type_C == 3) + { + float32x2_t _c = vdup_n_f32(pC[0]); + _c = vset_lane_f32(pC[c_hstep], _c, 1); + _out0 = beta == 1.f ? vadd_f32(_out0, _c) : vmla_n_f32(_out0, _c, beta); + } + if (broadcast_type_C == 4) + { + _out0 = vadd_f32(_out0, vdup_n_f32(beta == 1.f ? pC[0] : pC[0] * beta)); + } + } + + if (alpha != 1.f) + { + _out0 = vmul_n_f32(_out0, alpha); + } + + vst1_lane_f32(outptr0, _out0, 0); + vst1_lane_f32(outptr1, _out0, 1); + + outptr0++; + outptr1++; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC++; + pp += 2; + } + } + for (; ii < max_ii; ii++) + { + pC = (const float*)C; + float* outptr0 = (float*)top_blob + (size_t)(i + ii) * out_hstep + j; + float32x4_t _c0; + float c0 = 0.f; + if (pC) + { + if (broadcast_type_C == 0) + { + c0 = pC[0]; + if (beta != 1.f) + c0 *= beta; + _c0 = vdupq_n_f32(c0); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + c0 = pC[i + ii]; + if (beta != 1.f) + c0 *= beta; + _c0 = vdupq_n_f32(c0); + } + if (broadcast_type_C == 3) + pC = (const float*)C + (i + ii) * c_hstep + j; + if (broadcast_type_C == 4) + pC = (const float*)C + j; + } + + int jj = 0; +#if __aarch64__ + for (; jj + 7 < max_jj; jj += 8) + { + float32x4_t _out0 = vld1q_f32(pp + 0); + float32x4_t _out1 = vld1q_f32(pp + 4); + + if (pC) + { + if (broadcast_type_C == 0) + { + _out0 = vaddq_f32(_out0, _c0); + _out1 = vaddq_f32(_out1, _c0); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _out0 = vaddq_f32(_out0, _c0); + _out1 = vaddq_f32(_out1, _c0); + } + if (broadcast_type_C == 3) + { + _out0 = vaddq_f32(_out0, (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta))); + _out1 = vaddq_f32(_out1, (beta == 1.f ? vld1q_f32(pC + 4) : vmulq_n_f32(vld1q_f32(pC + 4), beta))); + } + if (broadcast_type_C == 4) + { + const float32x4_t _c0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); + _out0 = vaddq_f32(_out0, _c0); + const float32x4_t _c1 = (beta == 1.f ? vld1q_f32(pC + 4) : vmulq_n_f32(vld1q_f32(pC + 4), beta)); + _out1 = vaddq_f32(_out1, _c1); + } + } + + if (alpha != 1.f) + { + _out0 = vmulq_n_f32(_out0, alpha); + _out1 = vmulq_n_f32(_out1, alpha); + } + + vst1q_f32(outptr0, _out0); + vst1q_f32(outptr0 + 4, _out1); + + outptr0 += 8; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 8; + pp += 8; + } +#endif // __aarch64__ + for (; jj + 3 < max_jj; jj += 4) + { + float32x4_t _out0 = vld1q_f32(pp + 0); + + if (pC) + { + if (broadcast_type_C == 0) + { + _out0 = vaddq_f32(_out0, _c0); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _out0 = vaddq_f32(_out0, _c0); + } + if (broadcast_type_C == 3) + { + _out0 = vaddq_f32(_out0, (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta))); + } + if (broadcast_type_C == 4) + { + const float32x4_t _c0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); + _out0 = vaddq_f32(_out0, _c0); + } + } + + if (alpha != 1.f) + { + _out0 = vmulq_n_f32(_out0, alpha); + } + + vst1q_f32(outptr0, _out0); + + outptr0 += 4; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 4; + pp += 4; + } + for (; jj + 1 < max_jj; jj += 2) + { + float32x2_t _out0 = vld1_f32(pp); + + if (pC) + { + if (broadcast_type_C == 0) + { + _out0 = vadd_f32(_out0, vdup_n_f32(c0)); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _out0 = vadd_f32(_out0, vdup_n_f32(c0)); + } + if (broadcast_type_C == 3) + { + const float32x2_t _c0 = vld1_f32(pC); + _out0 = beta == 1.f ? vadd_f32(_out0, _c0) : vmla_n_f32(_out0, _c0, beta); + } + if (broadcast_type_C == 4) + { + const float32x2_t _c0 = beta == 1.f ? vld1_f32(pC) : vmul_n_f32(vld1_f32(pC), beta); + _out0 = vadd_f32(_out0, _c0); + } + } + + if (alpha != 1.f) + { + _out0 = vmul_n_f32(_out0, alpha); + } + + vst1_f32(outptr0, _out0); + + outptr0 += 2; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 2; + pp += 2; + } + for (; jj < max_jj; jj += 1) + { + float out0 = pp[0]; + + if (pC) + { + if (broadcast_type_C == 0) + { + out0 += c0; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + out0 += c0; + } + if (broadcast_type_C == 3) + { + out0 += beta == 1.f ? pC[0] : pC[0] * beta; + } + if (broadcast_type_C == 4) + { + out0 += beta == 1.f ? pC[0] : pC[0] * beta; + } + } + + if (alpha != 1.f) + { + out0 *= alpha; + } + + outptr0[0] = out0; + + outptr0++; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC++; + pp += 1; + } + } +#else + for (; ii + 1 < max_ii; ii += 2) + { + pC = (const float*)C; + float* outptr0 = (float*)top_blob + (size_t)(i + ii) * out_hstep + j; + float* outptr1 = outptr0 + out_hstep; + float c0 = 0.f; + float c1 = 0.f; + if (pC) + { + if (broadcast_type_C == 0) + { + float c = pC[0]; + if (beta != 1.f) + c *= beta; + c0 = c; + c1 = c; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + c0 = pC[i + ii]; + if (beta != 1.f) + c0 *= beta; + c1 = pC[i + ii + 1]; + if (beta != 1.f) + c1 *= beta; + } + if (broadcast_type_C == 3) + pC = (const float*)C + (i + ii) * c_hstep + j; + if (broadcast_type_C == 4) + pC = (const float*)C + j; + } + + int jj = 0; + for (; jj + 1 < max_jj; jj += 2) + { + float out00 = pp[0]; + float out01 = pp[1]; + float out10 = pp[2]; + float out11 = pp[3]; + + if (pC) + { + if (broadcast_type_C == 0) + { + out00 += c0; + out01 += c0; + out10 += c0; + out11 += c0; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + out00 += c0; + out01 += c0; + out10 += c1; + out11 += c1; + } + if (broadcast_type_C == 3) + { + out00 += beta == 1.f ? pC[0] : pC[0] * beta; + out01 += beta == 1.f ? pC[1] : pC[1] * beta; + out10 += beta == 1.f ? pC[c_hstep] : pC[c_hstep] * beta; + out11 += beta == 1.f ? pC[c_hstep + 1] : pC[c_hstep + 1] * beta; + } + if (broadcast_type_C == 4) + { + out00 += beta == 1.f ? pC[0] : pC[0] * beta; + out01 += beta == 1.f ? pC[1] : pC[1] * beta; + out10 += beta == 1.f ? pC[0] : pC[0] * beta; + out11 += beta == 1.f ? pC[1] : pC[1] * beta; + } + } + + if (alpha != 1.f) + { + out00 *= alpha; + out01 *= alpha; + out10 *= alpha; + out11 *= alpha; + } + + outptr0[0] = out00; + outptr0[1] = out01; + outptr1[0] = out10; + outptr1[1] = out11; + + outptr0 += 2; + outptr1 += 2; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 2; + pp += 4; + } + for (; jj < max_jj; jj += 1) + { + float out00 = pp[0]; + float out10 = pp[1]; + + if (pC) + { + if (broadcast_type_C == 0) + { + out00 += c0; + out10 += c0; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + out00 += c0; + out10 += c1; + } + if (broadcast_type_C == 3) + { + out00 += beta == 1.f ? pC[0] : pC[0] * beta; + out10 += beta == 1.f ? pC[c_hstep] : pC[c_hstep] * beta; + } + if (broadcast_type_C == 4) + { + out00 += beta == 1.f ? pC[0] : pC[0] * beta; + out10 += beta == 1.f ? pC[0] : pC[0] * beta; + } + } + + if (alpha != 1.f) + { + out00 *= alpha; + out10 *= alpha; + } + + outptr0[0] = out00; + outptr1[0] = out10; + + outptr0++; + outptr1++; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC++; + pp += 2; + } + } + for (; ii < max_ii; ii++) + { + pC = (const float*)C; + float* outptr0 = (float*)top_blob + (size_t)(i + ii) * out_hstep + j; + float c0 = 0.f; + if (pC) + { + if (broadcast_type_C == 0) + { + float c = pC[0]; + if (beta != 1.f) + c *= beta; + c0 = c; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + c0 = pC[i + ii]; + if (beta != 1.f) + c0 *= beta; + } + if (broadcast_type_C == 3) + pC = (const float*)C + (i + ii) * c_hstep + j; + if (broadcast_type_C == 4) + pC = (const float*)C + j; + } + + int jj = 0; + for (; jj + 1 < max_jj; jj += 2) + { + float out00 = pp[0]; + float out01 = pp[1]; + + if (pC) + { + if (broadcast_type_C == 0) + { + out00 += c0; + out01 += c0; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + out00 += c0; + out01 += c0; + } + if (broadcast_type_C == 3) + { + out00 += beta == 1.f ? pC[0] : pC[0] * beta; + out01 += beta == 1.f ? pC[1] : pC[1] * beta; + } + if (broadcast_type_C == 4) + { + out00 += beta == 1.f ? pC[0] : pC[0] * beta; + out01 += beta == 1.f ? pC[1] : pC[1] * beta; + } + } + + if (alpha != 1.f) + { + out00 *= alpha; + out01 *= alpha; + } + + outptr0[0] = out00; + outptr0[1] = out01; + + outptr0 += 2; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 2; + pp += 2; + } + for (; jj < max_jj; jj += 1) + { + float out00 = pp[0]; + + if (pC) + { + if (broadcast_type_C == 0) + { + out00 += c0; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + out00 += c0; + } + if (broadcast_type_C == 3) + { + out00 += beta == 1.f ? pC[0] : pC[0] * beta; + } + if (broadcast_type_C == 4) + { + out00 += beta == 1.f ? pC[0] : pC[0] * beta; + } + } + + if (alpha != 1.f) + { + out00 *= alpha; + } + + outptr0[0] = out00; + + outptr0++; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC++; + pp += 1; + } + } +#endif // __ARM_NEON +} + +static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta) +{ + (void)N; + + const float* pC = C; + const float* pp = topT; + const size_t out_hstep = top_blob.dims == 3 ? top_blob.cstep : (size_t)top_blob.w; + const size_t c_hstep = C.dims == 3 ? C.cstep : (size_t)C.w; + + int ii = 0; +#if __ARM_NEON +#if __aarch64__ + for (; ii + 7 < max_ii; ii += 8) + { + float* outptr0 = (float*)top_blob + (size_t)j * out_hstep + i + ii; + float* outptr1 = outptr0 + out_hstep; + float* outptr2 = outptr1 + out_hstep; + float* outptr3 = outptr2 + out_hstep; + pC = (const float*)C; + float c0 = 0.f; + float c1 = 0.f; + float c2 = 0.f; + float c3 = 0.f; + float c4 = 0.f; + float c5 = 0.f; + float c6 = 0.f; + float c7 = 0.f; + if (pC) + { + if (broadcast_type_C == 0) + c0 = c1 = c2 = c3 = c4 = c5 = c6 = c7 = pC[0] * beta; + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + pC = (const float*)C + i + ii; + c0 = pC[0] * beta; + c1 = pC[1] * beta; + c2 = pC[2] * beta; + c3 = pC[3] * beta; + c4 = pC[4] * beta; + c5 = pC[5] * beta; + c6 = pC[6] * beta; + c7 = pC[7] * beta; + } + if (broadcast_type_C == 3) + pC = (const float*)C + (size_t)(i + ii) * c_hstep + j; + if (broadcast_type_C == 4) + pC = (const float*)C + j; + } + + int jj = 0; + for (; jj + 3 < max_jj; jj += 4) + { + float32x4_t _out0 = vld1q_f32(pp); + float32x4_t _out1 = vld1q_f32(pp + 4); + float32x4_t _out2 = vld1q_f32(pp + 8); + float32x4_t _out3 = vld1q_f32(pp + 12); + float32x4_t _out4 = vld1q_f32(pp + 16); + float32x4_t _out5 = vld1q_f32(pp + 20); + float32x4_t _out6 = vld1q_f32(pp + 24); + float32x4_t _out7 = vld1q_f32(pp + 28); + if (pC) + { + if (broadcast_type_C <= 2) + { + _out0 = vaddq_f32(_out0, vdupq_n_f32(c0)); + _out1 = vaddq_f32(_out1, vdupq_n_f32(c1)); + _out2 = vaddq_f32(_out2, vdupq_n_f32(c2)); + _out3 = vaddq_f32(_out3, vdupq_n_f32(c3)); + _out4 = vaddq_f32(_out4, vdupq_n_f32(c4)); + _out5 = vaddq_f32(_out5, vdupq_n_f32(c5)); + _out6 = vaddq_f32(_out6, vdupq_n_f32(c6)); + _out7 = vaddq_f32(_out7, vdupq_n_f32(c7)); + } + if (broadcast_type_C == 3) + { + _out0 = beta == 1.f ? vaddq_f32(_out0, vld1q_f32(pC)) : vmlaq_n_f32(_out0, vld1q_f32(pC), beta); + _out1 = beta == 1.f ? vaddq_f32(_out1, vld1q_f32(pC + c_hstep)) : vmlaq_n_f32(_out1, vld1q_f32(pC + c_hstep), beta); + _out2 = beta == 1.f ? vaddq_f32(_out2, vld1q_f32(pC + c_hstep * 2)) : vmlaq_n_f32(_out2, vld1q_f32(pC + c_hstep * 2), beta); + _out3 = beta == 1.f ? vaddq_f32(_out3, vld1q_f32(pC + c_hstep * 3)) : vmlaq_n_f32(_out3, vld1q_f32(pC + c_hstep * 3), beta); + _out4 = beta == 1.f ? vaddq_f32(_out4, vld1q_f32(pC + c_hstep * 4)) : vmlaq_n_f32(_out4, vld1q_f32(pC + c_hstep * 4), beta); + _out5 = beta == 1.f ? vaddq_f32(_out5, vld1q_f32(pC + c_hstep * 5)) : vmlaq_n_f32(_out5, vld1q_f32(pC + c_hstep * 5), beta); + _out6 = beta == 1.f ? vaddq_f32(_out6, vld1q_f32(pC + c_hstep * 6)) : vmlaq_n_f32(_out6, vld1q_f32(pC + c_hstep * 6), beta); + _out7 = beta == 1.f ? vaddq_f32(_out7, vld1q_f32(pC + c_hstep * 7)) : vmlaq_n_f32(_out7, vld1q_f32(pC + c_hstep * 7), beta); + pC += 4; + } + if (broadcast_type_C == 4) + { + const float32x4_t _c = beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta); + _out0 = vaddq_f32(_out0, _c); + _out1 = vaddq_f32(_out1, _c); + _out2 = vaddq_f32(_out2, _c); + _out3 = vaddq_f32(_out3, _c); + _out4 = vaddq_f32(_out4, _c); + _out5 = vaddq_f32(_out5, _c); + _out6 = vaddq_f32(_out6, _c); + _out7 = vaddq_f32(_out7, _c); + pC += 4; + } + } + if (alpha != 1.f) + { + _out0 = vmulq_n_f32(_out0, alpha); + _out1 = vmulq_n_f32(_out1, alpha); + _out2 = vmulq_n_f32(_out2, alpha); + _out3 = vmulq_n_f32(_out3, alpha); + _out4 = vmulq_n_f32(_out4, alpha); + _out5 = vmulq_n_f32(_out5, alpha); + _out6 = vmulq_n_f32(_out6, alpha); + _out7 = vmulq_n_f32(_out7, alpha); + } + transpose4x4_ps(_out0, _out1, _out2, _out3); + transpose4x4_ps(_out4, _out5, _out6, _out7); + vst1q_f32(outptr0, _out0); + vst1q_f32(outptr0 + 4, _out4); + vst1q_f32(outptr1, _out1); + vst1q_f32(outptr1 + 4, _out5); + vst1q_f32(outptr2, _out2); + vst1q_f32(outptr2 + 4, _out6); + vst1q_f32(outptr3, _out3); + vst1q_f32(outptr3 + 4, _out7); + outptr0 += out_hstep * 4; + outptr1 += out_hstep * 4; + outptr2 += out_hstep * 4; + outptr3 += out_hstep * 4; + pp += 32; + } + for (; jj + 1 < max_jj; jj += 2) + { + float32x2_t _out0 = vld1_f32(pp); + float32x2_t _out1 = vld1_f32(pp + 2); + float32x2_t _out2 = vld1_f32(pp + 4); + float32x2_t _out3 = vld1_f32(pp + 6); + float32x2_t _out4 = vld1_f32(pp + 8); + float32x2_t _out5 = vld1_f32(pp + 10); + float32x2_t _out6 = vld1_f32(pp + 12); + float32x2_t _out7 = vld1_f32(pp + 14); + if (pC) + { + if (broadcast_type_C <= 2) + { + _out0 = vadd_f32(_out0, vdup_n_f32(c0)); _out1 = vadd_f32(_out1, vdup_n_f32(c1)); + _out2 = vadd_f32(_out2, vdup_n_f32(c2)); _out3 = vadd_f32(_out3, vdup_n_f32(c3)); + _out4 = vadd_f32(_out4, vdup_n_f32(c4)); _out5 = vadd_f32(_out5, vdup_n_f32(c5)); + _out6 = vadd_f32(_out6, vdup_n_f32(c6)); _out7 = vadd_f32(_out7, vdup_n_f32(c7)); + } + if (broadcast_type_C == 3) + { + _out0 = beta == 1.f ? vadd_f32(_out0, vld1_f32(pC)) : vmla_n_f32(_out0, vld1_f32(pC), beta); + _out1 = beta == 1.f ? vadd_f32(_out1, vld1_f32(pC + c_hstep)) : vmla_n_f32(_out1, vld1_f32(pC + c_hstep), beta); + _out2 = beta == 1.f ? vadd_f32(_out2, vld1_f32(pC + c_hstep * 2)) : vmla_n_f32(_out2, vld1_f32(pC + c_hstep * 2), beta); + _out3 = beta == 1.f ? vadd_f32(_out3, vld1_f32(pC + c_hstep * 3)) : vmla_n_f32(_out3, vld1_f32(pC + c_hstep * 3), beta); + _out4 = beta == 1.f ? vadd_f32(_out4, vld1_f32(pC + c_hstep * 4)) : vmla_n_f32(_out4, vld1_f32(pC + c_hstep * 4), beta); + _out5 = beta == 1.f ? vadd_f32(_out5, vld1_f32(pC + c_hstep * 5)) : vmla_n_f32(_out5, vld1_f32(pC + c_hstep * 5), beta); + _out6 = beta == 1.f ? vadd_f32(_out6, vld1_f32(pC + c_hstep * 6)) : vmla_n_f32(_out6, vld1_f32(pC + c_hstep * 6), beta); + _out7 = beta == 1.f ? vadd_f32(_out7, vld1_f32(pC + c_hstep * 7)) : vmla_n_f32(_out7, vld1_f32(pC + c_hstep * 7), beta); + pC += 2; + } + if (broadcast_type_C == 4) + { + const float32x2_t _c = beta == 1.f ? vld1_f32(pC) : vmul_n_f32(vld1_f32(pC), beta); + _out0 = vadd_f32(_out0, _c); _out1 = vadd_f32(_out1, _c); + _out2 = vadd_f32(_out2, _c); _out3 = vadd_f32(_out3, _c); + _out4 = vadd_f32(_out4, _c); _out5 = vadd_f32(_out5, _c); + _out6 = vadd_f32(_out6, _c); _out7 = vadd_f32(_out7, _c); + pC += 2; + } + } + if (alpha != 1.f) + { + _out0 = vmul_n_f32(_out0, alpha); _out1 = vmul_n_f32(_out1, alpha); + _out2 = vmul_n_f32(_out2, alpha); _out3 = vmul_n_f32(_out3, alpha); + _out4 = vmul_n_f32(_out4, alpha); _out5 = vmul_n_f32(_out5, alpha); + _out6 = vmul_n_f32(_out6, alpha); _out7 = vmul_n_f32(_out7, alpha); + } + float32x4x2_t _t0 = vuzpq_f32(vcombine_f32(_out0, _out1), vcombine_f32(_out2, _out3)); + float32x4x2_t _t1 = vuzpq_f32(vcombine_f32(_out4, _out5), vcombine_f32(_out6, _out7)); + vst1q_f32(outptr0, _t0.val[0]); vst1q_f32(outptr0 + 4, _t1.val[0]); + vst1q_f32(outptr1, _t0.val[1]); vst1q_f32(outptr1 + 4, _t1.val[1]); + outptr0 += out_hstep * 2; + outptr1 += out_hstep * 2; + pp += 16; + } + for (; jj < max_jj; jj++) + { + float32x4_t _out0 = vld1q_f32(pp); + float32x4_t _out1 = vld1q_f32(pp + 4); + if (pC) + { + if (broadcast_type_C <= 2) + { + const float32x4_t _c0 = {c0, c1, c2, c3}; + const float32x4_t _c1 = {c4, c5, c6, c7}; + _out0 = vaddq_f32(_out0, _c0); + _out1 = vaddq_f32(_out1, _c1); + } + if (broadcast_type_C == 3) + { + float32x4_t _c0 = {pC[0], pC[c_hstep], pC[c_hstep * 2], pC[c_hstep * 3]}; + float32x4_t _c1 = {pC[c_hstep * 4], pC[c_hstep * 5], pC[c_hstep * 6], pC[c_hstep * 7]}; + _out0 = beta == 1.f ? vaddq_f32(_out0, _c0) : vmlaq_n_f32(_out0, _c0, beta); + _out1 = beta == 1.f ? vaddq_f32(_out1, _c1) : vmlaq_n_f32(_out1, _c1, beta); + pC++; + } + if (broadcast_type_C == 4) + { + const float32x4_t _c = vdupq_n_f32(beta == 1.f ? pC[0] : pC[0] * beta); + _out0 = vaddq_f32(_out0, _c); + _out1 = vaddq_f32(_out1, _c); + pC++; + } + } + if (alpha != 1.f) + { + _out0 = vmulq_n_f32(_out0, alpha); + _out1 = vmulq_n_f32(_out1, alpha); + } + vst1q_f32(outptr0, _out0); + vst1q_f32(outptr0 + 4, _out1); + outptr0 += out_hstep; + pp += 8; + } + } +#endif // __aarch64__ + for (; ii + 3 < max_ii; ii += 4) + { + pC = (const float*)C; + float* outptr0 = (float*)top_blob + (size_t)j * out_hstep + i + ii; + float32x4_t _c0123 = vdupq_n_f32(0.f); + if (pC) + { + if (broadcast_type_C == 0) + { + float c = pC[0]; + if (beta != 1.f) + c *= beta; + _c0123 = vdupq_n_f32(c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _c0123 = vld1q_f32(pC + i + ii); + if (beta != 1.f) + _c0123 = vmulq_n_f32(_c0123, beta); + } + if (broadcast_type_C == 3) + pC = (const float*)C + (i + ii) * c_hstep + j; + if (broadcast_type_C == 4) + pC = (const float*)C + j; + } + + int jj = 0; +#if __aarch64__ + for (; jj + 7 < max_jj; jj += 8) + { + float32x4_t _out0 = vld1q_f32(pp + 0); + float32x4_t _out1 = vld1q_f32(pp + 4); + float32x4_t _out2 = vld1q_f32(pp + 8); + float32x4_t _out3 = vld1q_f32(pp + 12); + float32x4_t _out4 = vld1q_f32(pp + 16); + float32x4_t _out5 = vld1q_f32(pp + 20); + float32x4_t _out6 = vld1q_f32(pp + 24); + float32x4_t _out7 = vld1q_f32(pp + 28); + + if (pC) + { + if (broadcast_type_C == 3) + { + _out0 = vaddq_f32(_out0, (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta))); + _out1 = vaddq_f32(_out1, (beta == 1.f ? vld1q_f32(pC + 4) : vmulq_n_f32(vld1q_f32(pC + 4), beta))); + _out2 = vaddq_f32(_out2, (beta == 1.f ? vld1q_f32(pC + c_hstep) : vmulq_n_f32(vld1q_f32(pC + c_hstep), beta))); + _out3 = vaddq_f32(_out3, (beta == 1.f ? vld1q_f32(pC + c_hstep + 4) : vmulq_n_f32(vld1q_f32(pC + c_hstep + 4), beta))); + _out4 = vaddq_f32(_out4, (beta == 1.f ? vld1q_f32(pC + c_hstep * 2) : vmulq_n_f32(vld1q_f32(pC + c_hstep * 2), beta))); + _out5 = vaddq_f32(_out5, (beta == 1.f ? vld1q_f32(pC + c_hstep * 2 + 4) : vmulq_n_f32(vld1q_f32(pC + c_hstep * 2 + 4), beta))); + _out6 = vaddq_f32(_out6, (beta == 1.f ? vld1q_f32(pC + c_hstep * 3) : vmulq_n_f32(vld1q_f32(pC + c_hstep * 3), beta))); + _out7 = vaddq_f32(_out7, (beta == 1.f ? vld1q_f32(pC + c_hstep * 3 + 4) : vmulq_n_f32(vld1q_f32(pC + c_hstep * 3 + 4), beta))); + } + if (broadcast_type_C == 4) + { + const float32x4_t _cc0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); + _out0 = vaddq_f32(_out0, _cc0); + _out2 = vaddq_f32(_out2, _cc0); + _out4 = vaddq_f32(_out4, _cc0); + _out6 = vaddq_f32(_out6, _cc0); + const float32x4_t _cc1 = (beta == 1.f ? vld1q_f32(pC + 4) : vmulq_n_f32(vld1q_f32(pC + 4), beta)); + _out1 = vaddq_f32(_out1, _cc1); + _out3 = vaddq_f32(_out3, _cc1); + _out5 = vaddq_f32(_out5, _cc1); + _out7 = vaddq_f32(_out7, _cc1); + } + } + + if (out_hstep == 4) + { + if (pC) + { + if (broadcast_type_C == 0) + { + _out0 = vaddq_f32(_out0, _c0123); + _out1 = vaddq_f32(_out1, _c0123); + _out2 = vaddq_f32(_out2, _c0123); + _out3 = vaddq_f32(_out3, _c0123); + _out4 = vaddq_f32(_out4, _c0123); + _out5 = vaddq_f32(_out5, _c0123); + _out6 = vaddq_f32(_out6, _c0123); + _out7 = vaddq_f32(_out7, _c0123); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + const float32x4_t _c0 = vdupq_lane_f32(vget_low_f32(_c0123), 0); + const float32x4_t _c1 = vdupq_lane_f32(vget_low_f32(_c0123), 1); + const float32x4_t _c2 = vdupq_lane_f32(vget_high_f32(_c0123), 0); + const float32x4_t _c3 = vdupq_lane_f32(vget_high_f32(_c0123), 1); + _out0 = vaddq_f32(_out0, _c0); + _out1 = vaddq_f32(_out1, _c0); + _out2 = vaddq_f32(_out2, _c1); + _out3 = vaddq_f32(_out3, _c1); + _out4 = vaddq_f32(_out4, _c2); + _out5 = vaddq_f32(_out5, _c2); + _out6 = vaddq_f32(_out6, _c3); + _out7 = vaddq_f32(_out7, _c3); + } + } + + if (alpha != 1.f) + { + _out0 = vmulq_n_f32(_out0, alpha); + _out1 = vmulq_n_f32(_out1, alpha); + _out2 = vmulq_n_f32(_out2, alpha); + _out3 = vmulq_n_f32(_out3, alpha); + _out4 = vmulq_n_f32(_out4, alpha); + _out5 = vmulq_n_f32(_out5, alpha); + _out6 = vmulq_n_f32(_out6, alpha); + _out7 = vmulq_n_f32(_out7, alpha); + } + + float32x4x4_t _r0; + _r0.val[0] = _out0; + _r0.val[1] = _out2; + _r0.val[2] = _out4; + _r0.val[3] = _out6; + vst4q_f32(outptr0, _r0); + float32x4x4_t _r1; + _r1.val[0] = _out1; + _r1.val[1] = _out3; + _r1.val[2] = _out5; + _r1.val[3] = _out7; + vst4q_f32(outptr0 + out_hstep * 4, _r1); + } + else + { + transpose4x4_ps(_out0, _out2, _out4, _out6); + transpose4x4_ps(_out1, _out3, _out5, _out7); + + if (pC && broadcast_type_C <= 2) + { + _out0 = vaddq_f32(_out0, _c0123); + _out1 = vaddq_f32(_out1, _c0123); + _out2 = vaddq_f32(_out2, _c0123); + _out3 = vaddq_f32(_out3, _c0123); + _out4 = vaddq_f32(_out4, _c0123); + _out5 = vaddq_f32(_out5, _c0123); + _out6 = vaddq_f32(_out6, _c0123); + _out7 = vaddq_f32(_out7, _c0123); + } + + if (alpha != 1.f) + { + _out0 = vmulq_n_f32(_out0, alpha); + _out1 = vmulq_n_f32(_out1, alpha); + _out2 = vmulq_n_f32(_out2, alpha); + _out3 = vmulq_n_f32(_out3, alpha); + _out4 = vmulq_n_f32(_out4, alpha); + _out5 = vmulq_n_f32(_out5, alpha); + _out6 = vmulq_n_f32(_out6, alpha); + _out7 = vmulq_n_f32(_out7, alpha); + } + + vst1q_f32(outptr0, _out0); + vst1q_f32(outptr0 + out_hstep, _out2); + vst1q_f32(outptr0 + out_hstep * 2, _out4); + vst1q_f32(outptr0 + out_hstep * 3, _out6); + vst1q_f32(outptr0 + out_hstep * 4, _out1); + vst1q_f32(outptr0 + out_hstep * 5, _out3); + vst1q_f32(outptr0 + out_hstep * 6, _out5); + vst1q_f32(outptr0 + out_hstep * 7, _out7); + } + + outptr0 += out_hstep * 8; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 8; + pp += 32; + } +#endif // __aarch64__ + for (; jj + 3 < max_jj; jj += 4) + { + float32x4_t _out0 = vld1q_f32(pp + 0); + float32x4_t _out1 = vld1q_f32(pp + 4); + float32x4_t _out2 = vld1q_f32(pp + 8); + float32x4_t _out3 = vld1q_f32(pp + 12); + + if (pC) + { + if (broadcast_type_C == 3) + { + _out0 = vaddq_f32(_out0, (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta))); + _out1 = vaddq_f32(_out1, (beta == 1.f ? vld1q_f32(pC + c_hstep) : vmulq_n_f32(vld1q_f32(pC + c_hstep), beta))); + _out2 = vaddq_f32(_out2, (beta == 1.f ? vld1q_f32(pC + c_hstep * 2) : vmulq_n_f32(vld1q_f32(pC + c_hstep * 2), beta))); + _out3 = vaddq_f32(_out3, (beta == 1.f ? vld1q_f32(pC + c_hstep * 3) : vmulq_n_f32(vld1q_f32(pC + c_hstep * 3), beta))); + } + if (broadcast_type_C == 4) + { + const float32x4_t _cc0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); + _out0 = vaddq_f32(_out0, _cc0); + _out1 = vaddq_f32(_out1, _cc0); + _out2 = vaddq_f32(_out2, _cc0); + _out3 = vaddq_f32(_out3, _cc0); + } + } + + if (out_hstep == 4) + { + if (pC) + { + if (broadcast_type_C == 0) + { + _out0 = vaddq_f32(_out0, _c0123); + _out1 = vaddq_f32(_out1, _c0123); + _out2 = vaddq_f32(_out2, _c0123); + _out3 = vaddq_f32(_out3, _c0123); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _out0 = vaddq_f32(_out0, vdupq_lane_f32(vget_low_f32(_c0123), 0)); + _out1 = vaddq_f32(_out1, vdupq_lane_f32(vget_low_f32(_c0123), 1)); + _out2 = vaddq_f32(_out2, vdupq_lane_f32(vget_high_f32(_c0123), 0)); + _out3 = vaddq_f32(_out3, vdupq_lane_f32(vget_high_f32(_c0123), 1)); + } + } + + if (alpha != 1.f) + { + _out0 = vmulq_n_f32(_out0, alpha); + _out1 = vmulq_n_f32(_out1, alpha); + _out2 = vmulq_n_f32(_out2, alpha); + _out3 = vmulq_n_f32(_out3, alpha); + } + + float32x4x4_t _r; + _r.val[0] = _out0; + _r.val[1] = _out1; + _r.val[2] = _out2; + _r.val[3] = _out3; + vst4q_f32(outptr0, _r); + } + else + { + transpose4x4_ps(_out0, _out1, _out2, _out3); + + if (pC && broadcast_type_C <= 2) + { + _out0 = vaddq_f32(_out0, _c0123); + _out1 = vaddq_f32(_out1, _c0123); + _out2 = vaddq_f32(_out2, _c0123); + _out3 = vaddq_f32(_out3, _c0123); + } + + if (alpha != 1.f) + { + _out0 = vmulq_n_f32(_out0, alpha); + _out1 = vmulq_n_f32(_out1, alpha); + _out2 = vmulq_n_f32(_out2, alpha); + _out3 = vmulq_n_f32(_out3, alpha); + } + + vst1q_f32(outptr0, _out0); + vst1q_f32(outptr0 + out_hstep, _out1); + vst1q_f32(outptr0 + out_hstep * 2, _out2); + vst1q_f32(outptr0 + out_hstep * 3, _out3); + } + + outptr0 += out_hstep * 4; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 4; + pp += 16; + } + for (; jj + 1 < max_jj; jj += 2) + { + float32x4_t _out0 = vld1q_f32(pp); + float32x4_t _out1 = vld1q_f32(pp + 4); + + if (pC) + { + if (broadcast_type_C == 3) + { + const float32x4_t _c01 = vcombine_f32(vld1_f32(pC), vld1_f32(pC + c_hstep)); + const float32x4_t _c23 = vcombine_f32(vld1_f32(pC + c_hstep * 2), vld1_f32(pC + c_hstep * 3)); + _out0 = beta == 1.f ? vaddq_f32(_out0, _c01) : vmlaq_n_f32(_out0, _c01, beta); + _out1 = beta == 1.f ? vaddq_f32(_out1, _c23) : vmlaq_n_f32(_out1, _c23, beta); + } + if (broadcast_type_C == 4) + { + float32x2_t _c = vld1_f32(pC); + if (beta != 1.f) + _c = vmul_n_f32(_c, beta); + const float32x4_t _cc0 = vcombine_f32(_c, _c); + _out0 = vaddq_f32(_out0, _cc0); + _out1 = vaddq_f32(_out1, _cc0); + } + } + + if (out_hstep == 4) + { + if (pC) + { + if (broadcast_type_C == 0) + { + _out0 = vaddq_f32(_out0, _c0123); + _out1 = vaddq_f32(_out1, _c0123); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + float32x4x2_t _c = vzipq_f32(_c0123, _c0123); + _out0 = vaddq_f32(_out0, _c.val[0]); + _out1 = vaddq_f32(_out1, _c.val[1]); + } + } + + if (alpha != 1.f) + { + _out0 = vmulq_n_f32(_out0, alpha); + _out1 = vmulq_n_f32(_out1, alpha); + } + + float32x2x4_t _r; + _r.val[0] = vget_low_f32(_out0); + _r.val[1] = vget_high_f32(_out0); + _r.val[2] = vget_low_f32(_out1); + _r.val[3] = vget_high_f32(_out1); + vst4_f32(outptr0, _r); + } + else + { + float32x4x2_t _t = vuzpq_f32(_out0, _out1); + + if (pC && broadcast_type_C <= 2) + { + _t.val[0] = vaddq_f32(_t.val[0], _c0123); + _t.val[1] = vaddq_f32(_t.val[1], _c0123); + } + + if (alpha != 1.f) + { + _t.val[0] = vmulq_n_f32(_t.val[0], alpha); + _t.val[1] = vmulq_n_f32(_t.val[1], alpha); + } + + vst1q_f32(outptr0, _t.val[0]); + vst1q_f32(outptr0 + out_hstep, _t.val[1]); + } + + outptr0 += out_hstep * 2; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 2; + pp += 8; + } + for (; jj < max_jj; jj += 1) + { + float32x4_t _out0 = vld1q_f32(pp); + + if (pC) + { + if (broadcast_type_C == 0) + { + _out0 = vaddq_f32(_out0, _c0123); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _out0 = vaddq_f32(_out0, _c0123); + } + if (broadcast_type_C == 3) + { + float32x4_t _c = vdupq_n_f32(pC[0]); + _c = vsetq_lane_f32(pC[c_hstep], _c, 1); + _c = vsetq_lane_f32(pC[c_hstep * 2], _c, 2); + _c = vsetq_lane_f32(pC[c_hstep * 3], _c, 3); + _out0 = beta == 1.f ? vaddq_f32(_out0, _c) : vmlaq_n_f32(_out0, _c, beta); + } + if (broadcast_type_C == 4) + { + const float32x4_t _cc0 = vdupq_n_f32(beta == 1.f ? pC[0] : pC[0] * beta); + _out0 = vaddq_f32(_out0, _cc0); + } + } + + if (alpha != 1.f) + { + _out0 = vmulq_n_f32(_out0, alpha); + } + + vst1q_f32(outptr0, _out0); + + outptr0 += out_hstep; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC++; + pp += 4; + } + } + for (; ii + 1 < max_ii; ii += 2) + { + pC = (const float*)C; + float* outptr0 = (float*)top_blob + (size_t)j * out_hstep + i + ii; + float32x4_t _c01 = vdupq_n_f32(0.f); + if (pC) + { + if (broadcast_type_C == 0) + { + float c = pC[0]; + if (beta != 1.f) + c *= beta; + _c01 = vdupq_n_f32(c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + float32x2_t _c = vld1_f32(pC + i + ii); + if (beta != 1.f) + _c = vmul_n_f32(_c, beta); + _c01 = vcombine_f32(_c, _c); + } + if (broadcast_type_C == 3) + pC = (const float*)C + (i + ii) * c_hstep + j; + if (broadcast_type_C == 4) + pC = (const float*)C + j; + } + + int jj = 0; +#if __aarch64__ + for (; jj + 7 < max_jj; jj += 8) + { + float32x4_t _out0 = vld1q_f32(pp + 0); + float32x4_t _out1 = vld1q_f32(pp + 4); + float32x4_t _out2 = vld1q_f32(pp + 8); + float32x4_t _out3 = vld1q_f32(pp + 12); + + if (pC) + { + if (broadcast_type_C == 3) + { + _out0 = vaddq_f32(_out0, (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta))); + _out1 = vaddq_f32(_out1, (beta == 1.f ? vld1q_f32(pC + 4) : vmulq_n_f32(vld1q_f32(pC + 4), beta))); + _out2 = vaddq_f32(_out2, (beta == 1.f ? vld1q_f32(pC + c_hstep) : vmulq_n_f32(vld1q_f32(pC + c_hstep), beta))); + _out3 = vaddq_f32(_out3, (beta == 1.f ? vld1q_f32(pC + c_hstep + 4) : vmulq_n_f32(vld1q_f32(pC + c_hstep + 4), beta))); + } + if (broadcast_type_C == 4) + { + const float32x4_t _c0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); + _out0 = vaddq_f32(_out0, _c0); + _out2 = vaddq_f32(_out2, _c0); + const float32x4_t _c1 = (beta == 1.f ? vld1q_f32(pC + 4) : vmulq_n_f32(vld1q_f32(pC + 4), beta)); + _out1 = vaddq_f32(_out1, _c1); + _out3 = vaddq_f32(_out3, _c1); + } + } + + if (out_hstep == 2) + { + if (pC) + { + if (broadcast_type_C == 0) + { + _out0 = vaddq_f32(_out0, _c01); + _out1 = vaddq_f32(_out1, _c01); + _out2 = vaddq_f32(_out2, _c01); + _out3 = vaddq_f32(_out3, _c01); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + const float32x4_t _c0 = vdupq_lane_f32(vget_low_f32(_c01), 0); + const float32x4_t _c1 = vdupq_lane_f32(vget_low_f32(_c01), 1); + _out0 = vaddq_f32(_out0, _c0); + _out1 = vaddq_f32(_out1, _c0); + _out2 = vaddq_f32(_out2, _c1); + _out3 = vaddq_f32(_out3, _c1); + } + } + + if (alpha != 1.f) + { + _out0 = vmulq_n_f32(_out0, alpha); + _out1 = vmulq_n_f32(_out1, alpha); + _out2 = vmulq_n_f32(_out2, alpha); + _out3 = vmulq_n_f32(_out3, alpha); + } + + float32x4x2_t _r0; + _r0.val[0] = _out0; + _r0.val[1] = _out2; + vst2q_f32(outptr0, _r0); + float32x4x2_t _r1; + _r1.val[0] = _out1; + _r1.val[1] = _out3; + vst2q_f32(outptr0 + out_hstep * 4, _r1); + } + else + { + float32x4x2_t _t0 = vzipq_f32(_out0, _out2); + float32x4x2_t _t1 = vzipq_f32(_out1, _out3); + + if (pC && broadcast_type_C <= 2) + { + _t0.val[0] = vaddq_f32(_t0.val[0], _c01); + _t0.val[1] = vaddq_f32(_t0.val[1], _c01); + _t1.val[0] = vaddq_f32(_t1.val[0], _c01); + _t1.val[1] = vaddq_f32(_t1.val[1], _c01); + } + + if (alpha == 1.f) + { +#if __aarch64__ + vst1q_lane_u64((uint64_t*)(outptr0), vreinterpretq_u64_f32(_t0.val[0]), 0); + vst1q_lane_u64((uint64_t*)(outptr0 + out_hstep), vreinterpretq_u64_f32(_t0.val[0]), 1); + vst1q_lane_u64((uint64_t*)(outptr0 + out_hstep * 2), vreinterpretq_u64_f32(_t0.val[1]), 0); + vst1q_lane_u64((uint64_t*)(outptr0 + out_hstep * 3), vreinterpretq_u64_f32(_t0.val[1]), 1); + vst1q_lane_u64((uint64_t*)(outptr0 + out_hstep * 4), vreinterpretq_u64_f32(_t1.val[0]), 0); + vst1q_lane_u64((uint64_t*)(outptr0 + out_hstep * 5), vreinterpretq_u64_f32(_t1.val[0]), 1); + vst1q_lane_u64((uint64_t*)(outptr0 + out_hstep * 6), vreinterpretq_u64_f32(_t1.val[1]), 0); + vst1q_lane_u64((uint64_t*)(outptr0 + out_hstep * 7), vreinterpretq_u64_f32(_t1.val[1]), 1); +#else + vst1_f32(outptr0, vget_low_f32(_t0.val[0])); + vst1_f32(outptr0 + out_hstep, vget_high_f32(_t0.val[0])); + vst1_f32(outptr0 + out_hstep * 2, vget_low_f32(_t0.val[1])); + vst1_f32(outptr0 + out_hstep * 3, vget_high_f32(_t0.val[1])); + vst1_f32(outptr0 + out_hstep * 4, vget_low_f32(_t1.val[0])); + vst1_f32(outptr0 + out_hstep * 5, vget_high_f32(_t1.val[0])); + vst1_f32(outptr0 + out_hstep * 6, vget_low_f32(_t1.val[1])); + vst1_f32(outptr0 + out_hstep * 7, vget_high_f32(_t1.val[1])); +#endif + } + else + { + _t0.val[0] = vmulq_n_f32(_t0.val[0], alpha); + _t0.val[1] = vmulq_n_f32(_t0.val[1], alpha); + _t1.val[0] = vmulq_n_f32(_t1.val[0], alpha); + _t1.val[1] = vmulq_n_f32(_t1.val[1], alpha); + +#if __aarch64__ + vst1q_lane_u64((uint64_t*)(outptr0), vreinterpretq_u64_f32(_t0.val[0]), 0); + vst1q_lane_u64((uint64_t*)(outptr0 + out_hstep), vreinterpretq_u64_f32(_t0.val[0]), 1); + vst1q_lane_u64((uint64_t*)(outptr0 + out_hstep * 2), vreinterpretq_u64_f32(_t0.val[1]), 0); + vst1q_lane_u64((uint64_t*)(outptr0 + out_hstep * 3), vreinterpretq_u64_f32(_t0.val[1]), 1); + vst1q_lane_u64((uint64_t*)(outptr0 + out_hstep * 4), vreinterpretq_u64_f32(_t1.val[0]), 0); + vst1q_lane_u64((uint64_t*)(outptr0 + out_hstep * 5), vreinterpretq_u64_f32(_t1.val[0]), 1); + vst1q_lane_u64((uint64_t*)(outptr0 + out_hstep * 6), vreinterpretq_u64_f32(_t1.val[1]), 0); + vst1q_lane_u64((uint64_t*)(outptr0 + out_hstep * 7), vreinterpretq_u64_f32(_t1.val[1]), 1); +#else + vst1_f32(outptr0, vget_low_f32(_t0.val[0])); + vst1_f32(outptr0 + out_hstep, vget_high_f32(_t0.val[0])); + vst1_f32(outptr0 + out_hstep * 2, vget_low_f32(_t0.val[1])); + vst1_f32(outptr0 + out_hstep * 3, vget_high_f32(_t0.val[1])); + vst1_f32(outptr0 + out_hstep * 4, vget_low_f32(_t1.val[0])); + vst1_f32(outptr0 + out_hstep * 5, vget_high_f32(_t1.val[0])); + vst1_f32(outptr0 + out_hstep * 6, vget_low_f32(_t1.val[1])); + vst1_f32(outptr0 + out_hstep * 7, vget_high_f32(_t1.val[1])); +#endif + } + } + + outptr0 += out_hstep * 8; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 8; + pp += 16; + } +#endif // __aarch64__ + for (; jj + 3 < max_jj; jj += 4) + { + float32x4_t _out0 = vld1q_f32(pp + 0); + float32x4_t _out1 = vld1q_f32(pp + 4); + + if (pC) + { + if (broadcast_type_C == 3) + { + _out0 = vaddq_f32(_out0, (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta))); + _out1 = vaddq_f32(_out1, (beta == 1.f ? vld1q_f32(pC + c_hstep) : vmulq_n_f32(vld1q_f32(pC + c_hstep), beta))); + } + if (broadcast_type_C == 4) + { + const float32x4_t _c0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); + _out0 = vaddq_f32(_out0, _c0); + _out1 = vaddq_f32(_out1, _c0); + } + } + + if (out_hstep == 2) + { + if (pC) + { + if (broadcast_type_C == 0) + { + _out0 = vaddq_f32(_out0, _c01); + _out1 = vaddq_f32(_out1, _c01); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _out0 = vaddq_f32(_out0, vdupq_lane_f32(vget_low_f32(_c01), 0)); + _out1 = vaddq_f32(_out1, vdupq_lane_f32(vget_low_f32(_c01), 1)); + } + } + + if (alpha != 1.f) + { + _out0 = vmulq_n_f32(_out0, alpha); + _out1 = vmulq_n_f32(_out1, alpha); + } + + float32x4x2_t _r; + _r.val[0] = _out0; + _r.val[1] = _out1; + vst2q_f32(outptr0, _r); + } + else + { + float32x4x2_t _t0 = vzipq_f32(_out0, _out1); + + if (pC && broadcast_type_C <= 2) + { + _t0.val[0] = vaddq_f32(_t0.val[0], _c01); + _t0.val[1] = vaddq_f32(_t0.val[1], _c01); + } + + if (alpha == 1.f) + { +#if __aarch64__ + vst1q_lane_u64((uint64_t*)(outptr0), vreinterpretq_u64_f32(_t0.val[0]), 0); + vst1q_lane_u64((uint64_t*)(outptr0 + out_hstep), vreinterpretq_u64_f32(_t0.val[0]), 1); + vst1q_lane_u64((uint64_t*)(outptr0 + out_hstep * 2), vreinterpretq_u64_f32(_t0.val[1]), 0); + vst1q_lane_u64((uint64_t*)(outptr0 + out_hstep * 3), vreinterpretq_u64_f32(_t0.val[1]), 1); +#else + vst1_f32(outptr0, vget_low_f32(_t0.val[0])); + vst1_f32(outptr0 + out_hstep, vget_high_f32(_t0.val[0])); + vst1_f32(outptr0 + out_hstep * 2, vget_low_f32(_t0.val[1])); + vst1_f32(outptr0 + out_hstep * 3, vget_high_f32(_t0.val[1])); +#endif + } + else + { + _t0.val[0] = vmulq_n_f32(_t0.val[0], alpha); + _t0.val[1] = vmulq_n_f32(_t0.val[1], alpha); + +#if __aarch64__ + vst1q_lane_u64((uint64_t*)(outptr0), vreinterpretq_u64_f32(_t0.val[0]), 0); + vst1q_lane_u64((uint64_t*)(outptr0 + out_hstep), vreinterpretq_u64_f32(_t0.val[0]), 1); + vst1q_lane_u64((uint64_t*)(outptr0 + out_hstep * 2), vreinterpretq_u64_f32(_t0.val[1]), 0); + vst1q_lane_u64((uint64_t*)(outptr0 + out_hstep * 3), vreinterpretq_u64_f32(_t0.val[1]), 1); +#else + vst1_f32(outptr0, vget_low_f32(_t0.val[0])); + vst1_f32(outptr0 + out_hstep, vget_high_f32(_t0.val[0])); + vst1_f32(outptr0 + out_hstep * 2, vget_low_f32(_t0.val[1])); + vst1_f32(outptr0 + out_hstep * 3, vget_high_f32(_t0.val[1])); +#endif + } + } + + outptr0 += out_hstep * 4; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 4; + pp += 8; + } + for (; jj + 1 < max_jj; jj += 2) + { + float32x2_t _out0 = vld1_f32(pp); + float32x2_t _out1 = vld1_f32(pp + 2); + + if (pC) + { + if (broadcast_type_C == 3) + { + const float32x2_t _c0 = vld1_f32(pC); + const float32x2_t _c1 = vld1_f32(pC + c_hstep); + _out0 = beta == 1.f ? vadd_f32(_out0, _c0) : vmla_n_f32(_out0, _c0, beta); + _out1 = beta == 1.f ? vadd_f32(_out1, _c1) : vmla_n_f32(_out1, _c1, beta); + } + if (broadcast_type_C == 4) + { + const float32x2_t _c0 = beta == 1.f ? vld1_f32(pC) : vmul_n_f32(vld1_f32(pC), beta); + _out0 = vadd_f32(_out0, _c0); + _out1 = vadd_f32(_out1, _c0); + } + } + + if (out_hstep == 2) + { + if (pC) + { + if (broadcast_type_C == 0) + { + _out0 = vadd_f32(_out0, vget_low_f32(_c01)); + _out1 = vadd_f32(_out1, vget_low_f32(_c01)); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _out0 = vadd_f32(_out0, vdup_lane_f32(vget_low_f32(_c01), 0)); + _out1 = vadd_f32(_out1, vdup_lane_f32(vget_low_f32(_c01), 1)); + } + } + + if (alpha != 1.f) + { + _out0 = vmul_n_f32(_out0, alpha); + _out1 = vmul_n_f32(_out1, alpha); + } + + float32x2x2_t _r; + _r.val[0] = _out0; + _r.val[1] = _out1; + vst2_f32(outptr0, _r); + } + else + { + float32x2x2_t _t0 = vzip_f32(_out0, _out1); + + if (pC && broadcast_type_C <= 2) + { + _t0.val[0] = vadd_f32(_t0.val[0], vget_low_f32(_c01)); + _t0.val[1] = vadd_f32(_t0.val[1], vget_low_f32(_c01)); + } + + if (alpha == 1.f) + { + vst1_f32(outptr0, _t0.val[0]); + vst1_f32(outptr0 + out_hstep, _t0.val[1]); + } + else + { + _t0.val[0] = vmul_n_f32(_t0.val[0], alpha); + _t0.val[1] = vmul_n_f32(_t0.val[1], alpha); + + vst1_f32(outptr0, _t0.val[0]); + vst1_f32(outptr0 + out_hstep, _t0.val[1]); + } + } + + outptr0 += out_hstep * 2; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 2; + pp += 4; + } + for (; jj < max_jj; jj += 1) + { + float32x2_t _out0 = vld1_f32(pp); + + if (pC) + { + if (broadcast_type_C == 0) + { + _out0 = vadd_f32(_out0, vget_low_f32(_c01)); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _out0 = vadd_f32(_out0, vget_low_f32(_c01)); + } + if (broadcast_type_C == 3) + { + float32x2_t _c = vdup_n_f32(pC[0]); + _c = vset_lane_f32(pC[c_hstep], _c, 1); + _out0 = beta == 1.f ? vadd_f32(_out0, _c) : vmla_n_f32(_out0, _c, beta); + } + if (broadcast_type_C == 4) + { + _out0 = vadd_f32(_out0, vdup_n_f32(beta == 1.f ? pC[0] : pC[0] * beta)); + } + } + + if (alpha != 1.f) + { + _out0 = vmul_n_f32(_out0, alpha); + } + + vst1_f32(outptr0, _out0); + + outptr0 += out_hstep; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC++; + pp += 2; + } + } + for (; ii < max_ii; ii++) + { + pC = (const float*)C; + float* outptr0 = (float*)top_blob + (size_t)j * out_hstep + i + ii; + float32x4_t _c0; + float c0 = 0.f; + if (pC) + { + if (broadcast_type_C == 0) + { + c0 = pC[0]; + if (beta != 1.f) + c0 *= beta; + _c0 = vdupq_n_f32(c0); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + c0 = pC[i + ii]; + if (beta != 1.f) + c0 *= beta; + _c0 = vdupq_n_f32(c0); + } + if (broadcast_type_C == 3) + pC = (const float*)C + (i + ii) * c_hstep + j; + if (broadcast_type_C == 4) + pC = (const float*)C + j; + } + + int jj = 0; +#if __aarch64__ + for (; jj + 7 < max_jj; jj += 8) + { + float32x4_t _out0 = vld1q_f32(pp + 0); + float32x4_t _out1 = vld1q_f32(pp + 4); + + if (pC) + { + if (broadcast_type_C == 0) + { + _out0 = vaddq_f32(_out0, _c0); + _out1 = vaddq_f32(_out1, _c0); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _out0 = vaddq_f32(_out0, _c0); + _out1 = vaddq_f32(_out1, _c0); + } + if (broadcast_type_C == 3) + { + _out0 = vaddq_f32(_out0, (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta))); + _out1 = vaddq_f32(_out1, (beta == 1.f ? vld1q_f32(pC + 4) : vmulq_n_f32(vld1q_f32(pC + 4), beta))); + } + if (broadcast_type_C == 4) + { + const float32x4_t _c0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); + _out0 = vaddq_f32(_out0, _c0); + const float32x4_t _c1 = (beta == 1.f ? vld1q_f32(pC + 4) : vmulq_n_f32(vld1q_f32(pC + 4), beta)); + _out1 = vaddq_f32(_out1, _c1); + } + } + + if (alpha != 1.f) + { + _out0 = vmulq_n_f32(_out0, alpha); + _out1 = vmulq_n_f32(_out1, alpha); + } + + if (out_hstep == 1) + { + vst1q_f32(outptr0, _out0); + vst1q_f32(outptr0 + 4, _out1); + } + else + { + vst1q_lane_f32(outptr0, _out0, 0); + vst1q_lane_f32(outptr0 + out_hstep, _out0, 1); + vst1q_lane_f32(outptr0 + out_hstep * 2, _out0, 2); + vst1q_lane_f32(outptr0 + out_hstep * 3, _out0, 3); + vst1q_lane_f32(outptr0 + out_hstep * 4, _out1, 0); + vst1q_lane_f32(outptr0 + out_hstep * 5, _out1, 1); + vst1q_lane_f32(outptr0 + out_hstep * 6, _out1, 2); + vst1q_lane_f32(outptr0 + out_hstep * 7, _out1, 3); + } + + outptr0 += out_hstep * 8; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 8; + pp += 8; + } +#endif // __aarch64__ + for (; jj + 3 < max_jj; jj += 4) + { + float32x4_t _out0 = vld1q_f32(pp + 0); + + if (pC) + { + if (broadcast_type_C == 0) + { + _out0 = vaddq_f32(_out0, _c0); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _out0 = vaddq_f32(_out0, _c0); + } + if (broadcast_type_C == 3) + { + _out0 = vaddq_f32(_out0, (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta))); + } + if (broadcast_type_C == 4) + { + const float32x4_t _c0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); + _out0 = vaddq_f32(_out0, _c0); + } + } + + if (alpha != 1.f) + { + _out0 = vmulq_n_f32(_out0, alpha); + } + + if (out_hstep == 1) + { + vst1q_f32(outptr0, _out0); + } + else + { + vst1q_lane_f32(outptr0, _out0, 0); + vst1q_lane_f32(outptr0 + out_hstep, _out0, 1); + vst1q_lane_f32(outptr0 + out_hstep * 2, _out0, 2); + vst1q_lane_f32(outptr0 + out_hstep * 3, _out0, 3); + } + + outptr0 += out_hstep * 4; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 4; + pp += 4; + } + for (; jj + 1 < max_jj; jj += 2) + { + float32x2_t _out0 = vld1_f32(pp); + + if (pC) + { + if (broadcast_type_C == 0) + { + _out0 = vadd_f32(_out0, vdup_n_f32(c0)); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _out0 = vadd_f32(_out0, vdup_n_f32(c0)); + } + if (broadcast_type_C == 3) + { + const float32x2_t _c0 = vld1_f32(pC); + _out0 = beta == 1.f ? vadd_f32(_out0, _c0) : vmla_n_f32(_out0, _c0, beta); + } + if (broadcast_type_C == 4) + { + const float32x2_t _c0 = beta == 1.f ? vld1_f32(pC) : vmul_n_f32(vld1_f32(pC), beta); + _out0 = vadd_f32(_out0, _c0); + } + } + + if (alpha != 1.f) + { + _out0 = vmul_n_f32(_out0, alpha); + } + + if (out_hstep == 1) + { + vst1_f32(outptr0, _out0); + } + else + { + vst1_lane_f32(outptr0, _out0, 0); + vst1_lane_f32(outptr0 + out_hstep, _out0, 1); + } + + outptr0 += out_hstep * 2; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 2; + pp += 2; + } + for (; jj < max_jj; jj += 1) + { + float out0 = pp[0]; + + if (pC) + { + if (broadcast_type_C == 0) + { + out0 += c0; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + out0 += c0; + } + if (broadcast_type_C == 3) + { + out0 += beta == 1.f ? pC[0] : pC[0] * beta; + } + if (broadcast_type_C == 4) + { + out0 += beta == 1.f ? pC[0] : pC[0] * beta; + } + } + + if (alpha != 1.f) + { + out0 *= alpha; + } + + outptr0[0] = out0; + + outptr0 += out_hstep; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC++; + pp += 1; + } + } +#else + for (; ii + 1 < max_ii; ii += 2) + { + pC = (const float*)C; + float* outptr0 = (float*)top_blob + (size_t)j * out_hstep + i + ii; + float c0 = 0.f; + float c1 = 0.f; + if (pC) + { + if (broadcast_type_C == 0) + { + float c = pC[0]; + if (beta != 1.f) + c *= beta; + c0 = c; + c1 = c; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + c0 = pC[i + ii]; + if (beta != 1.f) + c0 *= beta; + c1 = pC[i + ii + 1]; + if (beta != 1.f) + c1 *= beta; + } + if (broadcast_type_C == 3) + pC = (const float*)C + (i + ii) * c_hstep + j; + if (broadcast_type_C == 4) + pC = (const float*)C + j; + } + + int jj = 0; + for (; jj + 1 < max_jj; jj += 2) + { + float out00 = pp[0]; + float out01 = pp[1]; + float out10 = pp[2]; + float out11 = pp[3]; + + if (pC) + { + if (broadcast_type_C == 0) + { + out00 += c0; + out01 += c0; + out10 += c0; + out11 += c0; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + out00 += c0; + out01 += c0; + out10 += c1; + out11 += c1; + } + if (broadcast_type_C == 3) + { + out00 += beta == 1.f ? pC[0] : pC[0] * beta; + out01 += beta == 1.f ? pC[1] : pC[1] * beta; + out10 += beta == 1.f ? pC[c_hstep] : pC[c_hstep] * beta; + out11 += beta == 1.f ? pC[c_hstep + 1] : pC[c_hstep + 1] * beta; + } + if (broadcast_type_C == 4) + { + out00 += beta == 1.f ? pC[0] : pC[0] * beta; + out01 += beta == 1.f ? pC[1] : pC[1] * beta; + out10 += beta == 1.f ? pC[0] : pC[0] * beta; + out11 += beta == 1.f ? pC[1] : pC[1] * beta; + } + } + + if (alpha != 1.f) + { + out00 *= alpha; + out01 *= alpha; + out10 *= alpha; + out11 *= alpha; + } + + outptr0[0] = out00; + outptr0[out_hstep] = out01; + outptr0[1] = out10; + outptr0[out_hstep + 1] = out11; + + outptr0 += out_hstep * 2; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 2; + pp += 4; + } + for (; jj < max_jj; jj += 1) + { + float out00 = pp[0]; + float out10 = pp[1]; + + if (pC) + { + if (broadcast_type_C == 0) + { + out00 += c0; + out10 += c0; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + out00 += c0; + out10 += c1; + } + if (broadcast_type_C == 3) + { + out00 += beta == 1.f ? pC[0] : pC[0] * beta; + out10 += beta == 1.f ? pC[c_hstep] : pC[c_hstep] * beta; + } + if (broadcast_type_C == 4) + { + out00 += beta == 1.f ? pC[0] : pC[0] * beta; + out10 += beta == 1.f ? pC[0] : pC[0] * beta; + } + } + + if (alpha != 1.f) + { + out00 *= alpha; + out10 *= alpha; + } + + outptr0[0] = out00; + outptr0[1] = out10; + + outptr0 += out_hstep; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC++; + pp += 2; + } + } + for (; ii < max_ii; ii++) + { + pC = (const float*)C; + float* outptr0 = (float*)top_blob + (size_t)j * out_hstep + i + ii; + float c0 = 0.f; + if (pC) + { + if (broadcast_type_C == 0) + { + float c = pC[0]; + if (beta != 1.f) + c *= beta; + c0 = c; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + c0 = pC[i + ii]; + if (beta != 1.f) + c0 *= beta; + } + if (broadcast_type_C == 3) + pC = (const float*)C + (i + ii) * c_hstep + j; + if (broadcast_type_C == 4) + pC = (const float*)C + j; + } + + int jj = 0; + for (; jj + 1 < max_jj; jj += 2) + { + float out00 = pp[0]; + float out01 = pp[1]; + + if (pC) + { + if (broadcast_type_C == 0) + { + out00 += c0; + out01 += c0; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + out00 += c0; + out01 += c0; + } + if (broadcast_type_C == 3) + { + out00 += beta == 1.f ? pC[0] : pC[0] * beta; + out01 += beta == 1.f ? pC[1] : pC[1] * beta; + } + if (broadcast_type_C == 4) + { + out00 += beta == 1.f ? pC[0] : pC[0] * beta; + out01 += beta == 1.f ? pC[1] : pC[1] * beta; + } + } + + if (alpha != 1.f) + { + out00 *= alpha; + out01 *= alpha; + } + + outptr0[0] = out00; + outptr0[out_hstep] = out01; + + outptr0 += out_hstep * 2; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 2; + pp += 2; + } + for (; jj < max_jj; jj += 1) + { + float out00 = pp[0]; + + if (pC) + { + if (broadcast_type_C == 0) + { + out00 += c0; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + out00 += c0; + } + if (broadcast_type_C == 3) + { + out00 += beta == 1.f ? pC[0] : pC[0] * beta; + } + if (broadcast_type_C == 4) + { + out00 += beta == 1.f ? pC[0] : pC[0] * beta; + } + } + + if (alpha != 1.f) + { + out00 *= alpha; + } + + outptr0[0] = out00; + + outptr0 += out_hstep; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC++; + pp += 1; + } + } +#endif // __ARM_NEON +} diff --git a/src/layer/arm/multiheadattention_arm.cpp b/src/layer/arm/multiheadattention_arm.cpp index a3bbaf98785f..6fa33f4386e2 100644 --- a/src/layer/arm/multiheadattention_arm.cpp +++ b/src/layer/arm/multiheadattention_arm.cpp @@ -32,10 +32,360 @@ MultiHeadAttention_arm::MultiHeadAttention_arm() qk_softmax = 0; } +#if NCNN_WEIGHT_QUANT +int MultiHeadAttention_arm::create_pipeline_wq_int8(const Option& _opt) +{ + if (q_gemm) + return 0; + + Option opt = _opt; + opt.use_fp16_storage &= support_fp16_storage; + opt.use_bf16_storage &= support_bf16_storage; + if (opt.use_fp16_storage) + opt.use_bf16_storage = false; + + Option opt_wq = opt; + opt_wq.use_packing_layout = false; + opt_wq.use_fp16_packed = false; + opt_wq.use_fp16_storage = false; + opt_wq.use_fp16_arithmetic = false; + opt_wq.use_bf16_packed = false; + opt_wq.use_bf16_storage = false; + + { + qk_softmax = ncnn::create_layer_cpu(ncnn::LayerType::Softmax); + if (!qk_softmax) + return -100; + ncnn::ParamDict pd; + pd.set(0, -1); + pd.set(1, 1); + int ret = qk_softmax->load_param(pd); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = qk_softmax->load_model(ModelBinFromMatArray(0)); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = qk_softmax->create_pipeline(opt); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + } + + const int qdim = weight_data_size / embed_dim; + + { + q_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!q_gemm) + { + destroy_pipeline(_opt); + return -100; + } + ncnn::ParamDict pd; + pd.set(0, scale); + pd.set(1, 1.f); + pd.set(2, 0); // transA + pd.set(3, 1); // transB + pd.set(4, 0); // constantA + pd.set(5, 1); // constantB + pd.set(6, 1); // constantC + pd.set(7, 0); // M + pd.set(8, embed_dim); // N + pd.set(9, qdim); // K + pd.set(10, 4); // constant_broadcast_type_C = null + pd.set(11, 0); // output_N1M + pd.set(12, 0); // output_elempack + pd.set(14, 1); // output_transpose + pd.set(18, quantize_term); + int ret = q_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + Mat weights[4]; + weights[0] = q_weight_data; + weights[1] = q_bias_data; + weights[2] = q_weight_data_quantize_scales; + weights[3] = q_weight_data_input_scales; + ret = q_gemm->load_model(ModelBinFromMatArray(weights)); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = q_gemm->create_pipeline(opt_wq); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + } + + { + k_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!k_gemm) + { + destroy_pipeline(_opt); + return -100; + } + ncnn::ParamDict pd; + pd.set(2, 0); // transA + pd.set(3, 1); // transB + pd.set(4, 0); // constantA + pd.set(5, 1); // constantB + pd.set(6, 1); // constantC + pd.set(7, 0); // M + pd.set(8, embed_dim); // N + pd.set(9, kdim); // K + pd.set(10, 4); // constant_broadcast_type_C = null + pd.set(11, 0); // output_N1M + pd.set(12, 0); // output_elempack + pd.set(14, 1); // output_transpose + pd.set(18, quantize_term); + int ret = k_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + Mat weights[4]; + weights[0] = k_weight_data; + weights[1] = k_bias_data; + weights[2] = k_weight_data_quantize_scales; + weights[3] = k_weight_data_input_scales; + ret = k_gemm->load_model(ModelBinFromMatArray(weights)); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = k_gemm->create_pipeline(opt_wq); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + } + + { + v_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!v_gemm) + { + destroy_pipeline(_opt); + return -100; + } + ncnn::ParamDict pd; + pd.set(2, 0); // transA + pd.set(3, 1); // transB + pd.set(4, 0); // constantA + pd.set(5, 1); // constantB + pd.set(6, 1); // constantC + pd.set(7, 0); // M + pd.set(8, embed_dim); // N + pd.set(9, vdim); // K + pd.set(10, 4); // constant_broadcast_type_C = null + pd.set(11, 0); // output_N1M + pd.set(12, 0); // output_elempack + pd.set(14, 1); // output_transpose + pd.set(18, quantize_term); + int ret = v_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + Mat weights[4]; + weights[0] = v_weight_data; + weights[1] = v_bias_data; + weights[2] = v_weight_data_quantize_scales; + weights[3] = v_weight_data_input_scales; + ret = v_gemm->load_model(ModelBinFromMatArray(weights)); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = v_gemm->create_pipeline(opt_wq); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + } + + { + o_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!o_gemm) + { + destroy_pipeline(_opt); + return -100; + } + ncnn::ParamDict pd; + pd.set(2, 1); // transA + pd.set(3, 1); // transB + pd.set(4, 0); // constantA + pd.set(5, 1); // constantB + pd.set(6, 1); // constantC + pd.set(7, 0); // M = outch + pd.set(8, qdim); // N = size + pd.set(9, embed_dim); // K = maxk*inch + pd.set(10, 4); // constant_broadcast_type_C = null + pd.set(11, 0); // output_N1M + pd.set(18, quantize_term); + int ret = o_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + Mat weights[4]; + weights[0] = out_weight_data; + weights[1] = out_bias_data; + weights[2] = out_weight_data_quantize_scales; + weights[3] = out_weight_data_input_scales; + ret = o_gemm->load_model(ModelBinFromMatArray(weights)); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = o_gemm->create_pipeline(opt_wq); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + } + + { + qk_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!qk_gemm) + { + destroy_pipeline(_opt); + return -100; + } + ncnn::ParamDict pd; + pd.set(2, 1); // transA + pd.set(3, 0); // transB + pd.set(4, 0); // constantA + pd.set(5, 0); // constantB + pd.set(6, attn_mask ? 0 : 1); // constantC + pd.set(7, 0); // M + pd.set(8, 0); // N + pd.set(9, 0); // K + pd.set(10, attn_mask ? 3 : -1); // constant_broadcast_type_C + pd.set(11, 0); // output_N1M + pd.set(12, 1); // output_elempack + pd.set(13, 1); // output_elemtype = fp32 + int ret = qk_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = qk_gemm->load_model(ModelBinFromMatArray(0)); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + Option opt1 = opt; + opt1.use_bf16_packed = false; + opt1.use_bf16_storage = false; + opt1.num_threads = 1; + ret = qk_gemm->create_pipeline(opt1); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + } + + { + qkv_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!qkv_gemm) + { + destroy_pipeline(_opt); + return -100; + } + ncnn::ParamDict pd; + pd.set(2, 0); // transA + pd.set(3, 1); // transB + pd.set(4, 0); // constantA + pd.set(5, 0); // constantB + pd.set(6, 1); // constantC + pd.set(7, 0); // M + pd.set(8, 0); // N + pd.set(9, 0); // K + pd.set(10, -1); // constant_broadcast_type_C + pd.set(11, 0); // output_N1M + pd.set(12, 1); // output_elempack + pd.set(13, 1); // output_elemtype = fp32 + pd.set(14, 1); // output_transpose + int ret = qkv_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = qkv_gemm->load_model(ModelBinFromMatArray(0)); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + Option opt1 = opt; + opt1.use_bf16_packed = false; + opt1.use_bf16_storage = false; + opt1.num_threads = 1; + ret = qkv_gemm->create_pipeline(opt1); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + } + + q_weight_data.release(); + q_bias_data.release(); + k_weight_data.release(); + k_bias_data.release(); + v_weight_data.release(); + v_bias_data.release(); + out_weight_data.release(); + out_bias_data.release(); + q_weight_data_quantize_scales.release(); + k_weight_data_quantize_scales.release(); + v_weight_data_quantize_scales.release(); + out_weight_data_quantize_scales.release(); + q_weight_data_input_scales.release(); + k_weight_data_input_scales.release(); + v_weight_data_input_scales.release(); + out_weight_data_input_scales.release(); + + return 0; +} +#endif // NCNN_WEIGHT_QUANT + int MultiHeadAttention_arm::create_pipeline(const Option& _opt) { +#if NCNN_WEIGHT_QUANT if (weight_block_quantize) - return 0; + { + if (quantize_term / 100 != 8) + return MultiHeadAttention::create_pipeline(_opt); + + return create_pipeline_wq_int8(_opt); + } +#endif Option opt = _opt; if (int8_scale_term) @@ -53,18 +403,40 @@ int MultiHeadAttention_arm::create_pipeline(const Option& _opt) { qk_softmax = ncnn::create_layer_cpu(ncnn::LayerType::Softmax); + if (!qk_softmax) + return -100; ncnn::ParamDict pd; pd.set(0, -1); pd.set(1, 1); - qk_softmax->load_param(pd); - qk_softmax->load_model(ModelBinFromMatArray(0)); - qk_softmax->create_pipeline(opt); + int ret = qk_softmax->load_param(pd); + if (ret != 0) + { + destroy_pipeline(opt); + return ret; + } + ret = qk_softmax->load_model(ModelBinFromMatArray(0)); + if (ret != 0) + { + destroy_pipeline(opt); + return ret; + } + ret = qk_softmax->create_pipeline(opt); + if (ret != 0) + { + destroy_pipeline(opt); + return ret; + } } const int qdim = weight_data_size / embed_dim; { q_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!q_gemm) + { + destroy_pipeline(opt); + return -100; + } ncnn::ParamDict pd; pd.set(0, scale); pd.set(1, 1.f); @@ -83,25 +455,39 @@ int MultiHeadAttention_arm::create_pipeline(const Option& _opt) #if NCNN_INT8 pd.set(18, int8_scale_term); #endif - q_gemm->load_param(pd); + int ret = q_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(opt); + return ret; + } Mat weights[3]; weights[0] = q_weight_data; weights[1] = q_bias_data; #if NCNN_INT8 weights[2] = q_weight_data_int8_scales; #endif - q_gemm->load_model(ModelBinFromMatArray(weights)); - q_gemm->create_pipeline(opt); - - if (opt.lightmode) + ret = q_gemm->load_model(ModelBinFromMatArray(weights)); + if (ret != 0) { - q_weight_data.release(); - q_bias_data.release(); + destroy_pipeline(opt); + return ret; + } + ret = q_gemm->create_pipeline(opt); + if (ret != 0) + { + destroy_pipeline(opt); + return ret; } } { k_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!k_gemm) + { + destroy_pipeline(opt); + return -100; + } ncnn::ParamDict pd; pd.set(2, 0); // transA pd.set(3, 1); // transB @@ -118,25 +504,39 @@ int MultiHeadAttention_arm::create_pipeline(const Option& _opt) #if NCNN_INT8 pd.set(18, int8_scale_term); #endif - k_gemm->load_param(pd); + int ret = k_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(opt); + return ret; + } Mat weights[3]; weights[0] = k_weight_data; weights[1] = k_bias_data; #if NCNN_INT8 weights[2] = k_weight_data_int8_scales; #endif - k_gemm->load_model(ModelBinFromMatArray(weights)); - k_gemm->create_pipeline(opt); - - if (opt.lightmode) + ret = k_gemm->load_model(ModelBinFromMatArray(weights)); + if (ret != 0) { - k_weight_data.release(); - k_bias_data.release(); + destroy_pipeline(opt); + return ret; + } + ret = k_gemm->create_pipeline(opt); + if (ret != 0) + { + destroy_pipeline(opt); + return ret; } } { v_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!v_gemm) + { + destroy_pipeline(opt); + return -100; + } ncnn::ParamDict pd; pd.set(2, 0); // transA pd.set(3, 1); // transB @@ -153,25 +553,39 @@ int MultiHeadAttention_arm::create_pipeline(const Option& _opt) #if NCNN_INT8 pd.set(18, int8_scale_term); #endif - v_gemm->load_param(pd); + int ret = v_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(opt); + return ret; + } Mat weights[3]; weights[0] = v_weight_data; weights[1] = v_bias_data; #if NCNN_INT8 weights[2] = v_weight_data_int8_scales; #endif - v_gemm->load_model(ModelBinFromMatArray(weights)); - v_gemm->create_pipeline(opt); - - if (opt.lightmode) + ret = v_gemm->load_model(ModelBinFromMatArray(weights)); + if (ret != 0) { - v_weight_data.release(); - v_bias_data.release(); + destroy_pipeline(opt); + return ret; + } + ret = v_gemm->create_pipeline(opt); + if (ret != 0) + { + destroy_pipeline(opt); + return ret; } } { o_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!o_gemm) + { + destroy_pipeline(opt); + return -100; + } ncnn::ParamDict pd; pd.set(2, 1); // transA pd.set(3, 1); // transB @@ -186,30 +600,49 @@ int MultiHeadAttention_arm::create_pipeline(const Option& _opt) #if NCNN_INT8 pd.set(18, int8_scale_term); #endif - o_gemm->load_param(pd); + int ret = o_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(opt); + return ret; + } Mat weights[3]; weights[0] = out_weight_data; weights[1] = out_bias_data; #if NCNN_INT8 Mat out_weight_data_int8_scales(1); + if (out_weight_data_int8_scales.empty()) + { + destroy_pipeline(opt); + return -100; + } out_weight_data_int8_scales[0] = out_weight_data_int8_scale; weights[2] = out_weight_data_int8_scales; #endif - o_gemm->load_model(ModelBinFromMatArray(weights)); + ret = o_gemm->load_model(ModelBinFromMatArray(weights)); + if (ret != 0) + { + destroy_pipeline(opt); + return ret; + } Option opt_fp32 = opt; opt_fp32.use_bf16_packed = false; opt_fp32.use_bf16_storage = false; - o_gemm->create_pipeline(opt_fp32); - - if (opt.lightmode) + ret = o_gemm->create_pipeline(opt_fp32); + if (ret != 0) { - out_weight_data.release(); - out_bias_data.release(); + destroy_pipeline(opt); + return ret; } } { qk_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!qk_gemm) + { + destroy_pipeline(opt); + return -100; + } ncnn::ParamDict pd; pd.set(2, 1); // transA pd.set(3, 0); // transB @@ -226,17 +659,37 @@ int MultiHeadAttention_arm::create_pipeline(const Option& _opt) #if NCNN_INT8 pd.set(18, int8_scale_term); #endif - qk_gemm->load_param(pd); - qk_gemm->load_model(ModelBinFromMatArray(0)); + int ret = qk_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(opt); + return ret; + } + ret = qk_gemm->load_model(ModelBinFromMatArray(0)); + if (ret != 0) + { + destroy_pipeline(opt); + return ret; + } Option opt1 = opt; opt1.use_bf16_packed = false; opt1.use_bf16_storage = false; opt1.num_threads = 1; - qk_gemm->create_pipeline(opt1); + ret = qk_gemm->create_pipeline(opt1); + if (ret != 0) + { + destroy_pipeline(opt); + return ret; + } } { qkv_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!qkv_gemm) + { + destroy_pipeline(opt); + return -100; + } ncnn::ParamDict pd; pd.set(2, 0); // transA pd.set(3, 1); // transB @@ -254,13 +707,40 @@ int MultiHeadAttention_arm::create_pipeline(const Option& _opt) #if NCNN_INT8 pd.set(18, int8_scale_term); #endif - qkv_gemm->load_param(pd); - qkv_gemm->load_model(ModelBinFromMatArray(0)); + int ret = qkv_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(opt); + return ret; + } + ret = qkv_gemm->load_model(ModelBinFromMatArray(0)); + if (ret != 0) + { + destroy_pipeline(opt); + return ret; + } Option opt1 = opt; opt1.use_bf16_packed = false; opt1.use_bf16_storage = false; opt1.num_threads = 1; - qkv_gemm->create_pipeline(opt1); + ret = qkv_gemm->create_pipeline(opt1); + if (ret != 0) + { + destroy_pipeline(opt); + return ret; + } + } + + if (opt.lightmode) + { + q_weight_data.release(); + q_bias_data.release(); + k_weight_data.release(); + k_bias_data.release(); + v_weight_data.release(); + v_bias_data.release(); + out_weight_data.release(); + out_bias_data.release(); } return 0; @@ -268,11 +748,11 @@ int MultiHeadAttention_arm::create_pipeline(const Option& _opt) int MultiHeadAttention_arm::destroy_pipeline(const Option& _opt) { - if (weight_block_quantize) - return 0; + if (weight_block_quantize && quantize_term / 100 != 8) + return MultiHeadAttention::destroy_pipeline(_opt); Option opt = _opt; - if (int8_scale_term) + if (int8_scale_term && !weight_block_quantize) { opt.use_packing_layout = false; // TODO enable packing } @@ -282,6 +762,17 @@ int MultiHeadAttention_arm::destroy_pipeline(const Option& _opt) if (opt.use_fp16_storage) opt.use_bf16_storage = false; + Option opt_wq = opt; + if (weight_block_quantize) + { + opt_wq.use_packing_layout = false; + opt_wq.use_fp16_packed = false; + opt_wq.use_fp16_storage = false; + opt_wq.use_fp16_arithmetic = false; + opt_wq.use_bf16_packed = false; + opt_wq.use_bf16_storage = false; + } + if (qk_softmax) { qk_softmax->destroy_pipeline(opt); @@ -291,28 +782,28 @@ int MultiHeadAttention_arm::destroy_pipeline(const Option& _opt) if (q_gemm) { - q_gemm->destroy_pipeline(opt); + q_gemm->destroy_pipeline(opt_wq); delete q_gemm; q_gemm = 0; } if (k_gemm) { - k_gemm->destroy_pipeline(opt); + k_gemm->destroy_pipeline(opt_wq); delete k_gemm; k_gemm = 0; } if (v_gemm) { - v_gemm->destroy_pipeline(opt); + v_gemm->destroy_pipeline(opt_wq); delete v_gemm; v_gemm = 0; } if (o_gemm) { - o_gemm->destroy_pipeline(opt); + o_gemm->destroy_pipeline(opt_wq); delete o_gemm; o_gemm = 0; } @@ -336,7 +827,7 @@ int MultiHeadAttention_arm::destroy_pipeline(const Option& _opt) int MultiHeadAttention_arm::forward(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& _opt) const { - if (weight_block_quantize) + if (weight_block_quantize && quantize_term / 100 != 8) return MultiHeadAttention::forward(bottom_blobs, top_blobs, _opt); int q_blob_i = 0; @@ -355,7 +846,7 @@ int MultiHeadAttention_arm::forward(const std::vector& bottom_blobs, std::v const Mat& cached_xv_blob = kv_cache ? bottom_blobs[cached_xv_i] : Mat(); Option opt = _opt; - if (int8_scale_term) + if (int8_scale_term && !weight_block_quantize) { opt.use_packing_layout = false; // TODO enable packing } @@ -365,6 +856,17 @@ int MultiHeadAttention_arm::forward(const std::vector& bottom_blobs, std::v if (opt.use_fp16_storage) opt.use_bf16_storage = false; + Option opt_wq = opt; + if (weight_block_quantize) + { + opt_wq.use_packing_layout = false; + opt_wq.use_fp16_packed = false; + opt_wq.use_fp16_storage = false; + opt_wq.use_fp16_arithmetic = false; + opt_wq.use_bf16_packed = false; + opt_wq.use_bf16_storage = false; + } + Mat attn_mask_blob_unpacked; if (attn_mask && attn_mask_blob.elempack != 1) { @@ -413,7 +915,7 @@ int MultiHeadAttention_arm::forward(const std::vector& bottom_blobs, std::v size_t workspace_elemsize = opt.use_bf16_storage ? 4u : elemsize; Mat q_affine; - int retq = q_gemm->forward(q_blob, q_affine, opt); + int retq = q_gemm->forward(q_blob, q_affine, opt_wq); if (retq != 0) return retq; @@ -423,7 +925,7 @@ int MultiHeadAttention_arm::forward(const std::vector& bottom_blobs, std::v if (q_blob_i == k_blob_i) { Mat k_affine_q; - int retk = k_gemm->forward(q_blob, k_affine_q, opt); + int retk = k_gemm->forward(q_blob, k_affine_q, opt_wq); if (retk != 0) return retk; @@ -451,7 +953,7 @@ int MultiHeadAttention_arm::forward(const std::vector& bottom_blobs, std::v } else { - int retk = k_gemm->forward(k_blob, k_affine, opt); + int retk = k_gemm->forward(k_blob, k_affine, opt_wq); if (retk != 0) return retk; } @@ -502,7 +1004,7 @@ int MultiHeadAttention_arm::forward(const std::vector& bottom_blobs, std::v if (q_blob_i == v_blob_i) { Mat v_affine_q; - int retk = v_gemm->forward(v_blob, v_affine_q, opt); + int retk = v_gemm->forward(v_blob, v_affine_q, opt_wq); if (retk != 0) return retk; @@ -530,7 +1032,7 @@ int MultiHeadAttention_arm::forward(const std::vector& bottom_blobs, std::v } else { - int retv = v_gemm->forward(v_blob, v_affine, opt); + int retv = v_gemm->forward(v_blob, v_affine, opt_wq); if (retv != 0) return retv; } @@ -576,7 +1078,7 @@ int MultiHeadAttention_arm::forward(const std::vector& bottom_blobs, std::v v_affine.release(); } - int reto = o_gemm->forward(qkv_cross, top_blobs[0], opt); + int reto = o_gemm->forward(qkv_cross, top_blobs[0], opt_wq); if (reto != 0) return reto; diff --git a/src/layer/arm/multiheadattention_arm.h b/src/layer/arm/multiheadattention_arm.h index c99988b03199..b70c9ceab558 100644 --- a/src/layer/arm/multiheadattention_arm.h +++ b/src/layer/arm/multiheadattention_arm.h @@ -18,6 +18,11 @@ class MultiHeadAttention_arm : public MultiHeadAttention virtual int forward(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& opt) const; +protected: +#if NCNN_WEIGHT_QUANT + int create_pipeline_wq_int8(const Option& opt); +#endif + public: Layer* q_gemm; Layer* k_gemm; diff --git a/src/layer/gemm.cpp b/src/layer/gemm.cpp index 2023bcf7011e..a43d3e500eb1 100644 --- a/src/layer/gemm.cpp +++ b/src/layer/gemm.cpp @@ -54,6 +54,144 @@ static int gemm_weight_quantize_packed_k_bytes(int constantK, int weight_bits) return (int)packed_k_bytes; } + +static inline signed char weight_block_quantize_float2int8(float v) +{ + int int32 = static_cast(round(v)); + if (int32 > 127) return 127; + if (int32 < -127) return -127; + return (signed char)int32; +} + +static void weight_block_quantize_activation_row_int8(const Mat& A, int transA, int i, signed char* outptr, float* descale_ptr, int K, int block_size, const float* input_scale_ptr) +{ + const int block_count = (K + block_size - 1) / block_size; + const size_t A_hstep = A.dims == 3 ? A.cstep : (size_t)A.w; + const float* ptrA = (const float*)A + i * A_hstep; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = block_size < K - k0 ? block_size : K - k0; + + float absmax = 0.f; + for (int kk = 0; kk < max_kk; kk++) + { + const int k = k0 + kk; + float v = transA ? ((const float*)A)[k * A_hstep + i] : ptrA[k]; + if (input_scale_ptr) + v *= input_scale_ptr[k]; + v = fabsf(v); + if (v > absmax) + absmax = v; + } + + if (absmax == 0.f) + { + descale_ptr[g] = 0.f; + for (int kk = 0; kk < max_kk; kk++) + outptr[k0 + kk] = 0; + continue; + } + + volatile double scale_fp64 = 127.0 / (double)absmax; + const float scale = (float)scale_fp64; + descale_ptr[g] = absmax / 127.f; + + for (int kk = 0; kk < max_kk; kk++) + { + const int k = k0 + kk; + float v = transA ? ((const float*)A)[k * A_hstep + i] : ptrA[k]; + if (input_scale_ptr) + { + v *= input_scale_ptr[k]; + volatile float v_ordered = v; + v = v_ordered; + } + outptr[k] = weight_block_quantize_float2int8(v * scale); + } + } +} + +static int weight_block_quantize_gemm_transB_int8(const Mat& A, int transA, const Mat& BT, const Mat& BT_scales, const Mat& input_scales, const Mat& C, Mat& top_blob, int M, int N, int K, int block_size, float alpha, float beta, int broadcast_type_C, int output_transpose, int output_m_offset, const Option& opt) +{ + const int block_count = (K + block_size - 1) / block_size; + + Mat A_int8; + A_int8.create(K, M, (size_t)1u, opt.workspace_allocator); + if (A_int8.empty()) + return -100; + + Mat A_descales; + A_descales.create(block_count, M, (size_t)4u, opt.workspace_allocator); + if (A_descales.empty()) + return -100; + + const float* input_scale_ptr = input_scales; + + #pragma omp parallel for num_threads(opt.num_threads) + for (int i = 0; i < M; i++) + { + signed char* outptr = A_int8.row(i); + float* descale_ptr = A_descales.row(i); + weight_block_quantize_activation_row_int8(A, transA, i, outptr, descale_ptr, K, block_size, input_scale_ptr); + } + + const float* ptrC = C; + + #pragma omp parallel for num_threads(opt.num_threads) + for (int mn = 0; mn < M * N; mn++) + { + const int i = mn / N; + const int j = mn % N; + + float sum = 0.f; + if (ptrC) + { + if (broadcast_type_C == 0) + sum = ptrC[0]; + if (broadcast_type_C == 1) + sum = ptrC[i]; + if (broadcast_type_C == 2) + sum = ptrC[i]; + if (broadcast_type_C == 3) + sum = ptrC[i * N + j]; + if (broadcast_type_C == 4) + sum = ptrC[j]; + + sum *= beta; + } + + const signed char* ptrA = A_int8.row(i); + const signed char* ptrB = BT.row(j); + const float* A_descale_ptr = A_descales.row(i); + const float* B_scale_ptr = BT_scales.row(j); + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = block_size < K - k0 ? block_size : K - k0; + + int sum_int32 = 0; + for (int kk = 0; kk < max_kk; kk++) + { + const int k = k0 + kk; + sum_int32 += ptrA[k] * ptrB[k]; + } + + sum += sum_int32 * A_descale_ptr[g] / B_scale_ptr[g]; + } + + sum *= alpha; + + if (output_transpose) + top_blob.row(j)[output_m_offset + i] = sum; + else + top_blob.row(output_m_offset + i)[j] = sum; + } + + return 0; +} #endif // NCNN_WEIGHT_QUANT Gemm::Gemm() @@ -102,13 +240,15 @@ int Gemm::load_param(const ParamDict& pd) if (weight_block_quantize) { #if NCNN_WEIGHT_QUANT - if (constantA != 0 || constantB != 1 || transA != 0 || transB != 1) + const int weight_bits = gemm_weight_quantize_bits(quantize_term); + + if (constantA != 0 || constantB != 1 || transB != 1 || (transA != 0 && (weight_bits != 8 || transA != 1))) { NCNN_LOGE("Gemm unsupported weight block quantize"); return -1; } - if (output_N1M != 0 || output_elempack != 0 || (output_elemtype != 0 && output_elemtype != 1) || output_transpose != 0) + if (output_N1M != 0 || output_elempack != 0 || (output_elemtype != 0 && output_elemtype != 1) || (output_transpose != 0 && (weight_bits != 8 || output_transpose != 1))) { NCNN_LOGE("Gemm unsupported weight block quantize"); return -1; @@ -344,7 +484,13 @@ int Gemm::forward_weight_block_quantize(const std::vector& bottom_blobs, st return -1; } - const int K = A.w; + if (transA && A.dims != 2) + { + NCNN_LOGE("Gemm unsupported input"); + return -1; + } + + const int K = transA ? A.h : A.w; if (K != constantK) { NCNN_LOGE("Gemm weight block quantize K mismatch"); @@ -358,7 +504,7 @@ int Gemm::forward_weight_block_quantize(const std::vector& bottom_blobs, st const bool has_input_scale = gemm_weight_quantize_has_input_scale(quantize_term); const float* input_scale_ptr = has_input_scale ? (const float*)B_data_input_scales : 0; - const int M = A.dims == 3 ? A.c : A.h; + const int M = transA ? A.w : A.dims == 3 ? A.c : A.h; const int N = constantN; Mat C; @@ -424,13 +570,21 @@ int Gemm::forward_weight_block_quantize(const std::vector& bottom_blobs, st } Mat& top_blob = top_blobs[0]; - top_blob.create(N, M, (size_t)4u, opt.blob_allocator); + if (output_transpose) + top_blob.create(M, N, (size_t)4u, opt.blob_allocator); + else + top_blob.create(N, M, (size_t)4u, opt.blob_allocator); if (top_blob.empty()) return -100; const size_t A_hstep = A.dims == 3 ? A.cstep : (size_t)A.w; const float* ptrC = C; + if (weight_bits == 8) + { + return weight_block_quantize_gemm_transB_int8(A, transA, B_data, B_data_quantize_scales, B_data_input_scales, C, top_blob, M, N, K, block_size, alpha, beta, broadcast_type_C, output_transpose, 0, opt); + } + #pragma omp parallel for num_threads(opt.num_threads) for (int i = 0; i < M; i++) { diff --git a/src/layer/loongarch/gemm_loongarch.cpp b/src/layer/loongarch/gemm_loongarch.cpp index 72832080bcf7..f0dd51074e8a 100644 --- a/src/layer/loongarch/gemm_loongarch.cpp +++ b/src/layer/loongarch/gemm_loongarch.cpp @@ -13,6 +13,10 @@ namespace ncnn { #include "gemm_int8.h" #endif +#if NCNN_WEIGHT_QUANT +#include "gemm_wq_int8.h" +#endif + #if NCNN_BF16 #include "gemm_bf16s.h" #endif @@ -7365,15 +7369,130 @@ static int gemm_AT_BT_loongarch(const Mat& AT, const Mat& BT, const Mat& C, Mat& return 0; } +#if NCNN_WEIGHT_QUANT +static int gemm_BT_loongarch_wq_int8(const Mat& A, const Mat& packed_B, const Mat& packed_B_descales, const Mat& input_scales, const Mat& C, Mat& top_blob, int broadcast_type_C, int N, int K, int block_size, int transA, int output_transpose, float alpha, float beta, int constant_TILE_M, int constant_TILE_N, int constant_TILE_K, int nT, const Option& opt) +{ + const int M = transA ? A.w : (A.dims == 3 ? A.c : A.h) * A.elempack; + const int block_count = (K + block_size - 1) / block_size; + const Mat BT = packed_B.reshape(K, N); + const Mat BT_descales = packed_B_descales.reshape(block_count, N); + int TILE_M, TILE_N, TILE_K; + get_optimal_tile_mnk_wq_int8(M, N, K, constant_TILE_M, constant_TILE_N, constant_TILE_K, TILE_M, TILE_N, TILE_K, nT); + + (void)TILE_K; + const int mr = std::min(M, TILE_M); + const int nr = std::min(N, TILE_N); + const int nn_M = (M + TILE_M - 1) / TILE_M; + const int nn_N = (N + TILE_N - 1) / TILE_N; + const float* input_scale_ptr = input_scales; + + Mat topT(nr * mr, 1, nT, (size_t)4u, opt.workspace_allocator); + if (topT.empty()) + return -100; + + if (nT > nn_M) + { + Mat AT(K, mr, nn_M, (size_t)1u, opt.workspace_allocator); + Mat AT_descales(block_count, mr, nn_M, (size_t)4u, opt.workspace_allocator); + if (AT.empty() || AT_descales.empty()) + return -100; + + #pragma omp parallel for num_threads(nT) + for (int ppi = 0; ppi < nn_M; ppi++) + { + const int i = ppi * TILE_M; + const int max_ii = std::min(M - i, TILE_M); + + Mat AT_tile = AT.channel(i / TILE_M); + Mat AT_descales_tile = AT_descales.channel(i / TILE_M); + + if (transA) + transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); + else + quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); + } + + #pragma omp parallel for num_threads(nT) + for (int ppij = 0; ppij < nn_M * nn_N; ppij++) + { + const int ppi = ppij / nn_N; + const int ppj = ppij % nn_N; + const int i = ppi * TILE_M; + const int j = ppj * TILE_N; + const int max_ii = std::min(M - i, TILE_M); + const int max_jj = std::min(N - j, TILE_N); + Mat AT_tile = AT.channel(i / TILE_M); + Mat AT_descales_tile = AT_descales.channel(i / TILE_M); + Mat BT_tile = BT.row_range(j, max_jj); + Mat BT_descales_tile = BT_descales.row_range(j, max_jj); + Mat topT_tile = topT.channel(get_omp_thread_num()); + gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); + if (output_transpose) + transpose_unpack_output_tile_wq_int8(topT_tile, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, M, alpha, beta); + else + unpack_output_tile_wq_int8(topT_tile, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, N, alpha, beta); + } + } + else + { + Mat ATX(K, mr, nT, (size_t)1u, opt.workspace_allocator); + Mat ATX_descales(block_count, mr, nT, (size_t)4u, opt.workspace_allocator); + if (ATX.empty() || ATX_descales.empty()) + return -100; + + #pragma omp parallel for num_threads(nT) + for (int ppi = 0; ppi < nn_M; ppi++) + { + const int i = ppi * TILE_M; + const int max_ii = std::min(M - i, TILE_M); + + Mat AT_tile = ATX.channel(get_omp_thread_num()); + Mat AT_descales_tile = ATX_descales.channel(get_omp_thread_num()); + Mat topT_tile = topT.channel(get_omp_thread_num()); + + if (transA) + transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); + else + quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); + + for (int j = 0; j < N; j += TILE_N) + { + const int max_jj = std::min(N - j, TILE_N); + Mat BT_tile = BT.row_range(j, max_jj); + Mat BT_descales_tile = BT_descales.row_range(j, max_jj); + gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); + if (output_transpose) transpose_unpack_output_tile_wq_int8(topT_tile, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, M, alpha, beta); + else unpack_output_tile_wq_int8(topT_tile, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, N, alpha, beta); + } + } + } + + return 0; +} +#endif // NCNN_WEIGHT_QUANT + int Gemm_loongarch::create_pipeline(const Option& opt) { +#if NCNN_WEIGHT_QUANT + if (weight_block_quantize && quantize_term / 100 == 8 && !BT_data_wq_int8.empty()) + return 0; +#endif + AT_data.release(); BT_data.release(); CT_data.release(); +#if NCNN_WEIGHT_QUANT + BT_data_wq_int8.release(); + BT_data_wq_int8_descales.release(); +#endif nT = 0; if (weight_block_quantize) { +#if NCNN_WEIGHT_QUANT + if (quantize_term / 100 == 8) + return create_pipeline_wq_int8(opt); +#endif return 0; } @@ -7525,10 +7644,129 @@ int Gemm_loongarch::create_pipeline(const Option& opt) return 0; } +int Gemm_loongarch::destroy_pipeline(const Option& /*opt*/) +{ + AT_data.release(); + BT_data.release(); + CT_data.release(); +#if NCNN_WEIGHT_QUANT + BT_data_wq_int8.release(); + BT_data_wq_int8_descales.release(); +#endif + nT = 0; + + return 0; +} + +#if NCNN_WEIGHT_QUANT +int Gemm_loongarch::create_pipeline_wq_int8(const Option& opt) +{ + if (B_data.empty() || B_data_quantize_scales.empty()) + return -100; + + const int block_size_code = quantize_term % 10; + const int block_size = block_size_code == 0 ? 32 : block_size_code == 1 ? 64 : 128; + + Mat BT_data_packed; + Mat BT_data_packed_descales; + int ret = pack_B_wq_int8(B_data, B_data_quantize_scales, BT_data_packed, BT_data_packed_descales, constantN, constantK, block_size, opt.num_threads); + if (ret != 0) + return ret; + + BT_data_wq_int8 = BT_data_packed; + BT_data_wq_int8_descales = BT_data_packed_descales; + + B_data.release(); + B_data_quantize_scales.release(); + + return 0; +} + +int Gemm_loongarch::forward_wq_int8(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& opt) const +{ + const Mat& A = bottom_blobs[0]; + if (A.elemsize != 4u || A.elempack != 1) + { + NCNN_LOGE("Gemm unsupported input"); + return -1; + } + if (transA && A.dims != 2) + { + NCNN_LOGE("Gemm unsupported input"); + return -1; + } + + const int K = transA ? A.h : A.w; + if (K != constantK) + { + NCNN_LOGE("Gemm weight block quantize K mismatch"); + return -1; + } + const int M = transA ? A.w : A.dims == 3 ? A.c : A.h; + const int N = constantN; + + Mat C; + int broadcast_type_C = -1; + if (constantC) + { + C = C_data; + broadcast_type_C = constant_broadcast_type_C; + } + else + { + if (bottom_blobs.size() == 2) + C = bottom_blobs[1]; + + if (!C.empty()) + { + bool matched = false; + if (C.dims == 1 && C.w == 1) + broadcast_type_C = 0, matched = true; + if (C.dims == 1 && C.w == M) + broadcast_type_C = 1, matched = true; + if (C.dims == 1 && C.w == N) + broadcast_type_C = 4, matched = true; + if (C.dims == 2 && C.w == 1 && C.h == M) + broadcast_type_C = 2, matched = true; + if (C.dims == 2 && C.w == N && C.h == M) + broadcast_type_C = 3, matched = true; + if (C.dims == 2 && C.w == N && C.h == 1) + broadcast_type_C = 4, matched = true; + if (!matched || C.elemsize != 4u || C.elempack != 1) + { + NCNN_LOGE("Gemm unsupported C"); + return -1; + } + } + } + if (!C.empty() && (C.elemsize != 4u || C.elempack != 1)) + { + NCNN_LOGE("Gemm unsupported C"); + return -1; + } + + Mat& top_blob = top_blobs[0]; + if (output_transpose) + top_blob.create(M, N, (size_t)4u, opt.blob_allocator); + else + top_blob.create(N, M, (size_t)4u, opt.blob_allocator); + if (top_blob.empty()) + return -100; + + const int block_size_code = quantize_term % 10; + const int block_size = block_size_code == 0 ? 32 : block_size_code == 1 ? 64 : 128; + return gemm_BT_loongarch_wq_int8(A, BT_data_wq_int8, BT_data_wq_int8_descales, B_data_input_scales, C, top_blob, broadcast_type_C, N, K, block_size, transA, output_transpose, alpha, beta, constant_TILE_M, constant_TILE_N, constant_TILE_K, opt.num_threads, opt); +} +#endif // NCNN_WEIGHT_QUANT + int Gemm_loongarch::forward(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& opt) const { if (weight_block_quantize) { +#if NCNN_WEIGHT_QUANT + if (quantize_term / 100 == 8 && !BT_data_wq_int8.empty()) + return forward_wq_int8(bottom_blobs, top_blobs, opt); +#endif return Gemm::forward(bottom_blobs, top_blobs, opt); } diff --git a/src/layer/loongarch/gemm_loongarch.h b/src/layer/loongarch/gemm_loongarch.h index 050557e5ba4d..3435e6357460 100644 --- a/src/layer/loongarch/gemm_loongarch.h +++ b/src/layer/loongarch/gemm_loongarch.h @@ -14,6 +14,7 @@ class Gemm_loongarch : public Gemm Gemm_loongarch(); virtual int create_pipeline(const Option& opt); + virtual int destroy_pipeline(const Option& opt); virtual int forward(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& opt) const; @@ -23,6 +24,11 @@ class Gemm_loongarch : public Gemm int forward_int8(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& opt) const; #endif +#if NCNN_WEIGHT_QUANT + int create_pipeline_wq_int8(const Option& opt); + int forward_wq_int8(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& opt) const; +#endif + #if NCNN_BF16 int create_pipeline_bf16s(const Option& opt); int forward_bf16s(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& opt) const; @@ -33,6 +39,11 @@ class Gemm_loongarch : public Gemm Mat AT_data; Mat BT_data; Mat CT_data; + +#if NCNN_WEIGHT_QUANT + Mat BT_data_wq_int8; + Mat BT_data_wq_int8_descales; +#endif }; // expose some gemm internal routines for convolution uses diff --git a/src/layer/loongarch/gemm_wq_int8.h b/src/layer/loongarch/gemm_wq_int8.h new file mode 100644 index 000000000000..72825207f51f --- /dev/null +++ b/src/layer/loongarch/gemm_wq_int8.h @@ -0,0 +1,6297 @@ +// Copyright 2026 Tencent +// SPDX-License-Identifier: BSD-3-Clause + +#include + +static int pack_B_wq_int8(const Mat& B, const Mat& B_scales, Mat& BT, Mat& BT_descales, int N, int K, int block_size, int num_threads) +{ + const int block_count = (K + block_size - 1) / block_size; + Mat BT_packed(N * K, (size_t)1u); + Mat BT_packed_descales(N * block_count, (size_t)4u); + if (BT_packed.empty() || BT_packed_descales.empty()) + return -100; + BT_packed.cstep = (size_t)N * K; + BT_packed_descales.cstep = (size_t)N * block_count; + + int panel_start = 0; + int panel_count = 0; +#if __loongarch_asx + const int nn8 = (N - panel_start) / 8; + const int panel_start8 = panel_start; + panel_start += nn8 * 8; + panel_count += nn8; +#endif +#if __loongarch_sx + const int nn4 = (N - panel_start) / 4; + const int panel_start4 = panel_start; + panel_start += nn4 * 4; + panel_count += nn4; +#endif + const int nn2 = (N - panel_start) / 2; + const int panel_start2 = panel_start; + panel_start += nn2 * 2; + panel_count += nn2; + const int nn1 = N - panel_start; + const int panel_start1 = panel_start; + panel_count += nn1; + + #pragma omp parallel for num_threads(num_threads) + for (int p = 0; p < panel_count; p++) + { + int q = p; + int j = 0; + int nr = 1; +#if __loongarch_asx + if (q < nn8) + { + j = panel_start8 + q * 8; + nr = 8; + } + else + { + q -= nn8; +#endif +#if __loongarch_sx + if (q < nn4) + { + j = panel_start4 + q * 4; + nr = 4; + } + else + { + q -= nn4; +#endif + if (q < nn2) + { + j = panel_start2 + q * 2; + nr = 2; + } + else + { + q -= nn2; + j = panel_start1 + q; + nr = 1; + } +#if __loongarch_sx + } +#endif +#if __loongarch_asx + } +#endif + + signed char* pp = (signed char*)BT_packed + j * K; + float* pd = (float*)BT_packed_descales + j * block_count; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + for (int jj = 0; jj < nr; jj++) + { + const signed char* pB = B.row(j + jj) + k0 + kk; + pp[0] = pB[0]; + pp[1] = pB[1]; + pp[2] = pB[2]; + pp[3] = pB[3]; + pp += 4; + } + } + if (kk + 1 < max_kk) + { + for (int jj = 0; jj < nr; jj++) + { + const signed char* pB = B.row(j + jj) + k0 + kk; + pp[0] = pB[0]; + pp[1] = pB[1]; + pp += 2; + } + kk += 2; + } + if (kk < max_kk) + { + for (int jj = 0; jj < nr; jj++) + *pp++ = B.row(j + jj)[k0 + kk]; + } + + for (int jj = 0; jj < nr; jj++) + pd[g * nr + jj] = 1.f / B_scales.row(j + jj)[g]; + } + } + + BT = BT_packed; + BT_descales = BT_packed_descales; + return 0; +} + +static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr) +{ + signed char* outptr = AT_tile; + const int out_hstep = AT_tile.w; + float* descales = AT_descales_tile; + const int descales_hstep = AT_descales_tile.w; + const int block_count = (K + block_size - 1) / block_size; + const size_t A_hstep = A.dims == 3 ? A.cstep : (size_t)A.w; + + int ii = 0; +#if __loongarch_sx + for (; ii + 7 < max_ii; ii += 8) + { + const float* p0 = (const float*)A + (i + ii) * A_hstep; + const float* p1 = p0 + A_hstep; + const float* p2 = p1 + A_hstep; + const float* p3 = p2 + A_hstep; + const float* p4 = p3 + A_hstep; + const float* p5 = p4 + A_hstep; + const float* p6 = p5 + A_hstep; + const float* p7 = p6 + A_hstep; + + signed char* pp = outptr + ii * out_hstep; + float* pd = descales + ii * descales_hstep; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + const __m128i _abs_mask = __lsx_vreplgr2vr_w(0x7fffffff); + __m128 _absmax0 = (__m128)__lsx_vldi(0); + __m128 _absmax1 = (__m128)__lsx_vldi(0); + __m128 _absmax2 = (__m128)__lsx_vldi(0); + __m128 _absmax3 = (__m128)__lsx_vldi(0); + __m128 _absmax4 = (__m128)__lsx_vldi(0); + __m128 _absmax5 = (__m128)__lsx_vldi(0); + __m128 _absmax6 = (__m128)__lsx_vldi(0); + __m128 _absmax7 = (__m128)__lsx_vldi(0); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __m128 _v0 = (__m128)__lsx_vld(p0 + k0 + kk, 0); + __m128 _v1 = (__m128)__lsx_vld(p1 + k0 + kk, 0); + __m128 _v2 = (__m128)__lsx_vld(p2 + k0 + kk, 0); + __m128 _v3 = (__m128)__lsx_vld(p3 + k0 + kk, 0); + __m128 _v4 = (__m128)__lsx_vld(p4 + k0 + kk, 0); + __m128 _v5 = (__m128)__lsx_vld(p5 + k0 + kk, 0); + __m128 _v6 = (__m128)__lsx_vld(p6 + k0 + kk, 0); + __m128 _v7 = (__m128)__lsx_vld(p7 + k0 + kk, 0); + if (input_scale_ptr) + { + const __m128 _s = (__m128)__lsx_vld(input_scale_ptr + k0 + kk, 0); + _v0 = __lsx_vfmul_s(_v0, _s); + _v1 = __lsx_vfmul_s(_v1, _s); + _v2 = __lsx_vfmul_s(_v2, _s); + _v3 = __lsx_vfmul_s(_v3, _s); + _v4 = __lsx_vfmul_s(_v4, _s); + _v5 = __lsx_vfmul_s(_v5, _s); + _v6 = __lsx_vfmul_s(_v6, _s); + _v7 = __lsx_vfmul_s(_v7, _s); + } + _absmax0 = __lsx_vfmax_s(_absmax0, (__m128)__lsx_vand_v((__m128i)_v0, _abs_mask)); + _absmax1 = __lsx_vfmax_s(_absmax1, (__m128)__lsx_vand_v((__m128i)_v1, _abs_mask)); + _absmax2 = __lsx_vfmax_s(_absmax2, (__m128)__lsx_vand_v((__m128i)_v2, _abs_mask)); + _absmax3 = __lsx_vfmax_s(_absmax3, (__m128)__lsx_vand_v((__m128i)_v3, _abs_mask)); + _absmax4 = __lsx_vfmax_s(_absmax4, (__m128)__lsx_vand_v((__m128i)_v4, _abs_mask)); + _absmax5 = __lsx_vfmax_s(_absmax5, (__m128)__lsx_vand_v((__m128i)_v5, _abs_mask)); + _absmax6 = __lsx_vfmax_s(_absmax6, (__m128)__lsx_vand_v((__m128i)_v6, _abs_mask)); + _absmax7 = __lsx_vfmax_s(_absmax7, (__m128)__lsx_vand_v((__m128i)_v7, _abs_mask)); + } + float absmax0 = __lsx_reduce_fmax_s(_absmax0); + float absmax1 = __lsx_reduce_fmax_s(_absmax1); + float absmax2 = __lsx_reduce_fmax_s(_absmax2); + float absmax3 = __lsx_reduce_fmax_s(_absmax3); + float absmax4 = __lsx_reduce_fmax_s(_absmax4); + float absmax5 = __lsx_reduce_fmax_s(_absmax5); + float absmax6 = __lsx_reduce_fmax_s(_absmax6); + float absmax7 = __lsx_reduce_fmax_s(_absmax7); + for (; kk < max_kk; kk++) + { + const int k = k0 + kk; + const float s = input_scale_ptr ? input_scale_ptr[k] : 1.f; + absmax0 = std::max(absmax0, fabsf(p0[k] * s)); + absmax1 = std::max(absmax1, fabsf(p1[k] * s)); + absmax2 = std::max(absmax2, fabsf(p2[k] * s)); + absmax3 = std::max(absmax3, fabsf(p3[k] * s)); + absmax4 = std::max(absmax4, fabsf(p4[k] * s)); + absmax5 = std::max(absmax5, fabsf(p5[k] * s)); + absmax6 = std::max(absmax6, fabsf(p6[k] * s)); + absmax7 = std::max(absmax7, fabsf(p7[k] * s)); + } + + pd[0] = absmax0 / 127.f; + pd[1] = absmax1 / 127.f; + pd[2] = absmax2 / 127.f; + pd[3] = absmax3 / 127.f; + pd[4] = absmax4 / 127.f; + pd[5] = absmax5 / 127.f; + pd[6] = absmax6 / 127.f; + pd[7] = absmax7 / 127.f; + pd += 8; + + volatile double scale0_fp64 = absmax0 == 0.f ? 0.0 : 127.0 / (double)absmax0; + volatile double scale1_fp64 = absmax1 == 0.f ? 0.0 : 127.0 / (double)absmax1; + volatile double scale2_fp64 = absmax2 == 0.f ? 0.0 : 127.0 / (double)absmax2; + volatile double scale3_fp64 = absmax3 == 0.f ? 0.0 : 127.0 / (double)absmax3; + volatile double scale4_fp64 = absmax4 == 0.f ? 0.0 : 127.0 / (double)absmax4; + volatile double scale5_fp64 = absmax5 == 0.f ? 0.0 : 127.0 / (double)absmax5; + volatile double scale6_fp64 = absmax6 == 0.f ? 0.0 : 127.0 / (double)absmax6; + volatile double scale7_fp64 = absmax7 == 0.f ? 0.0 : 127.0 / (double)absmax7; + const float scale0 = (float)scale0_fp64; + const float scale1 = (float)scale1_fp64; + const float scale2 = (float)scale2_fp64; + const float scale3 = (float)scale3_fp64; + const float scale4 = (float)scale4_fp64; + const float scale5 = (float)scale5_fp64; + const float scale6 = (float)scale6_fp64; + const float scale7 = (float)scale7_fp64; + const __m128 _scale0 = __lsx_vreplfr2vr_s(scale0); + const __m128 _scale1 = __lsx_vreplfr2vr_s(scale1); + const __m128 _scale2 = __lsx_vreplfr2vr_s(scale2); + const __m128 _scale3 = __lsx_vreplfr2vr_s(scale3); + const __m128 _scale4 = __lsx_vreplfr2vr_s(scale4); + const __m128 _scale5 = __lsx_vreplfr2vr_s(scale5); + const __m128 _scale6 = __lsx_vreplfr2vr_s(scale6); + const __m128 _scale7 = __lsx_vreplfr2vr_s(scale7); + kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + const int k = k0 + kk; + __m128 _v0 = (__m128)__lsx_vld(p0 + k, 0); + __m128 _v1 = (__m128)__lsx_vld(p1 + k, 0); + __m128 _v2 = (__m128)__lsx_vld(p2 + k, 0); + __m128 _v3 = (__m128)__lsx_vld(p3 + k, 0); + __m128 _v4 = (__m128)__lsx_vld(p4 + k, 0); + __m128 _v5 = (__m128)__lsx_vld(p5 + k, 0); + __m128 _v6 = (__m128)__lsx_vld(p6 + k, 0); + __m128 _v7 = (__m128)__lsx_vld(p7 + k, 0); + if (input_scale_ptr) + { + const __m128 _s = (__m128)__lsx_vld(input_scale_ptr + k, 0); + _v0 = __lsx_vfmul_s(_v0, _s); + _v1 = __lsx_vfmul_s(_v1, _s); + _v2 = __lsx_vfmul_s(_v2, _s); + _v3 = __lsx_vfmul_s(_v3, _s); + _v4 = __lsx_vfmul_s(_v4, _s); + _v5 = __lsx_vfmul_s(_v5, _s); + _v6 = __lsx_vfmul_s(_v6, _s); + _v7 = __lsx_vfmul_s(_v7, _s); + } + *((int64_t*)pp) = float2int8(__lsx_vfmul_s(_v0, _scale0), __lsx_vfmul_s(_v1, _scale1)); + *((int64_t*)(pp + 8)) = float2int8(__lsx_vfmul_s(_v2, _scale2), __lsx_vfmul_s(_v3, _scale3)); + *((int64_t*)(pp + 16)) = float2int8(__lsx_vfmul_s(_v4, _scale4), __lsx_vfmul_s(_v5, _scale5)); + *((int64_t*)(pp + 24)) = float2int8(__lsx_vfmul_s(_v6, _scale6), __lsx_vfmul_s(_v7, _scale7)); + pp += 32; + } + for (; kk < max_kk; kk++) + { + const int k = k0 + kk; + const float s = input_scale_ptr ? input_scale_ptr[k] : 1.f; + pp[0] = float2int8(p0[k] * s * scale0); + pp[1] = float2int8(p1[k] * s * scale1); + pp[2] = float2int8(p2[k] * s * scale2); + pp[3] = float2int8(p3[k] * s * scale3); + pp[4] = float2int8(p4[k] * s * scale4); + pp[5] = float2int8(p5[k] * s * scale5); + pp[6] = float2int8(p6[k] * s * scale6); + pp[7] = float2int8(p7[k] * s * scale7); + pp += 8; + } + } + } +#endif // __loongarch_sx +#if __loongarch_sx + for (; ii + 3 < max_ii; ii += 4) + { + const float* p0 = (const float*)A + (i + ii) * A_hstep; + const float* p1 = p0 + A_hstep; + const float* p2 = p1 + A_hstep; + const float* p3 = p2 + A_hstep; + signed char* pp = outptr + ii * out_hstep; + float* pd = descales + ii * descales_hstep; + const __m128i _abs_mask = __lsx_vreplgr2vr_w(0x7fffffff); + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + __m128 _absmax0 = (__m128)__lsx_vldi(0); + __m128 _absmax1 = (__m128)__lsx_vldi(0); + __m128 _absmax2 = (__m128)__lsx_vldi(0); + __m128 _absmax3 = (__m128)__lsx_vldi(0); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __m128 _v0 = (__m128)__lsx_vld(p0 + k0 + kk, 0); + __m128 _v1 = (__m128)__lsx_vld(p1 + k0 + kk, 0); + __m128 _v2 = (__m128)__lsx_vld(p2 + k0 + kk, 0); + __m128 _v3 = (__m128)__lsx_vld(p3 + k0 + kk, 0); + if (input_scale_ptr) + { + const __m128 _s = (__m128)__lsx_vld(input_scale_ptr + k0 + kk, 0); + _v0 = __lsx_vfmul_s(_v0, _s); + _v1 = __lsx_vfmul_s(_v1, _s); + _v2 = __lsx_vfmul_s(_v2, _s); + _v3 = __lsx_vfmul_s(_v3, _s); + } + _absmax0 = __lsx_vfmax_s(_absmax0, (__m128)__lsx_vand_v((__m128i)_v0, _abs_mask)); + _absmax1 = __lsx_vfmax_s(_absmax1, (__m128)__lsx_vand_v((__m128i)_v1, _abs_mask)); + _absmax2 = __lsx_vfmax_s(_absmax2, (__m128)__lsx_vand_v((__m128i)_v2, _abs_mask)); + _absmax3 = __lsx_vfmax_s(_absmax3, (__m128)__lsx_vand_v((__m128i)_v3, _abs_mask)); + } + float absmax0 = __lsx_reduce_fmax_s(_absmax0); + float absmax1 = __lsx_reduce_fmax_s(_absmax1); + float absmax2 = __lsx_reduce_fmax_s(_absmax2); + float absmax3 = __lsx_reduce_fmax_s(_absmax3); + for (; kk < max_kk; kk++) + { + const int k = k0 + kk; + const float s = input_scale_ptr ? input_scale_ptr[k] : 1.f; + absmax0 = std::max(absmax0, fabsf(p0[k] * s)); + absmax1 = std::max(absmax1, fabsf(p1[k] * s)); + absmax2 = std::max(absmax2, fabsf(p2[k] * s)); + absmax3 = std::max(absmax3, fabsf(p3[k] * s)); + } + pd[0] = absmax0 / 127.f; + pd[1] = absmax1 / 127.f; + pd[2] = absmax2 / 127.f; + pd[3] = absmax3 / 127.f; + pd += 4; + volatile double scale0_fp64 = absmax0 == 0.f ? 0.0 : 127.0 / (double)absmax0; + volatile double scale1_fp64 = absmax1 == 0.f ? 0.0 : 127.0 / (double)absmax1; + volatile double scale2_fp64 = absmax2 == 0.f ? 0.0 : 127.0 / (double)absmax2; + volatile double scale3_fp64 = absmax3 == 0.f ? 0.0 : 127.0 / (double)absmax3; + const float scale0 = (float)scale0_fp64; + const float scale1 = (float)scale1_fp64; + const float scale2 = (float)scale2_fp64; + const float scale3 = (float)scale3_fp64; + const __m128 _scale0 = __lsx_vreplfr2vr_s(scale0); + const __m128 _scale1 = __lsx_vreplfr2vr_s(scale1); + const __m128 _scale2 = __lsx_vreplfr2vr_s(scale2); + const __m128 _scale3 = __lsx_vreplfr2vr_s(scale3); + kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __m128 _v0 = (__m128)__lsx_vld(p0 + k0 + kk, 0); + __m128 _v1 = (__m128)__lsx_vld(p1 + k0 + kk, 0); + __m128 _v2 = (__m128)__lsx_vld(p2 + k0 + kk, 0); + __m128 _v3 = (__m128)__lsx_vld(p3 + k0 + kk, 0); + if (input_scale_ptr) + { + const __m128 _s = (__m128)__lsx_vld(input_scale_ptr + k0 + kk, 0); + _v0 = __lsx_vfmul_s(_v0, _s); + _v1 = __lsx_vfmul_s(_v1, _s); + _v2 = __lsx_vfmul_s(_v2, _s); + _v3 = __lsx_vfmul_s(_v3, _s); + } + *((int64_t*)pp) = float2int8(__lsx_vfmul_s(_v0, _scale0), __lsx_vfmul_s(_v1, _scale1)); + *((int64_t*)(pp + 8)) = float2int8(__lsx_vfmul_s(_v2, _scale2), __lsx_vfmul_s(_v3, _scale3)); + pp += 16; + } + if (kk + 1 < max_kk) + { + const int k = k0 + kk; + const float s0 = input_scale_ptr ? input_scale_ptr[k] : 1.f; + const float s1 = input_scale_ptr ? input_scale_ptr[k + 1] : 1.f; + pp[0] = float2int8(p0[k] * s0 * scale0); + pp[1] = float2int8(p0[k + 1] * s1 * scale0); + pp[2] = float2int8(p1[k] * s0 * scale1); + pp[3] = float2int8(p1[k + 1] * s1 * scale1); + pp[4] = float2int8(p2[k] * s0 * scale2); + pp[5] = float2int8(p2[k + 1] * s1 * scale2); + pp[6] = float2int8(p3[k] * s0 * scale3); + pp[7] = float2int8(p3[k + 1] * s1 * scale3); + pp += 8; + kk += 2; + } + if (kk < max_kk) + { + const int k = k0 + kk; + const float s = input_scale_ptr ? input_scale_ptr[k] : 1.f; + pp[0] = float2int8(p0[k] * s * scale0); + pp[1] = float2int8(p1[k] * s * scale1); + pp[2] = float2int8(p2[k] * s * scale2); + pp[3] = float2int8(p3[k] * s * scale3); + pp += 4; + } + } + } + for (; ii + 1 < max_ii; ii += 2) + { + const float* p0 = (const float*)A + (i + ii) * A_hstep; + const float* p1 = p0 + A_hstep; + signed char* pp = outptr + ii * out_hstep; + float* pd = descales + ii * descales_hstep; + const __m128i _abs_mask = __lsx_vreplgr2vr_w(0x7fffffff); + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + __m128 _absmax0 = (__m128)__lsx_vldi(0); + __m128 _absmax1 = (__m128)__lsx_vldi(0); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __m128 _v0 = (__m128)__lsx_vld(p0 + k0 + kk, 0); + __m128 _v1 = (__m128)__lsx_vld(p1 + k0 + kk, 0); + if (input_scale_ptr) + { + const __m128 _s = (__m128)__lsx_vld(input_scale_ptr + k0 + kk, 0); + _v0 = __lsx_vfmul_s(_v0, _s); + _v1 = __lsx_vfmul_s(_v1, _s); + } + _absmax0 = __lsx_vfmax_s(_absmax0, (__m128)__lsx_vand_v((__m128i)_v0, _abs_mask)); + _absmax1 = __lsx_vfmax_s(_absmax1, (__m128)__lsx_vand_v((__m128i)_v1, _abs_mask)); + } + float absmax0 = __lsx_reduce_fmax_s(_absmax0); + float absmax1 = __lsx_reduce_fmax_s(_absmax1); + for (; kk < max_kk; kk++) + { + const int k = k0 + kk; + const float s = input_scale_ptr ? input_scale_ptr[k] : 1.f; + absmax0 = std::max(absmax0, fabsf(p0[k] * s)); + absmax1 = std::max(absmax1, fabsf(p1[k] * s)); + } + pd[0] = absmax0 / 127.f; + pd[1] = absmax1 / 127.f; + pd += 2; + volatile double scale0_fp64 = absmax0 == 0.f ? 0.0 : 127.0 / (double)absmax0; + volatile double scale1_fp64 = absmax1 == 0.f ? 0.0 : 127.0 / (double)absmax1; + const float scale0 = (float)scale0_fp64; + const float scale1 = (float)scale1_fp64; + const __m128 _scale0 = __lsx_vreplfr2vr_s(scale0); + const __m128 _scale1 = __lsx_vreplfr2vr_s(scale1); + kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __m128 _v0 = (__m128)__lsx_vld(p0 + k0 + kk, 0); + __m128 _v1 = (__m128)__lsx_vld(p1 + k0 + kk, 0); + if (input_scale_ptr) + { + const __m128 _s = (__m128)__lsx_vld(input_scale_ptr + k0 + kk, 0); + _v0 = __lsx_vfmul_s(_v0, _s); + _v1 = __lsx_vfmul_s(_v1, _s); + } + *((int64_t*)pp) = float2int8(__lsx_vfmul_s(_v0, _scale0), __lsx_vfmul_s(_v1, _scale1)); + pp += 8; + } + if (kk + 1 < max_kk) + { + const int k = k0 + kk; + const float s0 = input_scale_ptr ? input_scale_ptr[k] : 1.f; + const float s1 = input_scale_ptr ? input_scale_ptr[k + 1] : 1.f; + pp[0] = float2int8(p0[k] * s0 * scale0); + pp[1] = float2int8(p0[k + 1] * s1 * scale0); + pp[2] = float2int8(p1[k] * s0 * scale1); + pp[3] = float2int8(p1[k + 1] * s1 * scale1); + pp += 4; + kk += 2; + } + if (kk < max_kk) + { + const int k = k0 + kk; + const float s = input_scale_ptr ? input_scale_ptr[k] : 1.f; + pp[0] = float2int8(p0[k] * s * scale0); + pp[1] = float2int8(p1[k] * s * scale1); + pp += 2; + } + } + } +#endif // __loongarch_sx + for (; ii < max_ii; ii++) + { + const float* ptrA = (const float*)A + (i + ii) * A_hstep; + signed char* outptr0 = outptr + ii * out_hstep; + float* descale_ptr = descales + ii * descales_hstep; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + float absmax = 0.f; + int kk = 0; +#if __loongarch_sx + const __m128i _abs_mask = __lsx_vreplgr2vr_w(0x7fffffff); +#if __loongarch_asx + const __m256i _abs_mask256 = __lasx_xvreplgr2vr_w(0x7fffffff); + __m256 _absmax256 = (__m256)__lasx_xvldi(0); + for (; kk + 7 < max_kk; kk += 8) + { + __m256 _v = (__m256)__lasx_xvld(ptrA + k0 + kk, 0); + if (input_scale_ptr) + _v = __lasx_xvfmul_s(_v, (__m256)__lasx_xvld(input_scale_ptr + k0 + kk, 0)); + _v = (__m256)__lasx_xvand_v((__m256i)_v, _abs_mask256); + _absmax256 = __lasx_xvfmax_s(_absmax256, _v); + } + absmax = __lasx_reduce_fmax_s(_absmax256); +#endif + __m128 _absmax128 = __lsx_vreplfr2vr_s(absmax); + for (; kk + 3 < max_kk; kk += 4) + { + __m128 _v = (__m128)__lsx_vld(ptrA + k0 + kk, 0); + if (input_scale_ptr) + _v = __lsx_vfmul_s(_v, (__m128)__lsx_vld(input_scale_ptr + k0 + kk, 0)); + _v = (__m128)__lsx_vand_v((__m128i)_v, _abs_mask); + _absmax128 = __lsx_vfmax_s(_absmax128, _v); + } + absmax = __lsx_reduce_fmax_s(_absmax128); +#endif + for (; kk < max_kk; kk++) + { + const int k = k0 + kk; + float v = ptrA[k]; + if (input_scale_ptr) + v *= input_scale_ptr[k]; + absmax = std::max(absmax, fabsf(v)); + } + + if (absmax == 0.f) + { + descale_ptr[g] = 0.f; + for (int k = 0; k < max_kk; k++) + outptr0[k0 + k] = 0; + continue; + } + + volatile double scale_fp64 = 127.0 / (double)absmax; + const float scale = (float)scale_fp64; + descale_ptr[g] = absmax / 127.f; + kk = 0; +#if __loongarch_sx +#if __loongarch_asx + const __m256 _scale256 = (__m256)__lasx_xvreplfr2vr_s(scale); + for (; kk + 7 < max_kk; kk += 8) + { + __m256 _v = (__m256)__lasx_xvld(ptrA + k0 + kk, 0); + if (input_scale_ptr) + _v = __lasx_xvfmul_s(_v, (__m256)__lasx_xvld(input_scale_ptr + k0 + kk, 0)); + _v = __lasx_xvfmul_s(_v, _scale256); + __lsx_vstelm_d(__lasx_extract_128_lo(float2int8(_v)), outptr0 + k0 + kk, 0, 0); + } +#endif + const __m128 _scale128 = __lsx_vreplfr2vr_s(scale); + for (; kk + 3 < max_kk; kk += 4) + { + __m128 _v = (__m128)__lsx_vld(ptrA + k0 + kk, 0); + if (input_scale_ptr) + _v = __lsx_vfmul_s(_v, (__m128)__lsx_vld(input_scale_ptr + k0 + kk, 0)); + _v = __lsx_vfmul_s(_v, _scale128); + __lsx_vstelm_w(float2int8(_v), outptr0 + k0 + kk, 0, 0); + } +#endif + for (; kk < max_kk; kk++) + { + const int k = k0 + kk; + float v = ptrA[k]; + if (input_scale_ptr) + { + v *= input_scale_ptr[k]; + // preserve multiplication order for consistent rounding + asm volatile("" : "+f"(v)); + } + outptr0[k] = float2int8(v * scale); + } + } + } +} + +static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr) +{ + signed char* outptr = AT_tile; + const int out_hstep = AT_tile.w; + float* descales = AT_descales_tile; + const int descales_hstep = AT_descales_tile.w; + const int block_count = (K + block_size - 1) / block_size; + const size_t A_hstep = A.dims == 3 ? A.cstep : (size_t)A.w; + + int ii = 0; +#if __loongarch_sx + for (; ii + 7 < max_ii; ii += 8) + { + const float* ptrA = (const float*)A + i + ii; + signed char* pp = outptr + ii * out_hstep; + float* pd = descales + ii * descales_hstep; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + const __m128i _abs_mask = __lsx_vreplgr2vr_w(0x7fffffff); + __m128 _absmax0 = (__m128)__lsx_vldi(0); + __m128 _absmax1 = (__m128)__lsx_vldi(0); + for (int kk = 0; kk < max_kk; kk++) + { + const int k = k0 + kk; + const float* p = ptrA + (size_t)k * A_hstep; + __m128 _v0 = (__m128)__lsx_vld(p, 0); + __m128 _v1 = (__m128)__lsx_vld(p + 4, 0); + if (input_scale_ptr) + { + const __m128 _s = __lsx_vreplfr2vr_s(input_scale_ptr[k]); + _v0 = __lsx_vfmul_s(_v0, _s); + _v1 = __lsx_vfmul_s(_v1, _s); + } + _absmax0 = __lsx_vfmax_s(_absmax0, (__m128)__lsx_vand_v((__m128i)_v0, _abs_mask)); + _absmax1 = __lsx_vfmax_s(_absmax1, (__m128)__lsx_vand_v((__m128i)_v1, _abs_mask)); + } + + float absmax[8]; + __lsx_vst(_absmax0, absmax, 0); + __lsx_vst(_absmax1, absmax + 4, 0); + const float absmax0 = absmax[0]; + const float absmax1 = absmax[1]; + const float absmax2 = absmax[2]; + const float absmax3 = absmax[3]; + const float absmax4 = absmax[4]; + const float absmax5 = absmax[5]; + const float absmax6 = absmax[6]; + const float absmax7 = absmax[7]; + + pd[0] = absmax0 / 127.f; + pd[1] = absmax1 / 127.f; + pd[2] = absmax2 / 127.f; + pd[3] = absmax3 / 127.f; + pd[4] = absmax4 / 127.f; + pd[5] = absmax5 / 127.f; + pd[6] = absmax6 / 127.f; + pd[7] = absmax7 / 127.f; + pd += 8; + + volatile double scale0_fp64 = absmax0 == 0.f ? 0.0 : 127.0 / (double)absmax0; + volatile double scale1_fp64 = absmax1 == 0.f ? 0.0 : 127.0 / (double)absmax1; + volatile double scale2_fp64 = absmax2 == 0.f ? 0.0 : 127.0 / (double)absmax2; + volatile double scale3_fp64 = absmax3 == 0.f ? 0.0 : 127.0 / (double)absmax3; + volatile double scale4_fp64 = absmax4 == 0.f ? 0.0 : 127.0 / (double)absmax4; + volatile double scale5_fp64 = absmax5 == 0.f ? 0.0 : 127.0 / (double)absmax5; + volatile double scale6_fp64 = absmax6 == 0.f ? 0.0 : 127.0 / (double)absmax6; + volatile double scale7_fp64 = absmax7 == 0.f ? 0.0 : 127.0 / (double)absmax7; + const float scale0 = (float)scale0_fp64; + const float scale1 = (float)scale1_fp64; + const float scale2 = (float)scale2_fp64; + const float scale3 = (float)scale3_fp64; + const float scale4 = (float)scale4_fp64; + const float scale5 = (float)scale5_fp64; + const float scale6 = (float)scale6_fp64; + const float scale7 = (float)scale7_fp64; + const __m128 _scale0 = __lsx_vreplfr2vr_s(scale0); + const __m128 _scale1 = __lsx_vreplfr2vr_s(scale1); + const __m128 _scale2 = __lsx_vreplfr2vr_s(scale2); + const __m128 _scale3 = __lsx_vreplfr2vr_s(scale3); + const __m128 _scale4 = __lsx_vreplfr2vr_s(scale4); + const __m128 _scale5 = __lsx_vreplfr2vr_s(scale5); + const __m128 _scale6 = __lsx_vreplfr2vr_s(scale6); + const __m128 _scale7 = __lsx_vreplfr2vr_s(scale7); + const float scales0[4] = {scale0, scale1, scale2, scale3}; + const float scales1[4] = {scale4, scale5, scale6, scale7}; + const __m128 _scales0 = (__m128)__lsx_vld(scales0, 0); + const __m128 _scales1 = (__m128)__lsx_vld(scales1, 0); + + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + const int k = k0 + kk; + const float* p0 = ptrA + (size_t)k * A_hstep; + const float* p1 = p0 + A_hstep; + const float* p2 = p1 + A_hstep; + const float* p3 = p2 + A_hstep; + __m128 _p0 = (__m128)__lsx_vld(p0, 0); + __m128 _p1 = (__m128)__lsx_vld(p1, 0); + __m128 _p2 = (__m128)__lsx_vld(p2, 0); + __m128 _p3 = (__m128)__lsx_vld(p3, 0); + __m128 _p4 = (__m128)__lsx_vld(p0 + 4, 0); + __m128 _p5 = (__m128)__lsx_vld(p1 + 4, 0); + __m128 _p6 = (__m128)__lsx_vld(p2 + 4, 0); + __m128 _p7 = (__m128)__lsx_vld(p3 + 4, 0); + if (input_scale_ptr) + { + _p0 = __lsx_vfmul_s(_p0, __lsx_vreplfr2vr_s(input_scale_ptr[k])); + _p1 = __lsx_vfmul_s(_p1, __lsx_vreplfr2vr_s(input_scale_ptr[k + 1])); + _p2 = __lsx_vfmul_s(_p2, __lsx_vreplfr2vr_s(input_scale_ptr[k + 2])); + _p3 = __lsx_vfmul_s(_p3, __lsx_vreplfr2vr_s(input_scale_ptr[k + 3])); + _p4 = __lsx_vfmul_s(_p4, __lsx_vreplfr2vr_s(input_scale_ptr[k])); + _p5 = __lsx_vfmul_s(_p5, __lsx_vreplfr2vr_s(input_scale_ptr[k + 1])); + _p6 = __lsx_vfmul_s(_p6, __lsx_vreplfr2vr_s(input_scale_ptr[k + 2])); + _p7 = __lsx_vfmul_s(_p7, __lsx_vreplfr2vr_s(input_scale_ptr[k + 3])); + } + transpose4x4_ps(_p0, _p1, _p2, _p3); + transpose4x4_ps(_p4, _p5, _p6, _p7); + *((int64_t*)pp) = float2int8(__lsx_vfmul_s(_p0, _scale0), __lsx_vfmul_s(_p1, _scale1)); + *((int64_t*)(pp + 8)) = float2int8(__lsx_vfmul_s(_p2, _scale2), __lsx_vfmul_s(_p3, _scale3)); + *((int64_t*)(pp + 16)) = float2int8(__lsx_vfmul_s(_p4, _scale4), __lsx_vfmul_s(_p5, _scale5)); + *((int64_t*)(pp + 24)) = float2int8(__lsx_vfmul_s(_p6, _scale6), __lsx_vfmul_s(_p7, _scale7)); + pp += 32; + } + for (; kk < max_kk; kk++) + { + const int k = k0 + kk; + const float* p = ptrA + (size_t)k * A_hstep; + __m128 _p0 = (__m128)__lsx_vld(p, 0); + __m128 _p1 = (__m128)__lsx_vld(p + 4, 0); + if (input_scale_ptr) + { + const __m128 _s = __lsx_vreplfr2vr_s(input_scale_ptr[k]); + _p0 = __lsx_vfmul_s(_p0, _s); + _p1 = __lsx_vfmul_s(_p1, _s); + } + _p0 = __lsx_vfmul_s(_p0, _scales0); + _p1 = __lsx_vfmul_s(_p1, _scales1); + *((int64_t*)pp) = float2int8(_p0, _p1); + pp += 8; + } + } + } + for (; ii + 3 < max_ii; ii += 4) + { + const float* ptrA = (const float*)A + i + ii; + signed char* pp = outptr + ii * out_hstep; + float* pd = descales + ii * descales_hstep; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + + const __m128i _abs_mask = __lsx_vreplgr2vr_w(0x7fffffff); + __m128 _absmax = (__m128)__lsx_vldi(0); + for (int kk = 0; kk < max_kk; kk++) + { + const int k = k0 + kk; + __m128 _v = (__m128)__lsx_vld(ptrA + (size_t)k * A_hstep, 0); + if (input_scale_ptr) + _v = __lsx_vfmul_s(_v, __lsx_vreplfr2vr_s(input_scale_ptr[k])); + _v = (__m128)__lsx_vand_v((__m128i)_v, _abs_mask); + _absmax = __lsx_vfmax_s(_absmax, _v); + } + + float absmax[4]; + __lsx_vst(_absmax, absmax, 0); + pd[0] = absmax[0] / 127.f; + pd[1] = absmax[1] / 127.f; + pd[2] = absmax[2] / 127.f; + pd[3] = absmax[3] / 127.f; + pd += 4; + + volatile double scale0_fp64 = absmax[0] == 0.f ? 0.0 : 127.0 / (double)absmax[0]; + volatile double scale1_fp64 = absmax[1] == 0.f ? 0.0 : 127.0 / (double)absmax[1]; + volatile double scale2_fp64 = absmax[2] == 0.f ? 0.0 : 127.0 / (double)absmax[2]; + volatile double scale3_fp64 = absmax[3] == 0.f ? 0.0 : 127.0 / (double)absmax[3]; + const float scales[4] = { + (float)scale0_fp64, + (float)scale1_fp64, + (float)scale2_fp64, + (float)scale3_fp64 + }; + const __m128 _scale = (__m128)__lsx_vld(scales, 0); + + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + const int k = k0 + kk; + const float* p0 = ptrA + (size_t)k * A_hstep; + const float* p1 = p0 + A_hstep; + const float* p2 = p1 + A_hstep; + const float* p3 = p2 + A_hstep; + __m128 _v0 = (__m128)__lsx_vld(p0, 0); + __m128 _v1 = (__m128)__lsx_vld(p1, 0); + __m128 _v2 = (__m128)__lsx_vld(p2, 0); + __m128 _v3 = (__m128)__lsx_vld(p3, 0); + if (input_scale_ptr) + { + _v0 = __lsx_vfmul_s(_v0, __lsx_vreplfr2vr_s(input_scale_ptr[k])); + _v1 = __lsx_vfmul_s(_v1, __lsx_vreplfr2vr_s(input_scale_ptr[k + 1])); + _v2 = __lsx_vfmul_s(_v2, __lsx_vreplfr2vr_s(input_scale_ptr[k + 2])); + _v3 = __lsx_vfmul_s(_v3, __lsx_vreplfr2vr_s(input_scale_ptr[k + 3])); + } + transpose4x4_ps(_v0, _v1, _v2, _v3); + *((int64_t*)pp) = float2int8(__lsx_vfmul_s(_v0, __lsx_vreplfr2vr_s(scales[0])), __lsx_vfmul_s(_v1, __lsx_vreplfr2vr_s(scales[1]))); + *((int64_t*)(pp + 8)) = float2int8(__lsx_vfmul_s(_v2, __lsx_vreplfr2vr_s(scales[2])), __lsx_vfmul_s(_v3, __lsx_vreplfr2vr_s(scales[3]))); + pp += 16; + } + if (kk + 1 < max_kk) + { + const int k = k0 + kk; + __m128 _v0 = (__m128)__lsx_vld(ptrA + (size_t)k * A_hstep, 0); + __m128 _v1 = (__m128)__lsx_vld(ptrA + (size_t)(k + 1) * A_hstep, 0); + if (input_scale_ptr) + { + _v0 = __lsx_vfmul_s(_v0, __lsx_vreplfr2vr_s(input_scale_ptr[k])); + _v1 = __lsx_vfmul_s(_v1, __lsx_vreplfr2vr_s(input_scale_ptr[k + 1])); + } + const int q0 = __lsx_vpickve2gr_w(float2int8(__lsx_vfmul_s(_v0, _scale)), 0); + const int q1 = __lsx_vpickve2gr_w(float2int8(__lsx_vfmul_s(_v1, _scale)), 0); + pp[0] = (signed char)q0; + pp[1] = (signed char)q1; + pp[2] = (signed char)(q0 >> 8); + pp[3] = (signed char)(q1 >> 8); + pp[4] = (signed char)(q0 >> 16); + pp[5] = (signed char)(q1 >> 16); + pp[6] = (signed char)(q0 >> 24); + pp[7] = (signed char)(q1 >> 24); + pp += 8; + kk += 2; + } + if (kk < max_kk) + { + const int k = k0 + kk; + __m128 _v = (__m128)__lsx_vld(ptrA + (size_t)k * A_hstep, 0); + if (input_scale_ptr) + _v = __lsx_vfmul_s(_v, __lsx_vreplfr2vr_s(input_scale_ptr[k])); + const int q = __lsx_vpickve2gr_w(float2int8(__lsx_vfmul_s(_v, _scale)), 0); + pp[0] = (signed char)q; + pp[1] = (signed char)(q >> 8); + pp[2] = (signed char)(q >> 16); + pp[3] = (signed char)(q >> 24); + pp += 4; + } + } + } + for (; ii + 1 < max_ii; ii += 2) + { + const float* ptrA = (const float*)A + i + ii; + signed char* pp = outptr + ii * out_hstep; + float* pd = descales + ii * descales_hstep; + const __m128i _abs_mask = __lsx_vreplgr2vr_w(0x7fffffff); + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + __m128 _absmax = (__m128)__lsx_vldi(0); + for (int kk = 0; kk < max_kk; kk++) + { + const int k = k0 + kk; + __m128 _v = (__m128)__lsx_vldrepl_d(ptrA + (size_t)k * A_hstep, 0); + if (input_scale_ptr) + _v = __lsx_vfmul_s(_v, __lsx_vreplfr2vr_s(input_scale_ptr[k])); + _absmax = __lsx_vfmax_s(_absmax, (__m128)__lsx_vand_v((__m128i)_v, _abs_mask)); + } + float absmax[2]; + __lsx_vstelm_d((__m128i)_absmax, absmax, 0, 0); + pd[0] = absmax[0] / 127.f; + pd[1] = absmax[1] / 127.f; + pd += 2; + volatile double scale0_fp64 = absmax[0] == 0.f ? 0.0 : 127.0 / (double)absmax[0]; + volatile double scale1_fp64 = absmax[1] == 0.f ? 0.0 : 127.0 / (double)absmax[1]; + const float scale0 = (float)scale0_fp64; + const float scale1 = (float)scale1_fp64; + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + const int k = k0 + kk; + const float* p0 = ptrA + (size_t)k * A_hstep; + const float* p1 = p0 + A_hstep; + const float* p2 = p1 + A_hstep; + const float* p3 = p2 + A_hstep; + const float s0 = input_scale_ptr ? input_scale_ptr[k] : 1.f; + const float s1 = input_scale_ptr ? input_scale_ptr[k + 1] : 1.f; + const float s2 = input_scale_ptr ? input_scale_ptr[k + 2] : 1.f; + const float s3 = input_scale_ptr ? input_scale_ptr[k + 3] : 1.f; + pp[0] = float2int8(p0[0] * s0 * scale0); + pp[1] = float2int8(p1[0] * s1 * scale0); + pp[2] = float2int8(p2[0] * s2 * scale0); + pp[3] = float2int8(p3[0] * s3 * scale0); + pp[4] = float2int8(p0[1] * s0 * scale1); + pp[5] = float2int8(p1[1] * s1 * scale1); + pp[6] = float2int8(p2[1] * s2 * scale1); + pp[7] = float2int8(p3[1] * s3 * scale1); + pp += 8; + } + if (kk + 1 < max_kk) + { + const int k = k0 + kk; + const float* p0 = ptrA + (size_t)k * A_hstep; + const float* p1 = p0 + A_hstep; + const float s0 = input_scale_ptr ? input_scale_ptr[k] : 1.f; + const float s1 = input_scale_ptr ? input_scale_ptr[k + 1] : 1.f; + pp[0] = float2int8(p0[0] * s0 * scale0); + pp[1] = float2int8(p1[0] * s1 * scale0); + pp[2] = float2int8(p0[1] * s0 * scale1); + pp[3] = float2int8(p1[1] * s1 * scale1); + pp += 4; + kk += 2; + } + if (kk < max_kk) + { + const int k = k0 + kk; + const float* p = ptrA + (size_t)k * A_hstep; + const float s = input_scale_ptr ? input_scale_ptr[k] : 1.f; + pp[0] = float2int8(p[0] * s * scale0); + pp[1] = float2int8(p[1] * s * scale1); + pp += 2; + } + } + } +#endif // __loongarch_sx + for (; ii < max_ii; ii++) + { + const float* ptrA = (const float*)A + i + ii; + signed char* outptr0 = outptr + ii * out_hstep; + float* descale_ptr = descales + ii * descales_hstep; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + + float absmax = 0.f; + for (int kk = 0; kk < max_kk; kk++) + { + const int k = k0 + kk; + float v = ptrA[(size_t)k * A_hstep]; + if (input_scale_ptr) + v *= input_scale_ptr[k]; + absmax = std::max(absmax, fabsf(v)); + } + + if (absmax == 0.f) + { + descale_ptr[g] = 0.f; + for (int k = 0; k < max_kk; k++) + outptr0[k0 + k] = 0; + continue; + } + + volatile double scale_fp64 = 127.0 / (double)absmax; + const float scale = (float)scale_fp64; + descale_ptr[g] = absmax / 127.f; + + for (int kk = 0; kk < max_kk; kk++) + { + const int k = k0 + kk; + float v = ptrA[(size_t)k * A_hstep]; + if (input_scale_ptr) + { + v *= input_scale_ptr[k]; + // preserve multiplication order for consistent rounding + asm volatile("" : "+f"(v)); + } + outptr0[k] = float2int8(v * scale); + } + } + } +} + +static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int block_size) +{ + const signed char* pAT = AT_tile; + const int A_hstep = AT_tile.w; + const float* pAT_descales = AT_descales_tile; + const int A_descales_hstep = AT_descales_tile.w; + const signed char* pBT = BT_tile; + const float* pBT_descales = BT_descales_tile; + float* outptr = topT_tile; + + int ii = 0; +#if __loongarch_sx + for (; ii + 7 < max_ii; ii += 8) + { + const signed char* pB = pBT; + const float* pB_descales = pBT_descales; + + int jj = 0; +#if __loongarch_asx + for (; jj + 7 < max_jj; jj += 8) + { + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + __m256i _sum0 = __lasx_xvreplgr2vr_w(0); + __m256i _sum1 = __lasx_xvreplgr2vr_w(0); + __m256i _sum2 = __lasx_xvreplgr2vr_w(0); + __m256i _sum3 = __lasx_xvreplgr2vr_w(0); + __m256i _sum4 = __lasx_xvreplgr2vr_w(0); + __m256i _sum5 = __lasx_xvreplgr2vr_w(0); + __m256i _sum6 = __lasx_xvreplgr2vr_w(0); + __m256i _sum7 = __lasx_xvreplgr2vr_w(0); + const int max_kk = std::min(K - k, block_size); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __m256i _pA = __lasx_xvld(pA, 0); + __m256i _pA1 = __lasx_xvshuf4i_w(_pA, _LSX_SHUFFLE(0, 3, 2, 1)); + __m256i _pA2 = __lasx_xvpermi_q(_pA, _pA, _LSX_SHUFFLE(0, 0, 0, 1)); + __m256i _pA3 = __lasx_xvshuf4i_w(_pA2, _LSX_SHUFFLE(0, 3, 2, 1)); + __m256i _pB0 = __lasx_xvld(pB, 0); + __m256i _pB1 = __lasx_xvshuf4i_w(_pB0, _LSX_SHUFFLE(1, 0, 3, 2)); + + __m256i _s0 = __lasx_xvmulwev_h_b(_pA, _pB0); + __m256i _s1 = __lasx_xvmulwev_h_b(_pA1, _pB0); + __m256i _s2 = __lasx_xvmulwev_h_b(_pA, _pB1); + __m256i _s3 = __lasx_xvmulwev_h_b(_pA1, _pB1); + __m256i _s4 = __lasx_xvmulwev_h_b(_pA2, _pB0); + __m256i _s5 = __lasx_xvmulwev_h_b(_pA3, _pB0); + __m256i _s6 = __lasx_xvmulwev_h_b(_pA2, _pB1); + __m256i _s7 = __lasx_xvmulwev_h_b(_pA3, _pB1); + _s0 = __lasx_xvmaddwod_h_b(_s0, _pA, _pB0); + _s1 = __lasx_xvmaddwod_h_b(_s1, _pA1, _pB0); + _s2 = __lasx_xvmaddwod_h_b(_s2, _pA, _pB1); + _s3 = __lasx_xvmaddwod_h_b(_s3, _pA1, _pB1); + _s4 = __lasx_xvmaddwod_h_b(_s4, _pA2, _pB0); + _s5 = __lasx_xvmaddwod_h_b(_s5, _pA3, _pB0); + _s6 = __lasx_xvmaddwod_h_b(_s6, _pA2, _pB1); + _s7 = __lasx_xvmaddwod_h_b(_s7, _pA3, _pB1); + _sum0 = __lasx_xvadd_w(_sum0, __lasx_xvhaddw_w_h(_s0, _s0)); + _sum1 = __lasx_xvadd_w(_sum1, __lasx_xvhaddw_w_h(_s1, _s1)); + _sum2 = __lasx_xvadd_w(_sum2, __lasx_xvhaddw_w_h(_s2, _s2)); + _sum3 = __lasx_xvadd_w(_sum3, __lasx_xvhaddw_w_h(_s3, _s3)); + _sum4 = __lasx_xvadd_w(_sum4, __lasx_xvhaddw_w_h(_s4, _s4)); + _sum5 = __lasx_xvadd_w(_sum5, __lasx_xvhaddw_w_h(_s5, _s5)); + _sum6 = __lasx_xvadd_w(_sum6, __lasx_xvhaddw_w_h(_s6, _s6)); + _sum7 = __lasx_xvadd_w(_sum7, __lasx_xvhaddw_w_h(_s7, _s7)); + pB += 32; + pA += 32; + } + if (kk + 1 < max_kk) + { + __m128i _pAs = __lsx_vld(pA, 0); + __m128i _pA8 = __lsx_vbsrl_v(_pAs, 8); + __m256i _pA = __lasx_concat_128(__lsx_vilvl_b(_pA8, _pAs), __lsx_vilvl_b(_pA8, _pAs)); + __m256i _pA1 = __lasx_xvshuf4i_h(_pA, _LSX_SHUFFLE(0, 3, 2, 1)); + __m256i _pA2 = __lasx_xvshuf4i_w(_pA, _LSX_SHUFFLE(1, 0, 3, 2)); + __m256i _pA3 = __lasx_xvshuf4i_h(_pA2, _LSX_SHUFFLE(0, 3, 2, 1)); + __m128i _pBs = __lsx_vld(pB, 0); + __m256i _pB0 = __lasx_concat_128(_pBs, _pBs); + __m256i _pB1 = __lasx_xvshuf4i_h(_pB0, _LSX_SHUFFLE(1, 0, 3, 2)); + __m256i _s0 = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA, _pB0), _pA, _pB0); + __m256i _s1 = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA1, _pB0), _pA1, _pB0); + __m256i _s2 = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA, _pB1), _pA, _pB1); + __m256i _s3 = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA1, _pB1), _pA1, _pB1); + __m256i _s4 = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA2, _pB0), _pA2, _pB0); + __m256i _s5 = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA3, _pB0), _pA3, _pB0); + __m256i _s6 = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA2, _pB1), _pA2, _pB1); + __m256i _s7 = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA3, _pB1), _pA3, _pB1); + _sum0 = __lasx_xvadd_w(_sum0, __lasx_vext2xv_w_h(_s0)); + _sum1 = __lasx_xvadd_w(_sum1, __lasx_vext2xv_w_h(_s1)); + _sum2 = __lasx_xvadd_w(_sum2, __lasx_vext2xv_w_h(_s2)); + _sum3 = __lasx_xvadd_w(_sum3, __lasx_vext2xv_w_h(_s3)); + _sum4 = __lasx_xvadd_w(_sum4, __lasx_vext2xv_w_h(_s4)); + _sum5 = __lasx_xvadd_w(_sum5, __lasx_vext2xv_w_h(_s5)); + _sum6 = __lasx_xvadd_w(_sum6, __lasx_vext2xv_w_h(_s6)); + _sum7 = __lasx_xvadd_w(_sum7, __lasx_vext2xv_w_h(_s7)); + pB += 16; + pA += 16; + kk += 2; + } + if (kk < max_kk) + { + __m256i _pA = __lasx_xvldrepl_d(pA, 0); + _pA = __lasx_xvilvl_b(__lasx_xvslti_b(_pA, 0), _pA); + __m256i _pA1 = __lasx_xvshuf4i_h(_pA, _LSX_SHUFFLE(0, 3, 2, 1)); + __m256i _pA2 = __lasx_xvshuf4i_w(_pA, _LSX_SHUFFLE(1, 0, 3, 2)); + __m256i _pA3 = __lasx_xvshuf4i_h(_pA2, _LSX_SHUFFLE(0, 3, 2, 1)); + __m256i _pB0 = __lasx_xvldrepl_d(pB, 0); + _pB0 = __lasx_xvilvl_b(__lasx_xvslti_b(_pB0, 0), _pB0); + __m256i _pB1 = __lasx_xvshuf4i_h(_pB0, _LSX_SHUFFLE(1, 0, 3, 2)); + __m256i _s0 = __lasx_xvmul_h(_pA, _pB0); + __m256i _s1 = __lasx_xvmul_h(_pA1, _pB0); + __m256i _s2 = __lasx_xvmul_h(_pA, _pB1); + __m256i _s3 = __lasx_xvmul_h(_pA1, _pB1); + __m256i _s4 = __lasx_xvmul_h(_pA2, _pB0); + __m256i _s5 = __lasx_xvmul_h(_pA3, _pB0); + __m256i _s6 = __lasx_xvmul_h(_pA2, _pB1); + __m256i _s7 = __lasx_xvmul_h(_pA3, _pB1); + _sum0 = __lasx_xvadd_w(_sum0, __lasx_vext2xv_w_h(_s0)); + _sum1 = __lasx_xvadd_w(_sum1, __lasx_vext2xv_w_h(_s1)); + _sum2 = __lasx_xvadd_w(_sum2, __lasx_vext2xv_w_h(_s2)); + _sum3 = __lasx_xvadd_w(_sum3, __lasx_vext2xv_w_h(_s3)); + _sum4 = __lasx_xvadd_w(_sum4, __lasx_vext2xv_w_h(_s4)); + _sum5 = __lasx_xvadd_w(_sum5, __lasx_vext2xv_w_h(_s5)); + _sum6 = __lasx_xvadd_w(_sum6, __lasx_vext2xv_w_h(_s6)); + _sum7 = __lasx_xvadd_w(_sum7, __lasx_vext2xv_w_h(_s7)); + pB += 8; + pA += 8; + } + + __m256 _bscale = (__m256)__lasx_xvld(pB_descales, 0); + __m256 _bscale1 = (__m256)__lasx_xvshuf4i_w((__m256i)_bscale, _LSX_SHUFFLE(1, 0, 3, 2)); + __m256 _ascale = (__m256)__lasx_xvld(pA_descales, 0); + __m256 _ascale1 = (__m256)__lasx_xvshuf4i_w((__m256i)_ascale, _LSX_SHUFFLE(0, 3, 2, 1)); + __m256 _ascale2 = (__m256)__lasx_xvpermi_q((__m256i)_ascale, (__m256i)_ascale, _LSX_SHUFFLE(0, 0, 0, 1)); + __m256 _ascale3 = (__m256)__lasx_xvshuf4i_w((__m256i)_ascale2, _LSX_SHUFFLE(0, 3, 2, 1)); + __m256 _out0 = k == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr, 0); + __m256 _out1 = k == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr + 8, 0); + __m256 _out2 = k == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr + 16, 0); + __m256 _out3 = k == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr + 24, 0); + __m256 _out4 = k == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr + 32, 0); + __m256 _out5 = k == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr + 40, 0); + __m256 _out6 = k == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr + 48, 0); + __m256 _out7 = k == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr + 56, 0); + _out0 = __lasx_xvfmadd_s((__m256)__lasx_xvffint_s_w(_sum0), __lasx_xvfmul_s(_ascale, _bscale), _out0); + _out1 = __lasx_xvfmadd_s((__m256)__lasx_xvffint_s_w(_sum1), __lasx_xvfmul_s(_ascale1, _bscale), _out1); + _out2 = __lasx_xvfmadd_s((__m256)__lasx_xvffint_s_w(_sum2), __lasx_xvfmul_s(_ascale, _bscale1), _out2); + _out3 = __lasx_xvfmadd_s((__m256)__lasx_xvffint_s_w(_sum3), __lasx_xvfmul_s(_ascale1, _bscale1), _out3); + _out4 = __lasx_xvfmadd_s((__m256)__lasx_xvffint_s_w(_sum4), __lasx_xvfmul_s(_ascale2, _bscale), _out4); + _out5 = __lasx_xvfmadd_s((__m256)__lasx_xvffint_s_w(_sum5), __lasx_xvfmul_s(_ascale3, _bscale), _out5); + _out6 = __lasx_xvfmadd_s((__m256)__lasx_xvffint_s_w(_sum6), __lasx_xvfmul_s(_ascale2, _bscale1), _out6); + _out7 = __lasx_xvfmadd_s((__m256)__lasx_xvffint_s_w(_sum7), __lasx_xvfmul_s(_ascale3, _bscale1), _out7); + __lasx_xvst(_out0, outptr, 0); + __lasx_xvst(_out1, outptr + 8, 0); + __lasx_xvst(_out2, outptr + 16, 0); + __lasx_xvst(_out3, outptr + 24, 0); + __lasx_xvst(_out4, outptr + 32, 0); + __lasx_xvst(_out5, outptr + 40, 0); + __lasx_xvst(_out6, outptr + 48, 0); + __lasx_xvst(_out7, outptr + 56, 0); + pA_descales += 8; + pB_descales += 8; + } + outptr += 64; + } +#endif // __loongarch_asx + for (; jj + 3 < max_jj; jj += 4) + { + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + __m128i _sum00 = __lsx_vreplgr2vr_w(0); + __m128i _sum01 = __lsx_vreplgr2vr_w(0); + __m128i _sum10 = __lsx_vreplgr2vr_w(0); + __m128i _sum11 = __lsx_vreplgr2vr_w(0); + __m128i _sum20 = __lsx_vreplgr2vr_w(0); + __m128i _sum21 = __lsx_vreplgr2vr_w(0); + __m128i _sum30 = __lsx_vreplgr2vr_w(0); + __m128i _sum31 = __lsx_vreplgr2vr_w(0); + const int max_kk = std::min(K - k, block_size); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __m128i _pA0 = __lsx_vld(pA, 0); + __m128i _pA1 = __lsx_vld(pA + 16, 0); + __m128i _pA0r = __lsx_vshuf4i_w(_pA0, _LSX_SHUFFLE(1, 0, 3, 2)); + __m128i _pA1r = __lsx_vshuf4i_w(_pA1, _LSX_SHUFFLE(1, 0, 3, 2)); + __m128i _pB0 = __lsx_vld(pB, 0); + __m128i _pB0r = __lsx_vshuf4i_w(_pB0, _LSX_SHUFFLE(0, 3, 2, 1)); + + __m128i _s0 = __lsx_vmulwev_h_b(_pA0, _pB0); + __m128i _s1 = __lsx_vmulwev_h_b(_pA1, _pB0); + _s0 = __lsx_vmaddwod_h_b(_s0, _pA0, _pB0); + _s1 = __lsx_vmaddwod_h_b(_s1, _pA1, _pB0); + _sum00 = __lsx_vadd_w(_sum00, __lsx_vhaddw_w_h(_s0, _s0)); + _sum01 = __lsx_vadd_w(_sum01, __lsx_vhaddw_w_h(_s1, _s1)); + + _s0 = __lsx_vmulwev_h_b(_pA0, _pB0r); + _s1 = __lsx_vmulwev_h_b(_pA1, _pB0r); + _s0 = __lsx_vmaddwod_h_b(_s0, _pA0, _pB0r); + _s1 = __lsx_vmaddwod_h_b(_s1, _pA1, _pB0r); + _sum10 = __lsx_vadd_w(_sum10, __lsx_vhaddw_w_h(_s0, _s0)); + _sum11 = __lsx_vadd_w(_sum11, __lsx_vhaddw_w_h(_s1, _s1)); + + _s0 = __lsx_vmulwev_h_b(_pA0r, _pB0); + _s1 = __lsx_vmulwev_h_b(_pA1r, _pB0); + _s0 = __lsx_vmaddwod_h_b(_s0, _pA0r, _pB0); + _s1 = __lsx_vmaddwod_h_b(_s1, _pA1r, _pB0); + _sum20 = __lsx_vadd_w(_sum20, __lsx_vhaddw_w_h(_s0, _s0)); + _sum21 = __lsx_vadd_w(_sum21, __lsx_vhaddw_w_h(_s1, _s1)); + + _s0 = __lsx_vmulwev_h_b(_pA0r, _pB0r); + _s1 = __lsx_vmulwev_h_b(_pA1r, _pB0r); + _s0 = __lsx_vmaddwod_h_b(_s0, _pA0r, _pB0r); + _s1 = __lsx_vmaddwod_h_b(_s1, _pA1r, _pB0r); + _sum30 = __lsx_vadd_w(_sum30, __lsx_vhaddw_w_h(_s0, _s0)); + _sum31 = __lsx_vadd_w(_sum31, __lsx_vhaddw_w_h(_s1, _s1)); + pB += 16; + pA += 32; + } + if (kk + 1 < max_kk) + { + __m128i _pAs = __lsx_vld(pA, 0); + __m128i _pA8 = __lsx_vbsrl_v(_pAs, 8); + __m128i _pA = __lsx_vilvl_b(_pA8, _pAs); + __m128i _pA0 = __lsx_vreplvei_d(_pA, 0); + __m128i _pA1 = __lsx_vreplvei_d(_pA, 1); + __m128i _pA0r = __lsx_vshuf4i_h(_pA0, _LSX_SHUFFLE(1, 0, 3, 2)); + __m128i _pA1r = __lsx_vshuf4i_h(_pA1, _LSX_SHUFFLE(1, 0, 3, 2)); + __m128i _pBs = __lsx_vldrepl_d(pB, 0); + __m128i _pB0r = __lsx_vshuf4i_h(_pBs, _LSX_SHUFFLE(0, 3, 2, 1)); + __m128i _s0 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pBs), _pA0, _pBs); + __m128i _s1 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA1, _pBs), _pA1, _pBs); + _sum00 = __lsx_vadd_w(_sum00, __lsx_vilvl_h(__lsx_vslti_h(_s0, 0), _s0)); + _sum01 = __lsx_vadd_w(_sum01, __lsx_vilvl_h(__lsx_vslti_h(_s1, 0), _s1)); + _s0 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB0r), _pA0, _pB0r); + _s1 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA1, _pB0r), _pA1, _pB0r); + _sum10 = __lsx_vadd_w(_sum10, __lsx_vilvl_h(__lsx_vslti_h(_s0, 0), _s0)); + _sum11 = __lsx_vadd_w(_sum11, __lsx_vilvl_h(__lsx_vslti_h(_s1, 0), _s1)); + _s0 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0r, _pBs), _pA0r, _pBs); + _s1 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA1r, _pBs), _pA1r, _pBs); + _sum20 = __lsx_vadd_w(_sum20, __lsx_vilvl_h(__lsx_vslti_h(_s0, 0), _s0)); + _sum21 = __lsx_vadd_w(_sum21, __lsx_vilvl_h(__lsx_vslti_h(_s1, 0), _s1)); + _s0 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0r, _pB0r), _pA0r, _pB0r); + _s1 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA1r, _pB0r), _pA1r, _pB0r); + _sum30 = __lsx_vadd_w(_sum30, __lsx_vilvl_h(__lsx_vslti_h(_s0, 0), _s0)); + _sum31 = __lsx_vadd_w(_sum31, __lsx_vilvl_h(__lsx_vslti_h(_s1, 0), _s1)); + pB += 8; + pA += 16; + kk += 2; + } + if (kk < max_kk) + { + __m128i _pA0 = __lsx_vldrepl_w(pA, 0); + _pA0 = __lsx_vilvl_b(__lsx_vslti_b(_pA0, 0), _pA0); + __m128i _pA1 = __lsx_vldrepl_w(pA + 4, 0); + _pA1 = __lsx_vilvl_b(__lsx_vslti_b(_pA1, 0), _pA1); + __m128i _pA0r = __lsx_vshuf4i_h(_pA0, _LSX_SHUFFLE(1, 0, 3, 2)); + __m128i _pA1r = __lsx_vshuf4i_h(_pA1, _LSX_SHUFFLE(1, 0, 3, 2)); + __m128i _pB0 = __lsx_vldrepl_w(pB, 0); + _pB0 = __lsx_vilvl_b(__lsx_vslti_b(_pB0, 0), _pB0); + __m128i _pB0r = __lsx_vshuf4i_h(_pB0, _LSX_SHUFFLE(0, 3, 2, 1)); + __m128i _s0 = __lsx_vmul_h(_pA0, _pB0); + __m128i _s1 = __lsx_vmul_h(_pA1, _pB0); + _sum00 = __lsx_vadd_w(_sum00, __lsx_vilvl_h(__lsx_vslti_h(_s0, 0), _s0)); + _sum01 = __lsx_vadd_w(_sum01, __lsx_vilvl_h(__lsx_vslti_h(_s1, 0), _s1)); + _s0 = __lsx_vmul_h(_pA0, _pB0r); + _s1 = __lsx_vmul_h(_pA1, _pB0r); + _sum10 = __lsx_vadd_w(_sum10, __lsx_vilvl_h(__lsx_vslti_h(_s0, 0), _s0)); + _sum11 = __lsx_vadd_w(_sum11, __lsx_vilvl_h(__lsx_vslti_h(_s1, 0), _s1)); + _s0 = __lsx_vmul_h(_pA0r, _pB0); + _s1 = __lsx_vmul_h(_pA1r, _pB0); + _sum20 = __lsx_vadd_w(_sum20, __lsx_vilvl_h(__lsx_vslti_h(_s0, 0), _s0)); + _sum21 = __lsx_vadd_w(_sum21, __lsx_vilvl_h(__lsx_vslti_h(_s1, 0), _s1)); + _s0 = __lsx_vmul_h(_pA0r, _pB0r); + _s1 = __lsx_vmul_h(_pA1r, _pB0r); + _sum30 = __lsx_vadd_w(_sum30, __lsx_vilvl_h(__lsx_vslti_h(_s0, 0), _s0)); + _sum31 = __lsx_vadd_w(_sum31, __lsx_vilvl_h(__lsx_vslti_h(_s1, 0), _s1)); + pB += 4; + pA += 8; + } + + __m128 _bscale = (__m128)__lsx_vld(pB_descales, 0); + __m128 _bscaler = (__m128)__lsx_vshuf4i_w((__m128i)_bscale, _LSX_SHUFFLE(0, 3, 2, 1)); + __m128 _ascale0 = (__m128)__lsx_vld(pA_descales, 0); + __m128 _ascale1 = (__m128)__lsx_vld(pA_descales + 4, 0); + __m128 _ascale0r = (__m128)__lsx_vshuf4i_w((__m128i)_ascale0, _LSX_SHUFFLE(1, 0, 3, 2)); + __m128 _ascale1r = (__m128)__lsx_vshuf4i_w((__m128i)_ascale1, _LSX_SHUFFLE(1, 0, 3, 2)); + __m128 _out00 = k == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr, 0); + __m128 _out01 = k == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 4, 0); + __m128 _out10 = k == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 8, 0); + __m128 _out11 = k == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 12, 0); + __m128 _out20 = k == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 16, 0); + __m128 _out21 = k == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 20, 0); + __m128 _out30 = k == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 24, 0); + __m128 _out31 = k == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 28, 0); + _out00 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum00), __lsx_vfmul_s(_ascale0, _bscale), _out00); + _out01 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum01), __lsx_vfmul_s(_ascale1, _bscale), _out01); + _out10 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum10), __lsx_vfmul_s(_ascale0, _bscaler), _out10); + _out11 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum11), __lsx_vfmul_s(_ascale1, _bscaler), _out11); + _out20 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum20), __lsx_vfmul_s(_ascale0r, _bscale), _out20); + _out21 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum21), __lsx_vfmul_s(_ascale1r, _bscale), _out21); + _out30 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum30), __lsx_vfmul_s(_ascale0r, _bscaler), _out30); + _out31 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum31), __lsx_vfmul_s(_ascale1r, _bscaler), _out31); + __lsx_vst((__m128i)_out00, outptr, 0); + __lsx_vst((__m128i)_out01, outptr + 4, 0); + __lsx_vst((__m128i)_out10, outptr + 8, 0); + __lsx_vst((__m128i)_out11, outptr + 12, 0); + __lsx_vst((__m128i)_out20, outptr + 16, 0); + __lsx_vst((__m128i)_out21, outptr + 20, 0); + __lsx_vst((__m128i)_out30, outptr + 24, 0); + __lsx_vst((__m128i)_out31, outptr + 28, 0); + pA_descales += 8; + pB_descales += 4; + } + outptr += 32; + } + + for (; jj + 1 < max_jj; jj += 2) + { + __m128 _out00 = (__m128)__lsx_vldi(0); + __m128 _out01 = (__m128)__lsx_vldi(0); + __m128 _out10 = (__m128)__lsx_vldi(0); + __m128 _out11 = (__m128)__lsx_vldi(0); + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + __m128i _sum00 = __lsx_vreplgr2vr_w(0); + __m128i _sum01 = __lsx_vreplgr2vr_w(0); + __m128i _sum10 = __lsx_vreplgr2vr_w(0); + __m128i _sum11 = __lsx_vreplgr2vr_w(0); + const int max_kk = std::min(K - k, block_size); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __m128i _pA0 = __lsx_vld(pA, 0); + __m128i _pA1 = __lsx_vld(pA + 16, 0); + __m128i _pB = __lsx_vldrepl_d(pB, 0); + __m128i _pB0 = __lsx_vreplvei_w(_pB, 0); + __m128i _pB1 = __lsx_vreplvei_w(_pB, 1); + __m128i _s0 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB0), _pA0, _pB0); + __m128i _s1 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA1, _pB0), _pA1, _pB0); + _sum00 = __lsx_vadd_w(_sum00, __lsx_vhaddw_w_h(_s0, _s0)); + _sum01 = __lsx_vadd_w(_sum01, __lsx_vhaddw_w_h(_s1, _s1)); + _s0 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB1), _pA0, _pB1); + _s1 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA1, _pB1), _pA1, _pB1); + _sum10 = __lsx_vadd_w(_sum10, __lsx_vhaddw_w_h(_s0, _s0)); + _sum11 = __lsx_vadd_w(_sum11, __lsx_vhaddw_w_h(_s1, _s1)); + pA += 32; + pB += 8; + } + if (kk + 1 < max_kk) + { + __m128i _pAs = __lsx_vld(pA, 0); + __m128i _pA8 = __lsx_vbsrl_v(_pAs, 8); + __m128i _pA = __lsx_vilvl_b(_pA8, _pAs); + __m128i _pA0 = __lsx_vreplvei_d(_pA, 0); + __m128i _pA1 = __lsx_vreplvei_d(_pA, 1); + __m128i _pB = __lsx_vreplgr2vr_w(*(const int*)pB); + __m128i _pB0 = __lsx_vreplvei_h(_pB, 0); + __m128i _pB1 = __lsx_vreplvei_h(_pB, 1); + __m128i _s0 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB0), _pA0, _pB0); + __m128i _s1 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA1, _pB0), _pA1, _pB0); + _sum00 = __lsx_vadd_w(_sum00, __lsx_vilvl_h(__lsx_vslti_h(_s0, 0), _s0)); + _sum01 = __lsx_vadd_w(_sum01, __lsx_vilvl_h(__lsx_vslti_h(_s1, 0), _s1)); + _s0 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB1), _pA0, _pB1); + _s1 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA1, _pB1), _pA1, _pB1); + _sum10 = __lsx_vadd_w(_sum10, __lsx_vilvl_h(__lsx_vslti_h(_s0, 0), _s0)); + _sum11 = __lsx_vadd_w(_sum11, __lsx_vilvl_h(__lsx_vslti_h(_s1, 0), _s1)); + pA += 16; + pB += 4; + kk += 2; + } + if (kk < max_kk) + { + __m128i _pA = __lsx_vldrepl_d(pA, 0); + _pA = __lsx_vilvl_b(__lsx_vslti_b(_pA, 0), _pA); + __m128i _pA0 = __lsx_vreplvei_d(_pA, 0); + __m128i _pA1 = __lsx_vreplvei_d(_pA, 1); + __m128i _pB0 = __lsx_vreplgr2vr_h((signed char)pB[0]); + __m128i _pB1 = __lsx_vreplgr2vr_h((signed char)pB[1]); + __m128i _s0 = __lsx_vmul_h(_pA0, _pB0); + __m128i _s1 = __lsx_vmul_h(_pA1, _pB0); + _sum00 = __lsx_vadd_w(_sum00, __lsx_vilvl_h(__lsx_vslti_h(_s0, 0), _s0)); + _sum01 = __lsx_vadd_w(_sum01, __lsx_vilvl_h(__lsx_vslti_h(_s1, 0), _s1)); + _s0 = __lsx_vmul_h(_pA0, _pB1); + _s1 = __lsx_vmul_h(_pA1, _pB1); + _sum10 = __lsx_vadd_w(_sum10, __lsx_vilvl_h(__lsx_vslti_h(_s0, 0), _s0)); + _sum11 = __lsx_vadd_w(_sum11, __lsx_vilvl_h(__lsx_vslti_h(_s1, 0), _s1)); + pA += 8; + pB += 2; + } + __m128 _ascale0 = (__m128)__lsx_vld(pA_descales, 0); + __m128 _ascale1 = (__m128)__lsx_vld(pA_descales + 4, 0); + __m128 _bscale = __lsx_vreplfr2vr_s(pB_descales[0]); + _out00 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum00), __lsx_vfmul_s(_ascale0, _bscale), _out00); + _out01 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum01), __lsx_vfmul_s(_ascale1, _bscale), _out01); + _bscale = __lsx_vreplfr2vr_s(pB_descales[1]); + _out10 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum10), __lsx_vfmul_s(_ascale0, _bscale), _out10); + _out11 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum11), __lsx_vfmul_s(_ascale1, _bscale), _out11); + pA_descales += 8; + pB_descales += 2; + } + __lsx_vst((__m128i)_out00, outptr, 0); + __lsx_vst((__m128i)_out01, outptr + 4, 0); + __lsx_vst((__m128i)_out10, outptr + 8, 0); + __lsx_vst((__m128i)_out11, outptr + 12, 0); + outptr += 16; + } + for (; jj < max_jj; jj++) + { + __m128 _out0 = (__m128)__lsx_vldi(0); + __m128 _out1 = (__m128)__lsx_vldi(0); + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + __m128i _sum0 = __lsx_vreplgr2vr_w(0); + __m128i _sum1 = __lsx_vreplgr2vr_w(0); + const int max_kk = std::min(K - k, block_size); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __m128i _pA0 = __lsx_vld(pA, 0); + __m128i _pA1 = __lsx_vld(pA + 16, 0); + __m128i _pB = __lsx_vreplgr2vr_w(*(const int*)pB); + __m128i _s0 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB), _pA0, _pB); + __m128i _s1 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA1, _pB), _pA1, _pB); + _sum0 = __lsx_vadd_w(_sum0, __lsx_vhaddw_w_h(_s0, _s0)); + _sum1 = __lsx_vadd_w(_sum1, __lsx_vhaddw_w_h(_s1, _s1)); + pA += 32; + pB += 4; + } + if (kk + 1 < max_kk) + { + __m128i _pAs = __lsx_vld(pA, 0); + __m128i _pA8 = __lsx_vbsrl_v(_pAs, 8); + __m128i _pA = __lsx_vilvl_b(_pA8, _pAs); + __m128i _pA0 = __lsx_vreplvei_d(_pA, 0); + __m128i _pA1 = __lsx_vreplvei_d(_pA, 1); + __m128i _pB = __lsx_vreplgr2vr_h((unsigned char)pB[0] | ((unsigned char)pB[1] << 8)); + __m128i _s0 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB), _pA0, _pB); + __m128i _s1 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA1, _pB), _pA1, _pB); + _sum0 = __lsx_vadd_w(_sum0, __lsx_vilvl_h(__lsx_vslti_h(_s0, 0), _s0)); + _sum1 = __lsx_vadd_w(_sum1, __lsx_vilvl_h(__lsx_vslti_h(_s1, 0), _s1)); + pA += 16; + pB += 2; + kk += 2; + } + if (kk < max_kk) + { + __m128i _pA = __lsx_vldrepl_d(pA, 0); + _pA = __lsx_vilvl_b(__lsx_vslti_b(_pA, 0), _pA); + __m128i _pA0 = __lsx_vreplvei_d(_pA, 0); + __m128i _pA1 = __lsx_vreplvei_d(_pA, 1); + __m128i _pB = __lsx_vreplgr2vr_h((signed char)pB[0]); + __m128i _s0 = __lsx_vmul_h(_pA0, _pB); + __m128i _s1 = __lsx_vmul_h(_pA1, _pB); + _sum0 = __lsx_vadd_w(_sum0, __lsx_vilvl_h(__lsx_vslti_h(_s0, 0), _s0)); + _sum1 = __lsx_vadd_w(_sum1, __lsx_vilvl_h(__lsx_vslti_h(_s1, 0), _s1)); + pA += 8; + pB++; + } + __m128 _bscale = __lsx_vreplfr2vr_s(*pB_descales++); + _out0 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum0), __lsx_vfmul_s((__m128)__lsx_vld(pA_descales, 0), _bscale), _out0); + _out1 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum1), __lsx_vfmul_s((__m128)__lsx_vld(pA_descales + 4, 0), _bscale), _out1); + pA_descales += 8; + } + __lsx_vst((__m128i)_out0, outptr, 0); + __lsx_vst((__m128i)_out1, outptr + 4, 0); + outptr += 8; + } + + pAT += K * 8; + pAT_descales += (K + block_size - 1) / block_size * 8; + } +#endif // __loongarch_sx +#if __loongarch_sx + for (; ii + 3 < max_ii; ii += 4) + { + const signed char* pB = pBT; + const float* pB_descales = pBT_descales; + + int jj = 0; +#if __loongarch_asx + for (; jj + 15 < max_jj; jj += 16) + { + const signed char* pB0 = pB; + const signed char* pB1 = pB + 8 * K; + const float* pB_descales0 = pB_descales; + const float* pB_descales1 = pB_descales + 8 * ((K + block_size - 1) / block_size); + __m256 _out00 = (__m256)__lasx_xvldi(0); + __m256 _out01 = (__m256)__lasx_xvldi(0); + __m256 _out10 = (__m256)__lasx_xvldi(0); + __m256 _out11 = (__m256)__lasx_xvldi(0); + __m256 _out20 = (__m256)__lasx_xvldi(0); + __m256 _out21 = (__m256)__lasx_xvldi(0); + __m256 _out30 = (__m256)__lasx_xvldi(0); + __m256 _out31 = (__m256)__lasx_xvldi(0); + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + __m256i _sum00 = __lasx_xvreplgr2vr_w(0); + __m256i _sum01 = __lasx_xvreplgr2vr_w(0); + __m256i _sum10 = __lasx_xvreplgr2vr_w(0); + __m256i _sum11 = __lasx_xvreplgr2vr_w(0); + __m256i _sum20 = __lasx_xvreplgr2vr_w(0); + __m256i _sum21 = __lasx_xvreplgr2vr_w(0); + __m256i _sum30 = __lasx_xvreplgr2vr_w(0); + __m256i _sum31 = __lasx_xvreplgr2vr_w(0); + const int max_kk = std::min(K - k, block_size); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __m256i _pB0 = __lasx_xvld(pB0, 0); + __m256i _pB1 = __lasx_xvld(pB1, 0); + __m128i _pA = __lsx_vld(pA, 0); + __m128i _pA0_128 = __lsx_vreplvei_w(_pA, 0); + __m128i _pA1_128 = __lsx_vreplvei_w(_pA, 1); + __m128i _pA2_128 = __lsx_vreplvei_w(_pA, 2); + __m128i _pA3_128 = __lsx_vreplvei_w(_pA, 3); + __m256i _pA0 = __lasx_concat_128(_pA0_128, _pA0_128); + __m256i _pA1 = __lasx_concat_128(_pA1_128, _pA1_128); + __m256i _pA2 = __lasx_concat_128(_pA2_128, _pA2_128); + __m256i _pA3 = __lasx_concat_128(_pA3_128, _pA3_128); + __m256i _s = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA0, _pB0), _pA0, _pB0); + _sum00 = __lasx_xvadd_w(_sum00, __lasx_xvhaddw_w_h(_s, _s)); + _s = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA0, _pB1), _pA0, _pB1); + _sum01 = __lasx_xvadd_w(_sum01, __lasx_xvhaddw_w_h(_s, _s)); + _s = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA1, _pB0), _pA1, _pB0); + _sum10 = __lasx_xvadd_w(_sum10, __lasx_xvhaddw_w_h(_s, _s)); + _s = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA1, _pB1), _pA1, _pB1); + _sum11 = __lasx_xvadd_w(_sum11, __lasx_xvhaddw_w_h(_s, _s)); + _s = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA2, _pB0), _pA2, _pB0); + _sum20 = __lasx_xvadd_w(_sum20, __lasx_xvhaddw_w_h(_s, _s)); + _s = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA2, _pB1), _pA2, _pB1); + _sum21 = __lasx_xvadd_w(_sum21, __lasx_xvhaddw_w_h(_s, _s)); + _s = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA3, _pB0), _pA3, _pB0); + _sum30 = __lasx_xvadd_w(_sum30, __lasx_xvhaddw_w_h(_s, _s)); + _s = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA3, _pB1), _pA3, _pB1); + _sum31 = __lasx_xvadd_w(_sum31, __lasx_xvhaddw_w_h(_s, _s)); + pB0 += 32; + pB1 += 32; + pA += 16; + } + if (kk + 1 < max_kk) + { + __m128i _pB0 = __lsx_vld(pB0, 0); + __m128i _pB1 = __lsx_vld(pB1, 0); + __m128i _pA = __lsx_vldrepl_d(pA, 0); + __m128i _pA0 = __lsx_vreplvei_h(_pA, 0); + __m128i _pA1 = __lsx_vreplvei_h(_pA, 1); + __m128i _pA2 = __lsx_vreplvei_h(_pA, 2); + __m128i _pA3 = __lsx_vreplvei_h(_pA, 3); + __m128i _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB0), _pA0, _pB0); + _sum00 = __lasx_xvadd_w(_sum00, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB1), _pA0, _pB1); + _sum01 = __lasx_xvadd_w(_sum01, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA1, _pB0), _pA1, _pB0); + _sum10 = __lasx_xvadd_w(_sum10, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA1, _pB1), _pA1, _pB1); + _sum11 = __lasx_xvadd_w(_sum11, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA2, _pB0), _pA2, _pB0); + _sum20 = __lasx_xvadd_w(_sum20, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA2, _pB1), _pA2, _pB1); + _sum21 = __lasx_xvadd_w(_sum21, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA3, _pB0), _pA3, _pB0); + _sum30 = __lasx_xvadd_w(_sum30, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA3, _pB1), _pA3, _pB1); + _sum31 = __lasx_xvadd_w(_sum31, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); + pB0 += 16; + pB1 += 16; + pA += 8; + kk += 2; + } + if (kk < max_kk) + { + __m128i _pB0 = __lsx_vldrepl_d(pB0, 0); + __m128i _pB1 = __lsx_vldrepl_d(pB1, 0); + _pB0 = __lsx_vilvl_b(__lsx_vslti_b(_pB0, 0), _pB0); + _pB1 = __lsx_vilvl_b(__lsx_vslti_b(_pB1, 0), _pB1); + const int a0123 = __lsx_vpickve2gr_w(__lsx_vldrepl_w(pA, 0), 0); + __m128i _s = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)a0123), _pB0); + _sum00 = __lasx_xvadd_w(_sum00, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); + _s = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)a0123), _pB1); + _sum01 = __lasx_xvadd_w(_sum01, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); + _s = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)(a0123 >> 8)), _pB0); + _sum10 = __lasx_xvadd_w(_sum10, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); + _s = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)(a0123 >> 8)), _pB1); + _sum11 = __lasx_xvadd_w(_sum11, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); + _s = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)(a0123 >> 16)), _pB0); + _sum20 = __lasx_xvadd_w(_sum20, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); + _s = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)(a0123 >> 16)), _pB1); + _sum21 = __lasx_xvadd_w(_sum21, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); + _s = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)(a0123 >> 24)), _pB0); + _sum30 = __lasx_xvadd_w(_sum30, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); + _s = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)(a0123 >> 24)), _pB1); + _sum31 = __lasx_xvadd_w(_sum31, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); + pB0 += 8; + pB1 += 8; + pA += 4; + } + + __m256 _bscale0 = (__m256)__lasx_xvld(pB_descales0, 0); + __m256 _bscale1 = (__m256)__lasx_xvld(pB_descales1, 0); + __m256 _ascale = (__m256)__lasx_xvreplfr2vr_s(pA_descales[0]); + _out00 = __lasx_xvfmadd_s((__m256)__lasx_xvffint_s_w(_sum00), __lasx_xvfmul_s(_bscale0, _ascale), _out00); + _out01 = __lasx_xvfmadd_s((__m256)__lasx_xvffint_s_w(_sum01), __lasx_xvfmul_s(_bscale1, _ascale), _out01); + _ascale = (__m256)__lasx_xvreplfr2vr_s(pA_descales[1]); + _out10 = __lasx_xvfmadd_s((__m256)__lasx_xvffint_s_w(_sum10), __lasx_xvfmul_s(_bscale0, _ascale), _out10); + _out11 = __lasx_xvfmadd_s((__m256)__lasx_xvffint_s_w(_sum11), __lasx_xvfmul_s(_bscale1, _ascale), _out11); + _ascale = (__m256)__lasx_xvreplfr2vr_s(pA_descales[2]); + _out20 = __lasx_xvfmadd_s((__m256)__lasx_xvffint_s_w(_sum20), __lasx_xvfmul_s(_bscale0, _ascale), _out20); + _out21 = __lasx_xvfmadd_s((__m256)__lasx_xvffint_s_w(_sum21), __lasx_xvfmul_s(_bscale1, _ascale), _out21); + _ascale = (__m256)__lasx_xvreplfr2vr_s(pA_descales[3]); + _out30 = __lasx_xvfmadd_s((__m256)__lasx_xvffint_s_w(_sum30), __lasx_xvfmul_s(_bscale0, _ascale), _out30); + _out31 = __lasx_xvfmadd_s((__m256)__lasx_xvffint_s_w(_sum31), __lasx_xvfmul_s(_bscale1, _ascale), _out31); + pA_descales += 4; + pB_descales0 += 8; + pB_descales1 += 8; + } + + pB = pB1; + pB_descales = pB_descales1; + + __lasx_xvst(_out00, outptr + (ii + 0) * max_jj + jj, 0); + __lasx_xvst(_out01, outptr + (ii + 0) * max_jj + jj + 8, 0); + __lasx_xvst(_out10, outptr + (ii + 1) * max_jj + jj, 0); + __lasx_xvst(_out11, outptr + (ii + 1) * max_jj + jj + 8, 0); + __lasx_xvst(_out20, outptr + (ii + 2) * max_jj + jj, 0); + __lasx_xvst(_out21, outptr + (ii + 2) * max_jj + jj + 8, 0); + __lasx_xvst(_out30, outptr + (ii + 3) * max_jj + jj, 0); + __lasx_xvst(_out31, outptr + (ii + 3) * max_jj + jj + 8, 0); + } + for (; jj + 7 < max_jj; jj += 8) + { + __m256 _out0 = (__m256)__lasx_xvldi(0); + __m256 _out1 = (__m256)__lasx_xvldi(0); + __m256 _out2 = (__m256)__lasx_xvldi(0); + __m256 _out3 = (__m256)__lasx_xvldi(0); + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + __m256i _sum0 = __lasx_xvreplgr2vr_w(0); + __m256i _sum1 = __lasx_xvreplgr2vr_w(0); + __m256i _sum2 = __lasx_xvreplgr2vr_w(0); + __m256i _sum3 = __lasx_xvreplgr2vr_w(0); + const int max_kk = std::min(K - k, block_size); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __m256i _pB = __lasx_xvld(pB, 0); + __m128i _pA = __lsx_vld(pA, 0); + __m128i _pA0_128 = __lsx_vreplvei_w(_pA, 0); + __m128i _pA1_128 = __lsx_vreplvei_w(_pA, 1); + __m128i _pA2_128 = __lsx_vreplvei_w(_pA, 2); + __m128i _pA3_128 = __lsx_vreplvei_w(_pA, 3); + __m256i _pA0 = __lasx_concat_128(_pA0_128, _pA0_128); + __m256i _pA1 = __lasx_concat_128(_pA1_128, _pA1_128); + __m256i _pA2 = __lasx_concat_128(_pA2_128, _pA2_128); + __m256i _pA3 = __lasx_concat_128(_pA3_128, _pA3_128); + __m256i _s0 = __lasx_xvmulwev_h_b(_pA0, _pB); + __m256i _s1 = __lasx_xvmulwev_h_b(_pA1, _pB); + __m256i _s2 = __lasx_xvmulwev_h_b(_pA2, _pB); + __m256i _s3 = __lasx_xvmulwev_h_b(_pA3, _pB); + _s0 = __lasx_xvmaddwod_h_b(_s0, _pA0, _pB); + _s1 = __lasx_xvmaddwod_h_b(_s1, _pA1, _pB); + _s2 = __lasx_xvmaddwod_h_b(_s2, _pA2, _pB); + _s3 = __lasx_xvmaddwod_h_b(_s3, _pA3, _pB); + _sum0 = __lasx_xvadd_w(_sum0, __lasx_xvhaddw_w_h(_s0, _s0)); + _sum1 = __lasx_xvadd_w(_sum1, __lasx_xvhaddw_w_h(_s1, _s1)); + _sum2 = __lasx_xvadd_w(_sum2, __lasx_xvhaddw_w_h(_s2, _s2)); + _sum3 = __lasx_xvadd_w(_sum3, __lasx_xvhaddw_w_h(_s3, _s3)); + pB += 32; + pA += 16; + } + if (kk + 1 < max_kk) + { + __m128i _pB = __lsx_vld(pB, 0); + __m128i _pA = __lsx_vldrepl_d(pA, 0); + __m128i _pA0 = __lsx_vreplvei_h(_pA, 0); + __m128i _pA1 = __lsx_vreplvei_h(_pA, 1); + __m128i _pA2 = __lsx_vreplvei_h(_pA, 2); + __m128i _pA3 = __lsx_vreplvei_h(_pA, 3); + __m128i _s0 = __lsx_vmulwev_h_b(_pA0, _pB); + __m128i _s1 = __lsx_vmulwev_h_b(_pA1, _pB); + __m128i _s2 = __lsx_vmulwev_h_b(_pA2, _pB); + __m128i _s3 = __lsx_vmulwev_h_b(_pA3, _pB); + _s0 = __lsx_vmaddwod_h_b(_s0, _pA0, _pB); + _s1 = __lsx_vmaddwod_h_b(_s1, _pA1, _pB); + _s2 = __lsx_vmaddwod_h_b(_s2, _pA2, _pB); + _s3 = __lsx_vmaddwod_h_b(_s3, _pA3, _pB); + _sum0 = __lasx_xvadd_w(_sum0, __lasx_vext2xv_w_h(__lasx_cast_128(_s0))); + _sum1 = __lasx_xvadd_w(_sum1, __lasx_vext2xv_w_h(__lasx_cast_128(_s1))); + _sum2 = __lasx_xvadd_w(_sum2, __lasx_vext2xv_w_h(__lasx_cast_128(_s2))); + _sum3 = __lasx_xvadd_w(_sum3, __lasx_vext2xv_w_h(__lasx_cast_128(_s3))); + pB += 16; + pA += 8; + kk += 2; + } + if (kk < max_kk) + { + __m128i _pB = __lsx_vldrepl_d(pB, 0); + _pB = __lsx_vilvl_b(__lsx_vslti_b(_pB, 0), _pB); + const int a0123 = __lsx_vpickve2gr_w(__lsx_vldrepl_w(pA, 0), 0); + __m128i _s0 = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)a0123), _pB); + __m128i _s1 = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)(a0123 >> 8)), _pB); + __m128i _s2 = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)(a0123 >> 16)), _pB); + __m128i _s3 = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)(a0123 >> 24)), _pB); + _sum0 = __lasx_xvadd_w(_sum0, __lasx_vext2xv_w_h(__lasx_cast_128(_s0))); + _sum1 = __lasx_xvadd_w(_sum1, __lasx_vext2xv_w_h(__lasx_cast_128(_s1))); + _sum2 = __lasx_xvadd_w(_sum2, __lasx_vext2xv_w_h(__lasx_cast_128(_s2))); + _sum3 = __lasx_xvadd_w(_sum3, __lasx_vext2xv_w_h(__lasx_cast_128(_s3))); + pB += 8; + pA += 4; + } + + __m256 _bscale = (__m256)__lasx_xvld(pB_descales, 0); + __m256 _scale0 = __lasx_xvfmul_s(_bscale, (__m256)__lasx_xvreplfr2vr_s(pA_descales[0])); + __m256 _scale1 = __lasx_xvfmul_s(_bscale, (__m256)__lasx_xvreplfr2vr_s(pA_descales[1])); + __m256 _scale2 = __lasx_xvfmul_s(_bscale, (__m256)__lasx_xvreplfr2vr_s(pA_descales[2])); + __m256 _scale3 = __lasx_xvfmul_s(_bscale, (__m256)__lasx_xvreplfr2vr_s(pA_descales[3])); + _out0 = __lasx_xvfmadd_s((__m256)__lasx_xvffint_s_w(_sum0), _scale0, _out0); + _out1 = __lasx_xvfmadd_s((__m256)__lasx_xvffint_s_w(_sum1), _scale1, _out1); + _out2 = __lasx_xvfmadd_s((__m256)__lasx_xvffint_s_w(_sum2), _scale2, _out2); + _out3 = __lasx_xvfmadd_s((__m256)__lasx_xvffint_s_w(_sum3), _scale3, _out3); + pA_descales += 4; + pB_descales += 8; + } + + __lasx_xvst(_out0, outptr + (ii + 0) * max_jj + jj, 0); + __lasx_xvst(_out1, outptr + (ii + 1) * max_jj + jj, 0); + __lasx_xvst(_out2, outptr + (ii + 2) * max_jj + jj, 0); + __lasx_xvst(_out3, outptr + (ii + 3) * max_jj + jj, 0); + } +#endif +#if __loongarch_sx + for (; jj + 7 < max_jj; jj += 8) + { + const signed char* pB0 = pB; + const signed char* pB1 = pB + 4 * K; + const float* pB_descales0 = pB_descales; + const float* pB_descales1 = pB_descales + 4 * ((K + block_size - 1) / block_size); + __m128 _out00 = (__m128)__lsx_vldi(0); + __m128 _out01 = (__m128)__lsx_vldi(0); + __m128 _out10 = (__m128)__lsx_vldi(0); + __m128 _out11 = (__m128)__lsx_vldi(0); + __m128 _out20 = (__m128)__lsx_vldi(0); + __m128 _out21 = (__m128)__lsx_vldi(0); + __m128 _out30 = (__m128)__lsx_vldi(0); + __m128 _out31 = (__m128)__lsx_vldi(0); + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + __m128i _sum00 = __lsx_vreplgr2vr_w(0); + __m128i _sum01 = __lsx_vreplgr2vr_w(0); + __m128i _sum10 = __lsx_vreplgr2vr_w(0); + __m128i _sum11 = __lsx_vreplgr2vr_w(0); + __m128i _sum20 = __lsx_vreplgr2vr_w(0); + __m128i _sum21 = __lsx_vreplgr2vr_w(0); + __m128i _sum30 = __lsx_vreplgr2vr_w(0); + __m128i _sum31 = __lsx_vreplgr2vr_w(0); + const int max_kk = std::min(K - k, block_size); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __m128i _pB0 = __lsx_vld(pB0, 0); + __m128i _pB1 = __lsx_vld(pB1, 0); + __m128i _pA = __lsx_vld(pA, 0); + __m128i _pA0 = __lsx_vreplvei_w(_pA, 0); + __m128i _pA1 = __lsx_vreplvei_w(_pA, 1); + __m128i _pA2 = __lsx_vreplvei_w(_pA, 2); + __m128i _pA3 = __lsx_vreplvei_w(_pA, 3); + __m128i _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB0), _pA0, _pB0); + _sum00 = __lsx_vadd_w(_sum00, __lsx_vhaddw_w_h(_s, _s)); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB1), _pA0, _pB1); + _sum01 = __lsx_vadd_w(_sum01, __lsx_vhaddw_w_h(_s, _s)); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA1, _pB0), _pA1, _pB0); + _sum10 = __lsx_vadd_w(_sum10, __lsx_vhaddw_w_h(_s, _s)); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA1, _pB1), _pA1, _pB1); + _sum11 = __lsx_vadd_w(_sum11, __lsx_vhaddw_w_h(_s, _s)); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA2, _pB0), _pA2, _pB0); + _sum20 = __lsx_vadd_w(_sum20, __lsx_vhaddw_w_h(_s, _s)); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA2, _pB1), _pA2, _pB1); + _sum21 = __lsx_vadd_w(_sum21, __lsx_vhaddw_w_h(_s, _s)); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA3, _pB0), _pA3, _pB0); + _sum30 = __lsx_vadd_w(_sum30, __lsx_vhaddw_w_h(_s, _s)); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA3, _pB1), _pA3, _pB1); + _sum31 = __lsx_vadd_w(_sum31, __lsx_vhaddw_w_h(_s, _s)); + pB0 += 16; + pB1 += 16; + pA += 16; + } + if (kk + 1 < max_kk) + { + __m128i _pB0 = __lsx_vldrepl_d(pB0, 0); + __m128i _pB1 = __lsx_vldrepl_d(pB1, 0); + __m128i _pA = __lsx_vldrepl_d(pA, 0); + __m128i _pA0 = __lsx_vreplvei_h(_pA, 0); + __m128i _pA1 = __lsx_vreplvei_h(_pA, 1); + __m128i _pA2 = __lsx_vreplvei_h(_pA, 2); + __m128i _pA3 = __lsx_vreplvei_h(_pA, 3); + __m128i _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB0), _pA0, _pB0); + _sum00 = __lsx_vadd_w(_sum00, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB1), _pA0, _pB1); + _sum01 = __lsx_vadd_w(_sum01, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA1, _pB0), _pA1, _pB0); + _sum10 = __lsx_vadd_w(_sum10, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA1, _pB1), _pA1, _pB1); + _sum11 = __lsx_vadd_w(_sum11, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA2, _pB0), _pA2, _pB0); + _sum20 = __lsx_vadd_w(_sum20, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA2, _pB1), _pA2, _pB1); + _sum21 = __lsx_vadd_w(_sum21, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA3, _pB0), _pA3, _pB0); + _sum30 = __lsx_vadd_w(_sum30, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA3, _pB1), _pA3, _pB1); + _sum31 = __lsx_vadd_w(_sum31, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); + pB0 += 8; + pB1 += 8; + pA += 8; + kk += 2; + } + if (kk < max_kk) + { + __m128i _pB0 = __lsx_vldrepl_w(pB0, 0); + __m128i _pB1 = __lsx_vldrepl_w(pB1, 0); + _pB0 = __lsx_vilvl_b(__lsx_vslti_b(_pB0, 0), _pB0); + _pB1 = __lsx_vilvl_b(__lsx_vslti_b(_pB1, 0), _pB1); + const int a0123 = __lsx_vpickve2gr_w(__lsx_vldrepl_w(pA, 0), 0); + __m128i _s = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)a0123), _pB0); + _sum00 = __lsx_vadd_w(_sum00, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); + _s = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)a0123), _pB1); + _sum01 = __lsx_vadd_w(_sum01, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); + _s = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)(a0123 >> 8)), _pB0); + _sum10 = __lsx_vadd_w(_sum10, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); + _s = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)(a0123 >> 8)), _pB1); + _sum11 = __lsx_vadd_w(_sum11, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); + _s = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)(a0123 >> 16)), _pB0); + _sum20 = __lsx_vadd_w(_sum20, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); + _s = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)(a0123 >> 16)), _pB1); + _sum21 = __lsx_vadd_w(_sum21, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); + _s = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)(a0123 >> 24)), _pB0); + _sum30 = __lsx_vadd_w(_sum30, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); + _s = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)(a0123 >> 24)), _pB1); + _sum31 = __lsx_vadd_w(_sum31, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); + pB0 += 4; + pB1 += 4; + pA += 4; + } + __m128 _bscale0 = (__m128)__lsx_vld(pB_descales0, 0); + __m128 _bscale1 = (__m128)__lsx_vld(pB_descales1, 0); + __m128 _ascale = __lsx_vreplfr2vr_s(pA_descales[0]); + _out00 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum00), __lsx_vfmul_s(_bscale0, _ascale), _out00); + _out01 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum01), __lsx_vfmul_s(_bscale1, _ascale), _out01); + _ascale = __lsx_vreplfr2vr_s(pA_descales[1]); + _out10 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum10), __lsx_vfmul_s(_bscale0, _ascale), _out10); + _out11 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum11), __lsx_vfmul_s(_bscale1, _ascale), _out11); + _ascale = __lsx_vreplfr2vr_s(pA_descales[2]); + _out20 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum20), __lsx_vfmul_s(_bscale0, _ascale), _out20); + _out21 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum21), __lsx_vfmul_s(_bscale1, _ascale), _out21); + _ascale = __lsx_vreplfr2vr_s(pA_descales[3]); + _out30 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum30), __lsx_vfmul_s(_bscale0, _ascale), _out30); + _out31 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum31), __lsx_vfmul_s(_bscale1, _ascale), _out31); + pA_descales += 4; + pB_descales0 += 4; + pB_descales1 += 4; + } + pB = pB1; + pB_descales = pB_descales1; + __lsx_vst((__m128i)_out00, outptr + (ii + 0) * max_jj + jj, 0); + __lsx_vst((__m128i)_out01, outptr + (ii + 0) * max_jj + jj + 4, 0); + __lsx_vst((__m128i)_out10, outptr + (ii + 1) * max_jj + jj, 0); + __lsx_vst((__m128i)_out11, outptr + (ii + 1) * max_jj + jj + 4, 0); + __lsx_vst((__m128i)_out20, outptr + (ii + 2) * max_jj + jj, 0); + __lsx_vst((__m128i)_out21, outptr + (ii + 2) * max_jj + jj + 4, 0); + __lsx_vst((__m128i)_out30, outptr + (ii + 3) * max_jj + jj, 0); + __lsx_vst((__m128i)_out31, outptr + (ii + 3) * max_jj + jj + 4, 0); + } + for (; jj + 3 < max_jj; jj += 4) + { + __m128 _out0 = (__m128)__lsx_vldi(0); + __m128 _out1 = (__m128)__lsx_vldi(0); + __m128 _out2 = (__m128)__lsx_vldi(0); + __m128 _out3 = (__m128)__lsx_vldi(0); + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + __m128i _sum0 = __lsx_vreplgr2vr_w(0); + __m128i _sum1 = __lsx_vreplgr2vr_w(0); + __m128i _sum2 = __lsx_vreplgr2vr_w(0); + __m128i _sum3 = __lsx_vreplgr2vr_w(0); + const int max_kk = std::min(K - k, block_size); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __m128i _pB = __lsx_vld(pB, 0); + __m128i _pA = __lsx_vld(pA, 0); + __m128i _pA0 = __lsx_vreplvei_w(_pA, 0); + __m128i _pA1 = __lsx_vreplvei_w(_pA, 1); + __m128i _pA2 = __lsx_vreplvei_w(_pA, 2); + __m128i _pA3 = __lsx_vreplvei_w(_pA, 3); + __m128i _s0 = __lsx_vmulwev_h_b(_pA0, _pB); + __m128i _s1 = __lsx_vmulwev_h_b(_pA1, _pB); + __m128i _s2 = __lsx_vmulwev_h_b(_pA2, _pB); + __m128i _s3 = __lsx_vmulwev_h_b(_pA3, _pB); + _s0 = __lsx_vmaddwod_h_b(_s0, _pA0, _pB); + _s1 = __lsx_vmaddwod_h_b(_s1, _pA1, _pB); + _s2 = __lsx_vmaddwod_h_b(_s2, _pA2, _pB); + _s3 = __lsx_vmaddwod_h_b(_s3, _pA3, _pB); + _sum0 = __lsx_vadd_w(_sum0, __lsx_vhaddw_w_h(_s0, _s0)); + _sum1 = __lsx_vadd_w(_sum1, __lsx_vhaddw_w_h(_s1, _s1)); + _sum2 = __lsx_vadd_w(_sum2, __lsx_vhaddw_w_h(_s2, _s2)); + _sum3 = __lsx_vadd_w(_sum3, __lsx_vhaddw_w_h(_s3, _s3)); + pB += 16; + pA += 16; + } + if (kk + 1 < max_kk) + { + __m128i _pB = __lsx_vldrepl_d(pB, 0); + __m128i _pA = __lsx_vldrepl_d(pA, 0); + __m128i _pA0 = __lsx_vreplvei_h(_pA, 0); + __m128i _pA1 = __lsx_vreplvei_h(_pA, 1); + __m128i _pA2 = __lsx_vreplvei_h(_pA, 2); + __m128i _pA3 = __lsx_vreplvei_h(_pA, 3); + __m128i _s0 = __lsx_vmulwev_h_b(_pA0, _pB); + __m128i _s1 = __lsx_vmulwev_h_b(_pA1, _pB); + __m128i _s2 = __lsx_vmulwev_h_b(_pA2, _pB); + __m128i _s3 = __lsx_vmulwev_h_b(_pA3, _pB); + _s0 = __lsx_vmaddwod_h_b(_s0, _pA0, _pB); + _s1 = __lsx_vmaddwod_h_b(_s1, _pA1, _pB); + _s2 = __lsx_vmaddwod_h_b(_s2, _pA2, _pB); + _s3 = __lsx_vmaddwod_h_b(_s3, _pA3, _pB); + _sum0 = __lsx_vadd_w(_sum0, __lsx_vilvl_h(__lsx_vslti_h(_s0, 0), _s0)); + _sum1 = __lsx_vadd_w(_sum1, __lsx_vilvl_h(__lsx_vslti_h(_s1, 0), _s1)); + _sum2 = __lsx_vadd_w(_sum2, __lsx_vilvl_h(__lsx_vslti_h(_s2, 0), _s2)); + _sum3 = __lsx_vadd_w(_sum3, __lsx_vilvl_h(__lsx_vslti_h(_s3, 0), _s3)); + pB += 8; + pA += 8; + kk += 2; + } + if (kk < max_kk) + { + __m128i _pB = __lsx_vldrepl_w(pB, 0); + _pB = __lsx_vilvl_b(__lsx_vslti_b(_pB, 0), _pB); + const int a0123 = __lsx_vpickve2gr_w(__lsx_vldrepl_w(pA, 0), 0); + __m128i _s0 = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)a0123), _pB); + __m128i _s1 = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)(a0123 >> 8)), _pB); + __m128i _s2 = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)(a0123 >> 16)), _pB); + __m128i _s3 = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)(a0123 >> 24)), _pB); + _sum0 = __lsx_vadd_w(_sum0, __lsx_vilvl_h(__lsx_vslti_h(_s0, 0), _s0)); + _sum1 = __lsx_vadd_w(_sum1, __lsx_vilvl_h(__lsx_vslti_h(_s1, 0), _s1)); + _sum2 = __lsx_vadd_w(_sum2, __lsx_vilvl_h(__lsx_vslti_h(_s2, 0), _s2)); + _sum3 = __lsx_vadd_w(_sum3, __lsx_vilvl_h(__lsx_vslti_h(_s3, 0), _s3)); + pB += 4; + pA += 4; + } + __m128 _bscale = (__m128)__lsx_vld(pB_descales, 0); + _out0 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum0), __lsx_vfmul_s(_bscale, __lsx_vreplfr2vr_s(pA_descales[0])), _out0); + _out1 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum1), __lsx_vfmul_s(_bscale, __lsx_vreplfr2vr_s(pA_descales[1])), _out1); + _out2 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum2), __lsx_vfmul_s(_bscale, __lsx_vreplfr2vr_s(pA_descales[2])), _out2); + _out3 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum3), __lsx_vfmul_s(_bscale, __lsx_vreplfr2vr_s(pA_descales[3])), _out3); + pA_descales += 4; + pB_descales += 4; + } + __lsx_vst((__m128i)_out0, outptr + (ii + 0) * max_jj + jj, 0); + __lsx_vst((__m128i)_out1, outptr + (ii + 1) * max_jj + jj, 0); + __lsx_vst((__m128i)_out2, outptr + (ii + 2) * max_jj + jj, 0); + __lsx_vst((__m128i)_out3, outptr + (ii + 3) * max_jj + jj, 0); + } +#endif +#if __loongarch_sx + for (; jj + 1 < max_jj; jj += 2) + { + __m128 _out0 = (__m128)__lsx_vldi(0); + __m128 _out1 = (__m128)__lsx_vldi(0); + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + __m128i _sum0 = __lsx_vreplgr2vr_w(0); + __m128i _sum1 = __lsx_vreplgr2vr_w(0); + const int max_kk = std::min(K - k, block_size); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __m128i _pA = __lsx_vld(pA, 0); + __m128i _pB0 = __lsx_vldrepl_d(pB, 0); + __m128i _pB1 = __lsx_vshuf4i_w(_pB0, _LSX_SHUFFLE(2, 3, 0, 1)); + __m128i _s0 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA, _pB0), _pA, _pB0); + __m128i _s1 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA, _pB1), _pA, _pB1); + _sum0 = __lsx_vadd_w(_sum0, __lsx_vhaddw_w_h(_s0, _s0)); + _sum1 = __lsx_vadd_w(_sum1, __lsx_vhaddw_w_h(_s1, _s1)); + pA += 16; + pB += 8; + } + if (kk + 1 < max_kk) + { + __m128i _pA = __lsx_vldrepl_d(pA, 0); + __m128i _pA0 = __lsx_vreplvei_w(__lsx_vpickev_b(_pA, _pA), 0); + __m128i _pA1 = __lsx_vreplvei_w(__lsx_vpickod_b(_pA, _pA), 0); + _pA0 = __lsx_vilvl_b(__lsx_vslti_b(_pA0, 0), _pA0); + _pA1 = __lsx_vilvl_b(__lsx_vslti_b(_pA1, 0), _pA1); + int b01 = (unsigned char)pB[0] | ((unsigned char)pB[2] << 8); + __m128i _pB0 = __lsx_vreplgr2vr_w(b01); + _pB0 = __lsx_vilvl_b(__lsx_vslti_b(_pB0, 0), _pB0); + _pB0 = __lsx_vshuf4i_h(_pB0, _LSX_SHUFFLE(1, 0, 1, 0)); + __m128i _pB1 = __lsx_vshuf4i_h(_pB0, _LSX_SHUFFLE(0, 3, 2, 1)); + __m128i _s0 = __lsx_vmul_h(_pA0, _pB0); + __m128i _s1 = __lsx_vmul_h(_pA0, _pB1); + _sum0 = __lsx_vadd_w(_sum0, __lsx_vilvl_h(__lsx_vslti_h(_s0, 0), _s0)); + _sum1 = __lsx_vadd_w(_sum1, __lsx_vilvl_h(__lsx_vslti_h(_s1, 0), _s1)); + b01 = (unsigned char)pB[1] | ((unsigned char)pB[3] << 8); + _pB0 = __lsx_vreplgr2vr_w(b01); + _pB0 = __lsx_vilvl_b(__lsx_vslti_b(_pB0, 0), _pB0); + _pB0 = __lsx_vshuf4i_h(_pB0, _LSX_SHUFFLE(1, 0, 1, 0)); + _pB1 = __lsx_vshuf4i_h(_pB0, _LSX_SHUFFLE(0, 3, 2, 1)); + _s0 = __lsx_vmul_h(_pA1, _pB0); + _s1 = __lsx_vmul_h(_pA1, _pB1); + _sum0 = __lsx_vadd_w(_sum0, __lsx_vilvl_h(__lsx_vslti_h(_s0, 0), _s0)); + _sum1 = __lsx_vadd_w(_sum1, __lsx_vilvl_h(__lsx_vslti_h(_s1, 0), _s1)); + pA += 8; + pB += 4; + kk += 2; + } + if (kk < max_kk) + { + __m128i _pA = __lsx_vldrepl_w(pA, 0); + _pA = __lsx_vilvl_b(__lsx_vslti_b(_pA, 0), _pA); + int b01 = (unsigned char)pB[0] | ((unsigned char)pB[1] << 8); + __m128i _pB0 = __lsx_vreplgr2vr_w(b01); + _pB0 = __lsx_vilvl_b(__lsx_vslti_b(_pB0, 0), _pB0); + _pB0 = __lsx_vshuf4i_h(_pB0, _LSX_SHUFFLE(1, 0, 1, 0)); + __m128i _pB1 = __lsx_vshuf4i_h(_pB0, _LSX_SHUFFLE(0, 3, 2, 1)); + __m128i _s0 = __lsx_vmul_h(_pA, _pB0); + __m128i _s1 = __lsx_vmul_h(_pA, _pB1); + _sum0 = __lsx_vadd_w(_sum0, __lsx_vilvl_h(__lsx_vslti_h(_s0, 0), _s0)); + _sum1 = __lsx_vadd_w(_sum1, __lsx_vilvl_h(__lsx_vslti_h(_s1, 0), _s1)); + pA += 4; + pB += 2; + } + __m128i _sum0e = __lsx_vshuf4i_w(_sum0, _LSX_SHUFFLE(3, 1, 2, 0)); + __m128i _sum0o = __lsx_vshuf4i_w(_sum0, _LSX_SHUFFLE(2, 0, 3, 1)); + __m128i _sum1e = __lsx_vshuf4i_w(_sum1, _LSX_SHUFFLE(3, 1, 2, 0)); + __m128i _sum1o = __lsx_vshuf4i_w(_sum1, _LSX_SHUFFLE(2, 0, 3, 1)); + __m128i _sumc0 = __lsx_vilvl_w(_sum1o, _sum0e); + __m128i _sumc1 = __lsx_vilvl_w(_sum0o, _sum1e); + const __m128 _ascale = (__m128)__lsx_vld(pA_descales, 0); + _out0 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sumc0), __lsx_vfmul_s(_ascale, __lsx_vreplfr2vr_s(pB_descales[0])), _out0); + _out1 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sumc1), __lsx_vfmul_s(_ascale, __lsx_vreplfr2vr_s(pB_descales[1])), _out1); + pA_descales += 4; + pB_descales += 2; + } + __lsx_vstelm_w((__m128i)_out0, outptr + (ii + 0) * max_jj + jj, 0, 0); + __lsx_vstelm_w((__m128i)_out1, outptr + (ii + 0) * max_jj + jj + 1, 0, 0); + __lsx_vstelm_w((__m128i)_out0, outptr + (ii + 1) * max_jj + jj, 0, 1); + __lsx_vstelm_w((__m128i)_out1, outptr + (ii + 1) * max_jj + jj + 1, 0, 1); + __lsx_vstelm_w((__m128i)_out0, outptr + (ii + 2) * max_jj + jj, 0, 2); + __lsx_vstelm_w((__m128i)_out1, outptr + (ii + 2) * max_jj + jj + 1, 0, 2); + __lsx_vstelm_w((__m128i)_out0, outptr + (ii + 3) * max_jj + jj, 0, 3); + __lsx_vstelm_w((__m128i)_out1, outptr + (ii + 3) * max_jj + jj + 1, 0, 3); + } + for (; jj < max_jj; jj++) + { + __m128 _out0 = (__m128)__lsx_vldi(0); + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + __m128i _sum0 = __lsx_vreplgr2vr_w(0); + const int max_kk = std::min(K - k, block_size); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __m128i _pA = __lsx_vld(pA, 0); + __m128i _pB = __lsx_vldrepl_w(pB, 0); + __m128i _s0 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA, _pB), _pA, _pB); + _sum0 = __lsx_vadd_w(_sum0, __lsx_vhaddw_w_h(_s0, _s0)); + pA += 16; + pB += 4; + } + if (kk + 1 < max_kk) + { + __m128i _pA = __lsx_vldrepl_d(pA, 0); + __m128i _pB = __lsx_vreplgr2vr_h((unsigned char)pB[0] | ((unsigned char)pB[1] << 8)); + __m128i _s0 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA, _pB), _pA, _pB); + _sum0 = __lsx_vadd_w(_sum0, __lsx_vilvl_h(__lsx_vslti_h(_s0, 0), _s0)); + pA += 8; + pB += 2; + kk += 2; + } + if (kk < max_kk) + { + __m128i _pA = __lsx_vldrepl_w(pA, 0); + _pA = __lsx_vilvl_b(__lsx_vslti_b(_pA, 0), _pA); + __m128i _s0 = __lsx_vmul_h(_pA, __lsx_vreplgr2vr_h(pB[0])); + _sum0 = __lsx_vadd_w(_sum0, __lsx_vilvl_h(__lsx_vslti_h(_s0, 0), _s0)); + pA += 4; + pB++; + } + const __m128 _scale = __lsx_vfmul_s((__m128)__lsx_vld(pA_descales, 0), __lsx_vreplfr2vr_s(*pB_descales++)); + _out0 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum0), _scale, _out0); + pA_descales += 4; + } + __lsx_vstelm_w((__m128i)_out0, outptr + (ii + 0) * max_jj + jj, 0, 0); + __lsx_vstelm_w((__m128i)_out0, outptr + (ii + 1) * max_jj + jj, 0, 1); + __lsx_vstelm_w((__m128i)_out0, outptr + (ii + 2) * max_jj + jj, 0, 2); + __lsx_vstelm_w((__m128i)_out0, outptr + (ii + 3) * max_jj + jj, 0, 3); + } +#endif + pAT += A_hstep * 4; + pAT_descales += A_descales_hstep * 4; + } + for (; ii + 1 < max_ii; ii += 2) + { + const signed char* pB = pBT; + const float* pB_descales = pBT_descales; + + int jj = 0; +#if __loongarch_asx + for (; jj + 15 < max_jj; jj += 16) + { + const signed char* pB0 = pB; + const signed char* pB1 = pB + 8 * K; + const float* pB_descales0 = pB_descales; + const float* pB_descales1 = pB_descales + 8 * ((K + block_size - 1) / block_size); + __m256 _out00 = (__m256)__lasx_xvldi(0); + __m256 _out01 = (__m256)__lasx_xvldi(0); + __m256 _out10 = (__m256)__lasx_xvldi(0); + __m256 _out11 = (__m256)__lasx_xvldi(0); + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + __m256i _sum00 = __lasx_xvreplgr2vr_w(0); + __m256i _sum01 = __lasx_xvreplgr2vr_w(0); + __m256i _sum10 = __lasx_xvreplgr2vr_w(0); + __m256i _sum11 = __lasx_xvreplgr2vr_w(0); + const int max_kk = std::min(K - k, block_size); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __m256i _pB0 = __lasx_xvld(pB0, 0); + __m256i _pB1 = __lasx_xvld(pB1, 0); + __m128i _pAs = __lsx_vldrepl_d(pA, 0); + __m256i _pA0 = __lasx_xvreplgr2vr_w(__lsx_vpickve2gr_w(_pAs, 0)); + __m256i _pA1 = __lasx_xvreplgr2vr_w(__lsx_vpickve2gr_w(_pAs, 1)); + __m256i _s = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA0, _pB0), _pA0, _pB0); + _sum00 = __lasx_xvadd_w(_sum00, __lasx_xvhaddw_w_h(_s, _s)); + _s = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA0, _pB1), _pA0, _pB1); + _sum01 = __lasx_xvadd_w(_sum01, __lasx_xvhaddw_w_h(_s, _s)); + _s = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA1, _pB0), _pA1, _pB0); + _sum10 = __lasx_xvadd_w(_sum10, __lasx_xvhaddw_w_h(_s, _s)); + _s = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA1, _pB1), _pA1, _pB1); + _sum11 = __lasx_xvadd_w(_sum11, __lasx_xvhaddw_w_h(_s, _s)); + pB0 += 32; + pB1 += 32; + pA += 8; + } + if (kk + 1 < max_kk) + { + __m128i _pB0 = __lsx_vld(pB0, 0); + __m128i _pB1 = __lsx_vld(pB1, 0); + __m128i _pAs = __lsx_vldrepl_w(pA, 0); + __m128i _pA0 = __lsx_vreplvei_h(_pAs, 0); + __m128i _pA1 = __lsx_vreplvei_h(_pAs, 1); + __m128i _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB0), _pA0, _pB0); + _sum00 = __lasx_xvadd_w(_sum00, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB1), _pA0, _pB1); + _sum01 = __lasx_xvadd_w(_sum01, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA1, _pB0), _pA1, _pB0); + _sum10 = __lasx_xvadd_w(_sum10, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA1, _pB1), _pA1, _pB1); + _sum11 = __lasx_xvadd_w(_sum11, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); + pB0 += 16; + pB1 += 16; + pA += 4; + kk += 2; + } + if (kk < max_kk) + { + __m128i _pB0 = __lsx_vldrepl_d(pB0, 0); + __m128i _pB1 = __lsx_vldrepl_d(pB1, 0); + _pB0 = __lsx_vilvl_b(__lsx_vslti_b(_pB0, 0), _pB0); + _pB1 = __lsx_vilvl_b(__lsx_vslti_b(_pB1, 0), _pB1); + __m128i _pAs = __lsx_vldrepl_h(pA, 0); + const int a0 = (signed char)__lsx_vpickve2gr_b(_pAs, 0); + const int a1 = (signed char)__lsx_vpickve2gr_b(_pAs, 1); + __m128i _s = __lsx_vmul_h(__lsx_vreplgr2vr_h(a0), _pB0); + _sum00 = __lasx_xvadd_w(_sum00, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); + _s = __lsx_vmul_h(__lsx_vreplgr2vr_h(a0), _pB1); + _sum01 = __lasx_xvadd_w(_sum01, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); + _s = __lsx_vmul_h(__lsx_vreplgr2vr_h(a1), _pB0); + _sum10 = __lasx_xvadd_w(_sum10, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); + _s = __lsx_vmul_h(__lsx_vreplgr2vr_h(a1), _pB1); + _sum11 = __lasx_xvadd_w(_sum11, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); + pB0 += 8; + pB1 += 8; + pA += 2; + } + __m256 _bscale0 = (__m256)__lasx_xvld(pB_descales0, 0); + __m256 _bscale1 = (__m256)__lasx_xvld(pB_descales1, 0); + __m256 _ascale = (__m256)__lasx_xvreplfr2vr_s(pA_descales[0]); + _out00 = __lasx_xvfmadd_s((__m256)__lasx_xvffint_s_w(_sum00), __lasx_xvfmul_s(_bscale0, _ascale), _out00); + _out01 = __lasx_xvfmadd_s((__m256)__lasx_xvffint_s_w(_sum01), __lasx_xvfmul_s(_bscale1, _ascale), _out01); + _ascale = (__m256)__lasx_xvreplfr2vr_s(pA_descales[1]); + _out10 = __lasx_xvfmadd_s((__m256)__lasx_xvffint_s_w(_sum10), __lasx_xvfmul_s(_bscale0, _ascale), _out10); + _out11 = __lasx_xvfmadd_s((__m256)__lasx_xvffint_s_w(_sum11), __lasx_xvfmul_s(_bscale1, _ascale), _out11); + pA_descales += 2; + pB_descales0 += 8; + pB_descales1 += 8; + } + pB = pB1; + pB_descales = pB_descales1; + __lasx_xvst(_out00, outptr + (ii + 0) * max_jj + jj, 0); + __lasx_xvst(_out01, outptr + (ii + 0) * max_jj + jj + 8, 0); + __lasx_xvst(_out10, outptr + (ii + 1) * max_jj + jj, 0); + __lasx_xvst(_out11, outptr + (ii + 1) * max_jj + jj + 8, 0); + } + for (; jj + 7 < max_jj; jj += 8) + { + __m256 _out0 = (__m256)__lasx_xvldi(0); + __m256 _out1 = (__m256)__lasx_xvldi(0); + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + __m256i _sum0 = __lasx_xvreplgr2vr_w(0); + __m256i _sum1 = __lasx_xvreplgr2vr_w(0); + const int max_kk = std::min(K - k, block_size); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __m256i _pB = __lasx_xvld(pB, 0); + __m128i _pAs = __lsx_vldrepl_d(pA, 0); + __m256i _pA0 = __lasx_xvreplgr2vr_w(__lsx_vpickve2gr_w(_pAs, 0)); + __m256i _pA1 = __lasx_xvreplgr2vr_w(__lsx_vpickve2gr_w(_pAs, 1)); + __m256i _s0 = __lasx_xvmulwev_h_b(_pA0, _pB); + __m256i _s1 = __lasx_xvmulwev_h_b(_pA1, _pB); + _s0 = __lasx_xvmaddwod_h_b(_s0, _pA0, _pB); + _s1 = __lasx_xvmaddwod_h_b(_s1, _pA1, _pB); + _sum0 = __lasx_xvadd_w(_sum0, __lasx_xvhaddw_w_h(_s0, _s0)); + _sum1 = __lasx_xvadd_w(_sum1, __lasx_xvhaddw_w_h(_s1, _s1)); + pB += 32; + pA += 8; + } + if (kk + 1 < max_kk) + { + __m128i _pB = __lsx_vld(pB, 0); + __m128i _pAs = __lsx_vldrepl_w(pA, 0); + __m128i _pA0 = __lsx_vreplvei_h(_pAs, 0); + __m128i _pA1 = __lsx_vreplvei_h(_pAs, 1); + __m128i _s0 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB), _pA0, _pB); + __m128i _s1 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA1, _pB), _pA1, _pB); + _sum0 = __lasx_xvadd_w(_sum0, __lasx_vext2xv_w_h(__lasx_cast_128(_s0))); + _sum1 = __lasx_xvadd_w(_sum1, __lasx_vext2xv_w_h(__lasx_cast_128(_s1))); + pB += 16; + pA += 4; + kk += 2; + } + if (kk < max_kk) + { + __m128i _pB = __lsx_vldrepl_d(pB, 0); + _pB = __lsx_vilvl_b(__lsx_vslti_b(_pB, 0), _pB); + __m128i _pAs = __lsx_vldrepl_h(pA, 0); + __m128i _s0 = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)__lsx_vpickve2gr_b(_pAs, 0)), _pB); + __m128i _s1 = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)__lsx_vpickve2gr_b(_pAs, 1)), _pB); + _sum0 = __lasx_xvadd_w(_sum0, __lasx_vext2xv_w_h(__lasx_cast_128(_s0))); + _sum1 = __lasx_xvadd_w(_sum1, __lasx_vext2xv_w_h(__lasx_cast_128(_s1))); + pB += 8; + pA += 2; + } + __m256 _bscale = (__m256)__lasx_xvld(pB_descales, 0); + _out0 = __lasx_xvfmadd_s((__m256)__lasx_xvffint_s_w(_sum0), __lasx_xvfmul_s(_bscale, (__m256)__lasx_xvreplfr2vr_s(pA_descales[0])), _out0); + _out1 = __lasx_xvfmadd_s((__m256)__lasx_xvffint_s_w(_sum1), __lasx_xvfmul_s(_bscale, (__m256)__lasx_xvreplfr2vr_s(pA_descales[1])), _out1); + pA_descales += 2; + pB_descales += 8; + } + __lasx_xvst(_out0, outptr + (ii + 0) * max_jj + jj, 0); + __lasx_xvst(_out1, outptr + (ii + 1) * max_jj + jj, 0); + } +#endif +#if __loongarch_sx + for (; jj + 7 < max_jj; jj += 8) + { + const signed char* pB0 = pB; + const signed char* pB1 = pB + 4 * K; + const float* pB_descales0 = pB_descales; + const float* pB_descales1 = pB_descales + 4 * ((K + block_size - 1) / block_size); + __m128 _out00 = (__m128)__lsx_vldi(0); + __m128 _out01 = (__m128)__lsx_vldi(0); + __m128 _out10 = (__m128)__lsx_vldi(0); + __m128 _out11 = (__m128)__lsx_vldi(0); + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + __m128i _sum00 = __lsx_vreplgr2vr_w(0); + __m128i _sum01 = __lsx_vreplgr2vr_w(0); + __m128i _sum10 = __lsx_vreplgr2vr_w(0); + __m128i _sum11 = __lsx_vreplgr2vr_w(0); + const int max_kk = std::min(K - k, block_size); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __m128i _pB0 = __lsx_vld(pB0, 0); + __m128i _pB1 = __lsx_vld(pB1, 0); + __m128i _pAs = __lsx_vldrepl_d(pA, 0); + __m128i _pA0 = __lsx_vreplvei_w(_pAs, 0); + __m128i _pA1 = __lsx_vreplvei_w(_pAs, 1); + __m128i _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB0), _pA0, _pB0); + _sum00 = __lsx_vadd_w(_sum00, __lsx_vhaddw_w_h(_s, _s)); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB1), _pA0, _pB1); + _sum01 = __lsx_vadd_w(_sum01, __lsx_vhaddw_w_h(_s, _s)); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA1, _pB0), _pA1, _pB0); + _sum10 = __lsx_vadd_w(_sum10, __lsx_vhaddw_w_h(_s, _s)); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA1, _pB1), _pA1, _pB1); + _sum11 = __lsx_vadd_w(_sum11, __lsx_vhaddw_w_h(_s, _s)); + pB0 += 16; + pB1 += 16; + pA += 8; + } + if (kk + 1 < max_kk) + { + __m128i _pB0 = __lsx_vldrepl_d(pB0, 0); + __m128i _pB1 = __lsx_vldrepl_d(pB1, 0); + __m128i _pAs = __lsx_vldrepl_w(pA, 0); + __m128i _pA0 = __lsx_vreplvei_h(_pAs, 0); + __m128i _pA1 = __lsx_vreplvei_h(_pAs, 1); + __m128i _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB0), _pA0, _pB0); + _sum00 = __lsx_vadd_w(_sum00, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB1), _pA0, _pB1); + _sum01 = __lsx_vadd_w(_sum01, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA1, _pB0), _pA1, _pB0); + _sum10 = __lsx_vadd_w(_sum10, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA1, _pB1), _pA1, _pB1); + _sum11 = __lsx_vadd_w(_sum11, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); + pB0 += 8; + pB1 += 8; + pA += 4; + kk += 2; + } + if (kk < max_kk) + { + __m128i _pB0 = __lsx_vldrepl_w(pB0, 0); + __m128i _pB1 = __lsx_vldrepl_w(pB1, 0); + _pB0 = __lsx_vilvl_b(__lsx_vslti_b(_pB0, 0), _pB0); + _pB1 = __lsx_vilvl_b(__lsx_vslti_b(_pB1, 0), _pB1); + __m128i _pAs = __lsx_vldrepl_h(pA, 0); + const int a0 = (signed char)__lsx_vpickve2gr_b(_pAs, 0); + const int a1 = (signed char)__lsx_vpickve2gr_b(_pAs, 1); + __m128i _s = __lsx_vmul_h(__lsx_vreplgr2vr_h(a0), _pB0); + _sum00 = __lsx_vadd_w(_sum00, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); + _s = __lsx_vmul_h(__lsx_vreplgr2vr_h(a0), _pB1); + _sum01 = __lsx_vadd_w(_sum01, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); + _s = __lsx_vmul_h(__lsx_vreplgr2vr_h(a1), _pB0); + _sum10 = __lsx_vadd_w(_sum10, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); + _s = __lsx_vmul_h(__lsx_vreplgr2vr_h(a1), _pB1); + _sum11 = __lsx_vadd_w(_sum11, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); + pB0 += 4; + pB1 += 4; + pA += 2; + } + __m128 _bscale0 = (__m128)__lsx_vld(pB_descales0, 0); + __m128 _bscale1 = (__m128)__lsx_vld(pB_descales1, 0); + __m128 _ascale = __lsx_vreplfr2vr_s(pA_descales[0]); + _out00 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum00), __lsx_vfmul_s(_bscale0, _ascale), _out00); + _out01 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum01), __lsx_vfmul_s(_bscale1, _ascale), _out01); + _ascale = __lsx_vreplfr2vr_s(pA_descales[1]); + _out10 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum10), __lsx_vfmul_s(_bscale0, _ascale), _out10); + _out11 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum11), __lsx_vfmul_s(_bscale1, _ascale), _out11); + pA_descales += 2; + pB_descales0 += 4; + pB_descales1 += 4; + } + pB = pB1; + pB_descales = pB_descales1; + __lsx_vst((__m128i)_out00, outptr + (ii + 0) * max_jj + jj, 0); + __lsx_vst((__m128i)_out01, outptr + (ii + 0) * max_jj + jj + 4, 0); + __lsx_vst((__m128i)_out10, outptr + (ii + 1) * max_jj + jj, 0); + __lsx_vst((__m128i)_out11, outptr + (ii + 1) * max_jj + jj + 4, 0); + } + for (; jj + 3 < max_jj; jj += 4) + { + __m128 _out0 = (__m128)__lsx_vldi(0); + __m128 _out1 = (__m128)__lsx_vldi(0); + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + __m128i _sum0 = __lsx_vreplgr2vr_w(0); + __m128i _sum1 = __lsx_vreplgr2vr_w(0); + const int max_kk = std::min(K - k, block_size); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __m128i _pB = __lsx_vld(pB, 0); + __m128i _pAs = __lsx_vldrepl_d(pA, 0); + __m128i _pA0 = __lsx_vreplvei_w(_pAs, 0); + __m128i _pA1 = __lsx_vreplvei_w(_pAs, 1); + __m128i _s0 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB), _pA0, _pB); + __m128i _s1 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA1, _pB), _pA1, _pB); + _sum0 = __lsx_vadd_w(_sum0, __lsx_vhaddw_w_h(_s0, _s0)); + _sum1 = __lsx_vadd_w(_sum1, __lsx_vhaddw_w_h(_s1, _s1)); + pB += 16; + pA += 8; + } + if (kk + 1 < max_kk) + { + __m128i _pB = __lsx_vldrepl_d(pB, 0); + __m128i _pAs = __lsx_vldrepl_w(pA, 0); + __m128i _pA0 = __lsx_vreplvei_h(_pAs, 0); + __m128i _pA1 = __lsx_vreplvei_h(_pAs, 1); + __m128i _s0 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB), _pA0, _pB); + __m128i _s1 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA1, _pB), _pA1, _pB); + _sum0 = __lsx_vadd_w(_sum0, __lsx_vilvl_h(__lsx_vslti_h(_s0, 0), _s0)); + _sum1 = __lsx_vadd_w(_sum1, __lsx_vilvl_h(__lsx_vslti_h(_s1, 0), _s1)); + pB += 8; + pA += 4; + kk += 2; + } + if (kk < max_kk) + { + __m128i _pB = __lsx_vldrepl_w(pB, 0); + _pB = __lsx_vilvl_b(__lsx_vslti_b(_pB, 0), _pB); + __m128i _pAs = __lsx_vldrepl_h(pA, 0); + __m128i _s0 = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)__lsx_vpickve2gr_b(_pAs, 0)), _pB); + __m128i _s1 = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)__lsx_vpickve2gr_b(_pAs, 1)), _pB); + _sum0 = __lsx_vadd_w(_sum0, __lsx_vilvl_h(__lsx_vslti_h(_s0, 0), _s0)); + _sum1 = __lsx_vadd_w(_sum1, __lsx_vilvl_h(__lsx_vslti_h(_s1, 0), _s1)); + pB += 4; + pA += 2; + } + __m128 _bscale = (__m128)__lsx_vld(pB_descales, 0); + _out0 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum0), __lsx_vfmul_s(_bscale, __lsx_vreplfr2vr_s(pA_descales[0])), _out0); + _out1 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum1), __lsx_vfmul_s(_bscale, __lsx_vreplfr2vr_s(pA_descales[1])), _out1); + pA_descales += 2; + pB_descales += 4; + } + __lsx_vst((__m128i)_out0, outptr + (ii + 0) * max_jj + jj, 0); + __lsx_vst((__m128i)_out1, outptr + (ii + 1) * max_jj + jj, 0); + } +#endif + for (; jj + 1 < max_jj; jj += 2) + { + float _out00 = 0.f; + float _out01 = 0.f; + float _out10 = 0.f; + float _out11 = 0.f; + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + int _sum00 = 0; + int _sum01 = 0; + int _sum10 = 0; + int _sum11 = 0; + const int max_kk = std::min(K - k, block_size); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + asm volatile("" + : + : + : "memory"); + + _sum00 += pA[0] * pB[0] + pA[1] * pB[1] + pA[2] * pB[2] + pA[3] * pB[3]; + _sum01 += pA[0] * pB[4] + pA[1] * pB[5] + pA[2] * pB[6] + pA[3] * pB[7]; + _sum10 += pA[4] * pB[0] + pA[5] * pB[1] + pA[6] * pB[2] + pA[7] * pB[3]; + _sum11 += pA[4] * pB[4] + pA[5] * pB[5] + pA[6] * pB[6] + pA[7] * pB[7]; + pB += 8; + pA += 8; + } + if (kk + 1 < max_kk) + { + _sum00 += pA[0] * pB[0] + pA[1] * pB[1]; + _sum01 += pA[0] * pB[2] + pA[1] * pB[3]; + _sum10 += pA[2] * pB[0] + pA[3] * pB[1]; + _sum11 += pA[2] * pB[2] + pA[3] * pB[3]; + pB += 4; + pA += 4; + kk += 2; + } + if (kk < max_kk) + { + _sum00 += pA[0] * pB[0]; + _sum01 += pA[0] * pB[1]; + _sum10 += pA[1] * pB[0]; + _sum11 += pA[1] * pB[1]; + pB += 2; + pA += 2; + } + const float bscale0 = pB_descales[0]; + const float bscale1 = pB_descales[1]; + const float ascale0 = pA_descales[0]; + const float ascale1 = pA_descales[1]; + _out00 += _sum00 * ascale0 * bscale0; + _out01 += _sum01 * ascale0 * bscale1; + _out10 += _sum10 * ascale1 * bscale0; + _out11 += _sum11 * ascale1 * bscale1; + pA_descales += 2; + pB_descales += 2; + } + outptr[(ii + 0) * max_jj + jj] = _out00; + outptr[(ii + 0) * max_jj + jj + 1] = _out01; + outptr[(ii + 1) * max_jj + jj] = _out10; + outptr[(ii + 1) * max_jj + jj + 1] = _out11; + } + for (; jj < max_jj; jj++) + { + float _out0 = 0.f; + float _out1 = 0.f; + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + int _sum0 = 0; + int _sum1 = 0; + const int max_kk = std::min(K - k, block_size); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + asm volatile("" + : + : + : "memory"); + + _sum0 += pA[0] * pB[0] + pA[1] * pB[1] + pA[2] * pB[2] + pA[3] * pB[3]; + _sum1 += pA[4] * pB[0] + pA[5] * pB[1] + pA[6] * pB[2] + pA[7] * pB[3]; + pB += 4; + pA += 8; + } + if (kk + 1 < max_kk) + { + _sum0 += pA[0] * pB[0] + pA[1] * pB[1]; + _sum1 += pA[2] * pB[0] + pA[3] * pB[1]; + pB += 2; + pA += 4; + kk += 2; + } + if (kk < max_kk) + { + _sum0 += pA[0] * pB[0]; + _sum1 += pA[1] * pB[0]; + pB++; + pA += 2; + } + const float bscale = *pB_descales++; + _out0 += _sum0 * pA_descales[0] * bscale; + _out1 += _sum1 * pA_descales[1] * bscale; + pA_descales += 2; + } + outptr[(ii + 0) * max_jj + jj] = _out0; + outptr[(ii + 1) * max_jj + jj] = _out1; + } + pAT += A_hstep * 2; + pAT_descales += A_descales_hstep * 2; + } +#endif // __loongarch_sx + for (; ii < max_ii; ii++) + { + const signed char* pB = pBT; + const float* pB_descales = pBT_descales; + + int jj = 0; +#if __loongarch_asx + for (; jj + 15 < max_jj; jj += 16) + { + const signed char* pB0 = pB; + const signed char* pB1 = pB + 8 * K; + const float* pB_descales0 = pB_descales; + const float* pB_descales1 = pB_descales + 8 * ((K + block_size - 1) / block_size); + __m256 _out00 = (__m256)__lasx_xvldi(0); + __m256 _out01 = (__m256)__lasx_xvldi(0); + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + __m256i _sum00 = __lasx_xvreplgr2vr_w(0); + __m256i _sum01 = __lasx_xvreplgr2vr_w(0); + const int max_kk = std::min(K - k, block_size); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __m256i _pB0 = __lasx_xvld(pB0, 0); + __m256i _pB1 = __lasx_xvld(pB1, 0); + __m256i _pA0 = __lasx_xvldrepl_w(pA, 0); + __m256i _s = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA0, _pB0), _pA0, _pB0); + _sum00 = __lasx_xvadd_w(_sum00, __lasx_xvhaddw_w_h(_s, _s)); + _s = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA0, _pB1), _pA0, _pB1); + _sum01 = __lasx_xvadd_w(_sum01, __lasx_xvhaddw_w_h(_s, _s)); + pB0 += 32; + pB1 += 32; + pA += 4; + } + if (kk + 1 < max_kk) + { + __m128i _pB0 = __lsx_vld(pB0, 0); + __m128i _pB1 = __lsx_vld(pB1, 0); + __m128i _pA0 = __lsx_vreplgr2vr_h((unsigned char)pA[0] | ((unsigned char)pA[1] << 8)); + __m128i _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB0), _pA0, _pB0); + _sum00 = __lasx_xvadd_w(_sum00, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB1), _pA0, _pB1); + _sum01 = __lasx_xvadd_w(_sum01, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); + pB0 += 16; + pB1 += 16; + pA += 2; + kk += 2; + } + if (kk < max_kk) + { + __m128i _pB0 = __lsx_vldrepl_d(pB0, 0); + __m128i _pB1 = __lsx_vldrepl_d(pB1, 0); + _pB0 = __lsx_vilvl_b(__lsx_vslti_b(_pB0, 0), _pB0); + _pB1 = __lsx_vilvl_b(__lsx_vslti_b(_pB1, 0), _pB1); + __m128i _pA0 = __lsx_vreplgr2vr_h(pA[0]); + __m128i _s = __lsx_vmul_h(_pA0, _pB0); + _sum00 = __lasx_xvadd_w(_sum00, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); + _s = __lsx_vmul_h(_pA0, _pB1); + _sum01 = __lasx_xvadd_w(_sum01, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); + pB0 += 8; + pB1 += 8; + pA++; + } + __m256 _bscale0 = (__m256)__lasx_xvld(pB_descales0, 0); + __m256 _bscale1 = (__m256)__lasx_xvld(pB_descales1, 0); + __m256 _ascale = (__m256)__lasx_xvreplfr2vr_s(*pA_descales++); + _out00 = __lasx_xvfmadd_s((__m256)__lasx_xvffint_s_w(_sum00), __lasx_xvfmul_s(_bscale0, _ascale), _out00); + _out01 = __lasx_xvfmadd_s((__m256)__lasx_xvffint_s_w(_sum01), __lasx_xvfmul_s(_bscale1, _ascale), _out01); + pB_descales0 += 8; + pB_descales1 += 8; + } + pB = pB1; + pB_descales = pB_descales1; + __lasx_xvst(_out00, outptr + ii * max_jj + jj, 0); + __lasx_xvst(_out01, outptr + ii * max_jj + jj + 8, 0); + } + for (; jj + 7 < max_jj; jj += 8) + { + __m256 _out0 = (__m256)__lasx_xvldi(0); + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + __m256i _sum0 = __lasx_xvreplgr2vr_w(0); + const int max_kk = std::min(K - k, block_size); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __m256i _pB = __lasx_xvld(pB, 0); + __m256i _pA0 = __lasx_xvldrepl_w(pA, 0); + __m256i _s0 = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA0, _pB), _pA0, _pB); + _sum0 = __lasx_xvadd_w(_sum0, __lasx_xvhaddw_w_h(_s0, _s0)); + pB += 32; + pA += 4; + } + if (kk + 1 < max_kk) + { + __m128i _pB = __lsx_vld(pB, 0); + __m128i _pA0 = __lsx_vreplgr2vr_h((unsigned char)pA[0] | ((unsigned char)pA[1] << 8)); + __m128i _s0 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB), _pA0, _pB); + _sum0 = __lasx_xvadd_w(_sum0, __lasx_vext2xv_w_h(__lasx_cast_128(_s0))); + pB += 16; + pA += 2; + kk += 2; + } + if (kk < max_kk) + { + __m128i _pB = __lsx_vldrepl_d(pB, 0); + _pB = __lsx_vilvl_b(__lsx_vslti_b(_pB, 0), _pB); + __m128i _s0 = __lsx_vmul_h(__lsx_vreplgr2vr_h(pA[0]), _pB); + _sum0 = __lasx_xvadd_w(_sum0, __lasx_vext2xv_w_h(__lasx_cast_128(_s0))); + pB += 8; + pA++; + } + __m256 _bscale = (__m256)__lasx_xvld(pB_descales, 0); + _out0 = __lasx_xvfmadd_s((__m256)__lasx_xvffint_s_w(_sum0), __lasx_xvfmul_s(_bscale, (__m256)__lasx_xvreplfr2vr_s(*pA_descales++)), _out0); + pB_descales += 8; + } + __lasx_xvst(_out0, outptr + ii * max_jj + jj, 0); + } +#endif +#if __loongarch_sx + for (; jj + 7 < max_jj; jj += 8) + { + const signed char* pB0 = pB; + const signed char* pB1 = pB + 4 * K; + const float* pB_descales0 = pB_descales; + const float* pB_descales1 = pB_descales + 4 * ((K + block_size - 1) / block_size); + __m128 _out00 = (__m128)__lsx_vldi(0); + __m128 _out01 = (__m128)__lsx_vldi(0); + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + __m128i _sum00 = __lsx_vreplgr2vr_w(0); + __m128i _sum01 = __lsx_vreplgr2vr_w(0); + const int max_kk = std::min(K - k, block_size); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __m128i _pB0 = __lsx_vld(pB0, 0); + __m128i _pB1 = __lsx_vld(pB1, 0); + __m128i _pA0 = __lsx_vldrepl_w(pA, 0); + __m128i _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB0), _pA0, _pB0); + _sum00 = __lsx_vadd_w(_sum00, __lsx_vhaddw_w_h(_s, _s)); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB1), _pA0, _pB1); + _sum01 = __lsx_vadd_w(_sum01, __lsx_vhaddw_w_h(_s, _s)); + pB0 += 16; + pB1 += 16; + pA += 4; + } + if (kk + 1 < max_kk) + { + __m128i _pB0 = __lsx_vldrepl_d(pB0, 0); + __m128i _pB1 = __lsx_vldrepl_d(pB1, 0); + __m128i _pA0 = __lsx_vreplgr2vr_h((unsigned char)pA[0] | ((unsigned char)pA[1] << 8)); + __m128i _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB0), _pA0, _pB0); + _sum00 = __lsx_vadd_w(_sum00, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB1), _pA0, _pB1); + _sum01 = __lsx_vadd_w(_sum01, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); + pB0 += 8; + pB1 += 8; + pA += 2; + kk += 2; + } + if (kk < max_kk) + { + __m128i _pB0 = __lsx_vldrepl_w(pB0, 0); + __m128i _pB1 = __lsx_vldrepl_w(pB1, 0); + _pB0 = __lsx_vilvl_b(__lsx_vslti_b(_pB0, 0), _pB0); + _pB1 = __lsx_vilvl_b(__lsx_vslti_b(_pB1, 0), _pB1); + __m128i _pA0 = __lsx_vreplgr2vr_h(pA[0]); + __m128i _s = __lsx_vmul_h(_pA0, _pB0); + _sum00 = __lsx_vadd_w(_sum00, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); + _s = __lsx_vmul_h(_pA0, _pB1); + _sum01 = __lsx_vadd_w(_sum01, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); + pB0 += 4; + pB1 += 4; + pA++; + } + __m128 _bscale0 = (__m128)__lsx_vld(pB_descales0, 0); + __m128 _bscale1 = (__m128)__lsx_vld(pB_descales1, 0); + __m128 _ascale = __lsx_vreplfr2vr_s(*pA_descales++); + _out00 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum00), __lsx_vfmul_s(_bscale0, _ascale), _out00); + _out01 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum01), __lsx_vfmul_s(_bscale1, _ascale), _out01); + pB_descales0 += 4; + pB_descales1 += 4; + } + pB = pB1; + pB_descales = pB_descales1; + __lsx_vst((__m128i)_out00, outptr + ii * max_jj + jj, 0); + __lsx_vst((__m128i)_out01, outptr + ii * max_jj + jj + 4, 0); + } + for (; jj + 3 < max_jj; jj += 4) + { + __m128 _out0 = (__m128)__lsx_vldi(0); + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + __m128i _sum0 = __lsx_vreplgr2vr_w(0); + const int max_kk = std::min(K - k, block_size); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __m128i _pB = __lsx_vld(pB, 0); + __m128i _pA0 = __lsx_vldrepl_w(pA, 0); + __m128i _s0 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB), _pA0, _pB); + _sum0 = __lsx_vadd_w(_sum0, __lsx_vhaddw_w_h(_s0, _s0)); + pB += 16; + pA += 4; + } + if (kk + 1 < max_kk) + { + __m128i _pB = __lsx_vldrepl_d(pB, 0); + __m128i _pA0 = __lsx_vreplgr2vr_h((unsigned char)pA[0] | ((unsigned char)pA[1] << 8)); + __m128i _s0 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB), _pA0, _pB); + _sum0 = __lsx_vadd_w(_sum0, __lsx_vilvl_h(__lsx_vslti_h(_s0, 0), _s0)); + pB += 8; + pA += 2; + kk += 2; + } + if (kk < max_kk) + { + __m128i _pB = __lsx_vldrepl_w(pB, 0); + _pB = __lsx_vilvl_b(__lsx_vslti_b(_pB, 0), _pB); + __m128i _s0 = __lsx_vmul_h(__lsx_vreplgr2vr_h(pA[0]), _pB); + _sum0 = __lsx_vadd_w(_sum0, __lsx_vilvl_h(__lsx_vslti_h(_s0, 0), _s0)); + pB += 4; + pA++; + } + __m128 _bscale = (__m128)__lsx_vld(pB_descales, 0); + _out0 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum0), __lsx_vfmul_s(_bscale, __lsx_vreplfr2vr_s(*pA_descales++)), _out0); + pB_descales += 4; + } + __lsx_vst((__m128i)_out0, outptr + ii * max_jj + jj, 0); + } +#endif + for (; jj + 1 < max_jj; jj += 2) + { + float _out0 = 0.f; + float _out1 = 0.f; + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + int _sum0 = 0; + int _sum1 = 0; + const int max_kk = std::min(K - k, block_size); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + asm volatile("" + : + : + : "memory"); + + _sum0 += pA[0] * pB[0] + pA[1] * pB[1] + pA[2] * pB[2] + pA[3] * pB[3]; + _sum1 += pA[0] * pB[4] + pA[1] * pB[5] + pA[2] * pB[6] + pA[3] * pB[7]; + pB += 8; + pA += 4; + } + if (kk + 1 < max_kk) + { + _sum0 += pA[0] * pB[0] + pA[1] * pB[1]; + _sum1 += pA[0] * pB[2] + pA[1] * pB[3]; + pB += 4; + pA += 2; + kk += 2; + } + if (kk < max_kk) + { + _sum0 += pA[0] * pB[0]; + _sum1 += pA[0] * pB[1]; + pB += 2; + pA++; + } + const float ascale = *pA_descales++; + _out0 += _sum0 * ascale * pB_descales[0]; + _out1 += _sum1 * ascale * pB_descales[1]; + pB_descales += 2; + } + outptr[ii * max_jj + jj] = _out0; + outptr[ii * max_jj + jj + 1] = _out1; + } + for (; jj < max_jj; jj++) + { + float _out0 = 0.f; + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + int _sum0 = 0; + const int max_kk = std::min(K - k, block_size); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + _sum0 += pA[0] * pB[0] + pA[1] * pB[1] + pA[2] * pB[2] + pA[3] * pB[3]; + pB += 4; + pA += 4; + } + if (kk + 1 < max_kk) + { + _sum0 += pA[0] * pB[0] + pA[1] * pB[1]; + pB += 2; + pA += 2; + kk += 2; + } + if (kk < max_kk) + { + _sum0 += pA[0] * pB[0]; + pB++; + pA++; + } + _out0 += _sum0 * *pA_descales++ * *pB_descales++; + } + outptr[ii * max_jj + jj] = _out0; + } + pAT += A_hstep; + pAT_descales += A_descales_hstep; + } +} + +static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta) +{ + const float* pp = topT; + const size_t c_hstep = C.dims == 3 ? C.cstep : (size_t)C.w; + const float* pC_base = C; + float* outptr = top_blob; + int ii = 0; +#if __loongarch_sx + for (; ii + 7 < max_ii; ii += 8) + { + float* p0 = outptr + (size_t)(i + ii) * N + j; + float* p1 = p0 + N; + float* p2 = p1 + N; + float* p3 = p2 + N; + float* p4 = p3 + N; + float* p5 = p4 + N; + float* p6 = p5 + N; + float* p7 = p6 + N; + + float c0 = 0.f; + float c1 = 0.f; + float c2 = 0.f; + float c3 = 0.f; + float c4 = 0.f; + float c5 = 0.f; + float c6 = 0.f; + float c7 = 0.f; + const float* pC = pC_base; + if (pC) + { + if (broadcast_type_C == 0) + { + c0 = pC[0] * beta; + c1 = c0; + c2 = c0; + c3 = c0; + c4 = c0; + c5 = c0; + c6 = c0; + c7 = c0; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + c0 = pC[i + ii] * beta; + c1 = pC[i + ii + 1] * beta; + c2 = pC[i + ii + 2] * beta; + c3 = pC[i + ii + 3] * beta; + c4 = pC[i + ii + 4] * beta; + c5 = pC[i + ii + 5] * beta; + c6 = pC[i + ii + 6] * beta; + c7 = pC[i + ii + 7] * beta; + } + if (broadcast_type_C == 3) + { + pC += (size_t)(i + ii) * c_hstep + j; + } + if (broadcast_type_C == 4) + pC += j; + } + + int jj = 0; +#if __loongarch_asx + const __m256 _alpha256 = (__m256)__lasx_xvreplfr2vr_s(alpha); + const __m256 _beta256 = (__m256)__lasx_xvreplfr2vr_s(beta); + for (; jj + 7 < max_jj; jj += 8) + { + __m256i _sum0 = __lasx_xvld(pp, 0); + __m256i _sum1 = __lasx_xvld(pp + 8, 0); + __m256i _sum2 = __lasx_xvld(pp + 16, 0); + __m256i _sum3 = __lasx_xvld(pp + 24, 0); + __m256i _sum4 = __lasx_xvld(pp + 32, 0); + __m256i _sum5 = __lasx_xvld(pp + 40, 0); + __m256i _sum6 = __lasx_xvld(pp + 48, 0); + __m256i _sum7 = __lasx_xvld(pp + 56, 0); + __m256i _tmp0 = _sum0; + __m256i _tmp1 = __lasx_xvshuf4i_w(_sum1, _LSX_SHUFFLE(2, 1, 0, 3)); + __m256i _tmp2 = _sum2; + __m256i _tmp3 = __lasx_xvshuf4i_w(_sum3, _LSX_SHUFFLE(2, 1, 0, 3)); + __m256i _tmp4 = _sum4; + __m256i _tmp5 = __lasx_xvshuf4i_w(_sum5, _LSX_SHUFFLE(2, 1, 0, 3)); + __m256i _tmp6 = _sum6; + __m256i _tmp7 = __lasx_xvshuf4i_w(_sum7, _LSX_SHUFFLE(2, 1, 0, 3)); + _sum0 = __lasx_xvilvl_w(_tmp3, _tmp0); + _sum1 = __lasx_xvilvh_w(_tmp3, _tmp0); + _sum2 = __lasx_xvilvl_w(_tmp1, _tmp2); + _sum3 = __lasx_xvilvh_w(_tmp1, _tmp2); + _sum4 = __lasx_xvilvl_w(_tmp7, _tmp4); + _sum5 = __lasx_xvilvh_w(_tmp7, _tmp4); + _sum6 = __lasx_xvilvl_w(_tmp5, _tmp6); + _sum7 = __lasx_xvilvh_w(_tmp5, _tmp6); + _tmp0 = __lasx_xvilvl_d(_sum2, _sum0); + _tmp1 = __lasx_xvilvh_d(_sum2, _sum0); + _tmp2 = __lasx_xvilvl_d(_sum1, _sum3); + _tmp3 = __lasx_xvilvh_d(_sum1, _sum3); + _tmp4 = __lasx_xvilvl_d(_sum6, _sum4); + _tmp5 = __lasx_xvilvh_d(_sum6, _sum4); + _tmp6 = __lasx_xvilvl_d(_sum5, _sum7); + _tmp7 = __lasx_xvilvh_d(_sum5, _sum7); + _tmp1 = __lasx_xvshuf4i_w(_tmp1, _LSX_SHUFFLE(2, 1, 0, 3)); + _tmp3 = __lasx_xvshuf4i_w(_tmp3, _LSX_SHUFFLE(2, 1, 0, 3)); + _tmp5 = __lasx_xvshuf4i_w(_tmp5, _LSX_SHUFFLE(2, 1, 0, 3)); + _tmp7 = __lasx_xvshuf4i_w(_tmp7, _LSX_SHUFFLE(2, 1, 0, 3)); + __m256 _f0 = (__m256)__lasx_xvpermi_q(_tmp4, _tmp0, _LSX_SHUFFLE(0, 3, 0, 0)); + __m256 _f1 = (__m256)__lasx_xvpermi_q(_tmp5, _tmp1, _LSX_SHUFFLE(0, 3, 0, 0)); + __m256 _f2 = (__m256)__lasx_xvpermi_q(_tmp6, _tmp2, _LSX_SHUFFLE(0, 3, 0, 0)); + __m256 _f3 = (__m256)__lasx_xvpermi_q(_tmp7, _tmp3, _LSX_SHUFFLE(0, 3, 0, 0)); + __m256 _f4 = (__m256)__lasx_xvpermi_q(_tmp0, _tmp4, _LSX_SHUFFLE(0, 3, 0, 0)); + __m256 _f5 = (__m256)__lasx_xvpermi_q(_tmp1, _tmp5, _LSX_SHUFFLE(0, 3, 0, 0)); + __m256 _f6 = (__m256)__lasx_xvpermi_q(_tmp2, _tmp6, _LSX_SHUFFLE(0, 3, 0, 0)); + __m256 _f7 = (__m256)__lasx_xvpermi_q(_tmp3, _tmp7, _LSX_SHUFFLE(0, 3, 0, 0)); + pp += 64; + if (pC) + { + if (broadcast_type_C == 0) + { + const __m256 _c = (__m256)__lasx_xvreplfr2vr_s(c0); + _f0 = __lasx_xvfadd_s(_f0, _c); + _f1 = __lasx_xvfadd_s(_f1, _c); + _f2 = __lasx_xvfadd_s(_f2, _c); + _f3 = __lasx_xvfadd_s(_f3, _c); + _f4 = __lasx_xvfadd_s(_f4, _c); + _f5 = __lasx_xvfadd_s(_f5, _c); + _f6 = __lasx_xvfadd_s(_f6, _c); + _f7 = __lasx_xvfadd_s(_f7, _c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = __lasx_xvfadd_s(_f0, (__m256)__lasx_xvreplfr2vr_s(c0)); + _f1 = __lasx_xvfadd_s(_f1, (__m256)__lasx_xvreplfr2vr_s(c1)); + _f2 = __lasx_xvfadd_s(_f2, (__m256)__lasx_xvreplfr2vr_s(c2)); + _f3 = __lasx_xvfadd_s(_f3, (__m256)__lasx_xvreplfr2vr_s(c3)); + _f4 = __lasx_xvfadd_s(_f4, (__m256)__lasx_xvreplfr2vr_s(c4)); + _f5 = __lasx_xvfadd_s(_f5, (__m256)__lasx_xvreplfr2vr_s(c5)); + _f6 = __lasx_xvfadd_s(_f6, (__m256)__lasx_xvreplfr2vr_s(c6)); + _f7 = __lasx_xvfadd_s(_f7, (__m256)__lasx_xvreplfr2vr_s(c7)); + } + if (broadcast_type_C == 3) + { + const __m256 _c0 = (__m256)__lasx_xvld(pC, 0); + const __m256 _c1 = (__m256)__lasx_xvld(pC + c_hstep, 0); + const __m256 _c2 = (__m256)__lasx_xvld(pC + c_hstep * 2, 0); + const __m256 _c3 = (__m256)__lasx_xvld(pC + c_hstep * 3, 0); + const __m256 _c4 = (__m256)__lasx_xvld(pC + c_hstep * 4, 0); + const __m256 _c5 = (__m256)__lasx_xvld(pC + c_hstep * 5, 0); + const __m256 _c6 = (__m256)__lasx_xvld(pC + c_hstep * 6, 0); + const __m256 _c7 = (__m256)__lasx_xvld(pC + c_hstep * 7, 0); + if (beta == 1.f) + { + _f0 = __lasx_xvfadd_s(_f0, _c0); + _f1 = __lasx_xvfadd_s(_f1, _c1); + _f2 = __lasx_xvfadd_s(_f2, _c2); + _f3 = __lasx_xvfadd_s(_f3, _c3); + _f4 = __lasx_xvfadd_s(_f4, _c4); + _f5 = __lasx_xvfadd_s(_f5, _c5); + _f6 = __lasx_xvfadd_s(_f6, _c6); + _f7 = __lasx_xvfadd_s(_f7, _c7); + } + else + { + _f0 = __lasx_xvfmadd_s(_c0, _beta256, _f0); + _f1 = __lasx_xvfmadd_s(_c1, _beta256, _f1); + _f2 = __lasx_xvfmadd_s(_c2, _beta256, _f2); + _f3 = __lasx_xvfmadd_s(_c3, _beta256, _f3); + _f4 = __lasx_xvfmadd_s(_c4, _beta256, _f4); + _f5 = __lasx_xvfmadd_s(_c5, _beta256, _f5); + _f6 = __lasx_xvfmadd_s(_c6, _beta256, _f6); + _f7 = __lasx_xvfmadd_s(_c7, _beta256, _f7); + } + } + if (broadcast_type_C == 4) + { + __m256 _c = (__m256)__lasx_xvld(pC, 0); + if (beta != 1.f) + _c = __lasx_xvfmul_s(_c, _beta256); + _f0 = __lasx_xvfadd_s(_f0, _c); + _f1 = __lasx_xvfadd_s(_f1, _c); + _f2 = __lasx_xvfadd_s(_f2, _c); + _f3 = __lasx_xvfadd_s(_f3, _c); + _f4 = __lasx_xvfadd_s(_f4, _c); + _f5 = __lasx_xvfadd_s(_f5, _c); + _f6 = __lasx_xvfadd_s(_f6, _c); + _f7 = __lasx_xvfadd_s(_f7, _c); + } + } + if (alpha != 1.f) + { + _f0 = __lasx_xvfmul_s(_f0, _alpha256); + _f1 = __lasx_xvfmul_s(_f1, _alpha256); + _f2 = __lasx_xvfmul_s(_f2, _alpha256); + _f3 = __lasx_xvfmul_s(_f3, _alpha256); + _f4 = __lasx_xvfmul_s(_f4, _alpha256); + _f5 = __lasx_xvfmul_s(_f5, _alpha256); + _f6 = __lasx_xvfmul_s(_f6, _alpha256); + _f7 = __lasx_xvfmul_s(_f7, _alpha256); + } + __lasx_xvst(_f0, p0, 0); + __lasx_xvst(_f1, p1, 0); + __lasx_xvst(_f2, p2, 0); + __lasx_xvst(_f3, p3, 0); + __lasx_xvst(_f4, p4, 0); + __lasx_xvst(_f5, p5, 0); + __lasx_xvst(_f6, p6, 0); + __lasx_xvst(_f7, p7, 0); + p0 += 8; + p1 += 8; + p2 += 8; + p3 += 8; + p4 += 8; + p5 += 8; + p6 += 8; + p7 += 8; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 8; + } +#endif + const __m128 _alpha128 = __lsx_vreplfr2vr_s(alpha); + const __m128 _beta128 = __lsx_vreplfr2vr_s(beta); + for (; jj + 3 < max_jj; jj += 4) + { + __m128i _sum0 = __lsx_vld(pp, 0); + __m128i _sum1 = __lsx_vld(pp + 8, 0); + __m128i _sum2 = __lsx_vld(pp + 16, 0); + __m128i _sum3 = __lsx_vld(pp + 24, 0); + _sum2 = __lsx_vshuf4i_w(_sum2, _LSX_SHUFFLE(1, 0, 3, 2)); + _sum3 = __lsx_vshuf4i_w(_sum3, _LSX_SHUFFLE(1, 0, 3, 2)); + transpose4x4_epi32(_sum0, _sum1, _sum2, _sum3); + _sum1 = __lsx_vshuf4i_w(_sum1, _LSX_SHUFFLE(2, 1, 0, 3)); + _sum2 = __lsx_vshuf4i_w(_sum2, _LSX_SHUFFLE(1, 0, 3, 2)); + _sum3 = __lsx_vshuf4i_w(_sum3, _LSX_SHUFFLE(0, 3, 2, 1)); + __m128i _sum4 = __lsx_vld(pp + 4, 0); + __m128i _sum5 = __lsx_vld(pp + 12, 0); + __m128i _sum6 = __lsx_vld(pp + 20, 0); + __m128i _sum7 = __lsx_vld(pp + 28, 0); + _sum6 = __lsx_vshuf4i_w(_sum6, _LSX_SHUFFLE(1, 0, 3, 2)); + _sum7 = __lsx_vshuf4i_w(_sum7, _LSX_SHUFFLE(1, 0, 3, 2)); + transpose4x4_epi32(_sum4, _sum5, _sum6, _sum7); + _sum5 = __lsx_vshuf4i_w(_sum5, _LSX_SHUFFLE(2, 1, 0, 3)); + _sum6 = __lsx_vshuf4i_w(_sum6, _LSX_SHUFFLE(1, 0, 3, 2)); + _sum7 = __lsx_vshuf4i_w(_sum7, _LSX_SHUFFLE(0, 3, 2, 1)); + __m128 _f0 = (__m128)_sum0; + __m128 _f1 = (__m128)_sum1; + __m128 _f2 = (__m128)_sum2; + __m128 _f3 = (__m128)_sum3; + __m128 _f4 = (__m128)_sum4; + __m128 _f5 = (__m128)_sum5; + __m128 _f6 = (__m128)_sum6; + __m128 _f7 = (__m128)_sum7; + pp += 32; + if (pC) + { + if (broadcast_type_C == 0) + { + const __m128 _c = __lsx_vreplfr2vr_s(c0); + _f0 = __lsx_vfadd_s(_f0, _c); + _f1 = __lsx_vfadd_s(_f1, _c); + _f2 = __lsx_vfadd_s(_f2, _c); + _f3 = __lsx_vfadd_s(_f3, _c); + _f4 = __lsx_vfadd_s(_f4, _c); + _f5 = __lsx_vfadd_s(_f5, _c); + _f6 = __lsx_vfadd_s(_f6, _c); + _f7 = __lsx_vfadd_s(_f7, _c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = __lsx_vfadd_s(_f0, __lsx_vreplfr2vr_s(c0)); + _f1 = __lsx_vfadd_s(_f1, __lsx_vreplfr2vr_s(c1)); + _f2 = __lsx_vfadd_s(_f2, __lsx_vreplfr2vr_s(c2)); + _f3 = __lsx_vfadd_s(_f3, __lsx_vreplfr2vr_s(c3)); + _f4 = __lsx_vfadd_s(_f4, __lsx_vreplfr2vr_s(c4)); + _f5 = __lsx_vfadd_s(_f5, __lsx_vreplfr2vr_s(c5)); + _f6 = __lsx_vfadd_s(_f6, __lsx_vreplfr2vr_s(c6)); + _f7 = __lsx_vfadd_s(_f7, __lsx_vreplfr2vr_s(c7)); + } + if (broadcast_type_C == 3) + { + const __m128 _c0 = (__m128)__lsx_vld(pC, 0); + const __m128 _c1 = (__m128)__lsx_vld(pC + c_hstep, 0); + const __m128 _c2 = (__m128)__lsx_vld(pC + c_hstep * 2, 0); + const __m128 _c3 = (__m128)__lsx_vld(pC + c_hstep * 3, 0); + const __m128 _c4 = (__m128)__lsx_vld(pC + c_hstep * 4, 0); + const __m128 _c5 = (__m128)__lsx_vld(pC + c_hstep * 5, 0); + const __m128 _c6 = (__m128)__lsx_vld(pC + c_hstep * 6, 0); + const __m128 _c7 = (__m128)__lsx_vld(pC + c_hstep * 7, 0); + if (beta == 1.f) + { + _f0 = __lsx_vfadd_s(_f0, _c0); + _f1 = __lsx_vfadd_s(_f1, _c1); + _f2 = __lsx_vfadd_s(_f2, _c2); + _f3 = __lsx_vfadd_s(_f3, _c3); + _f4 = __lsx_vfadd_s(_f4, _c4); + _f5 = __lsx_vfadd_s(_f5, _c5); + _f6 = __lsx_vfadd_s(_f6, _c6); + _f7 = __lsx_vfadd_s(_f7, _c7); + } + else + { + _f0 = __lsx_vfmadd_s(_c0, _beta128, _f0); + _f1 = __lsx_vfmadd_s(_c1, _beta128, _f1); + _f2 = __lsx_vfmadd_s(_c2, _beta128, _f2); + _f3 = __lsx_vfmadd_s(_c3, _beta128, _f3); + _f4 = __lsx_vfmadd_s(_c4, _beta128, _f4); + _f5 = __lsx_vfmadd_s(_c5, _beta128, _f5); + _f6 = __lsx_vfmadd_s(_c6, _beta128, _f6); + _f7 = __lsx_vfmadd_s(_c7, _beta128, _f7); + } + } + if (broadcast_type_C == 4) + { + __m128 _c = (__m128)__lsx_vld(pC, 0); + if (beta != 1.f) + _c = __lsx_vfmul_s(_c, _beta128); + _f0 = __lsx_vfadd_s(_f0, _c); + _f1 = __lsx_vfadd_s(_f1, _c); + _f2 = __lsx_vfadd_s(_f2, _c); + _f3 = __lsx_vfadd_s(_f3, _c); + _f4 = __lsx_vfadd_s(_f4, _c); + _f5 = __lsx_vfadd_s(_f5, _c); + _f6 = __lsx_vfadd_s(_f6, _c); + _f7 = __lsx_vfadd_s(_f7, _c); + } + } + if (alpha != 1.f) + { + _f0 = __lsx_vfmul_s(_f0, _alpha128); + _f1 = __lsx_vfmul_s(_f1, _alpha128); + _f2 = __lsx_vfmul_s(_f2, _alpha128); + _f3 = __lsx_vfmul_s(_f3, _alpha128); + _f4 = __lsx_vfmul_s(_f4, _alpha128); + _f5 = __lsx_vfmul_s(_f5, _alpha128); + _f6 = __lsx_vfmul_s(_f6, _alpha128); + _f7 = __lsx_vfmul_s(_f7, _alpha128); + } + __lsx_vst((__m128i)_f0, p0, 0); + __lsx_vst((__m128i)_f1, p1, 0); + __lsx_vst((__m128i)_f2, p2, 0); + __lsx_vst((__m128i)_f3, p3, 0); + __lsx_vst((__m128i)_f4, p4, 0); + __lsx_vst((__m128i)_f5, p5, 0); + __lsx_vst((__m128i)_f6, p6, 0); + __lsx_vst((__m128i)_f7, p7, 0); + p0 += 4; + p1 += 4; + p2 += 4; + p3 += 4; + p4 += 4; + p5 += 4; + p6 += 4; + p7 += 4; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 4; + } + for (; jj + 1 < max_jj; jj += 2) + { + __m128i _fi0 = __lsx_vldrepl_w(pp, 0); + __m128i _fi1 = __lsx_vldrepl_w(pp + 1, 0); + __m128i _fi2 = __lsx_vldrepl_w(pp + 2, 0); + __m128i _fi3 = __lsx_vldrepl_w(pp + 3, 0); + __m128i _fi4 = __lsx_vldrepl_w(pp + 4, 0); + __m128i _fi5 = __lsx_vldrepl_w(pp + 5, 0); + __m128i _fi6 = __lsx_vldrepl_w(pp + 6, 0); + __m128i _fi7 = __lsx_vldrepl_w(pp + 7, 0); + _fi0 = __lsx_vinsgr2vr_w(_fi0, ((const int*)pp)[8], 1); + _fi1 = __lsx_vinsgr2vr_w(_fi1, ((const int*)pp)[9], 1); + _fi2 = __lsx_vinsgr2vr_w(_fi2, ((const int*)pp)[10], 1); + _fi3 = __lsx_vinsgr2vr_w(_fi3, ((const int*)pp)[11], 1); + _fi4 = __lsx_vinsgr2vr_w(_fi4, ((const int*)pp)[12], 1); + _fi5 = __lsx_vinsgr2vr_w(_fi5, ((const int*)pp)[13], 1); + _fi6 = __lsx_vinsgr2vr_w(_fi6, ((const int*)pp)[14], 1); + _fi7 = __lsx_vinsgr2vr_w(_fi7, ((const int*)pp)[15], 1); + __m128 _f0 = (__m128)_fi0; + __m128 _f1 = (__m128)_fi1; + __m128 _f2 = (__m128)_fi2; + __m128 _f3 = (__m128)_fi3; + __m128 _f4 = (__m128)_fi4; + __m128 _f5 = (__m128)_fi5; + __m128 _f6 = (__m128)_fi6; + __m128 _f7 = (__m128)_fi7; + pp += 16; + if (pC) + { + if (broadcast_type_C == 0) + { + const __m128 _c = __lsx_vreplfr2vr_s(c0); + _f0 = __lsx_vfadd_s(_f0, _c); + _f1 = __lsx_vfadd_s(_f1, _c); + _f2 = __lsx_vfadd_s(_f2, _c); + _f3 = __lsx_vfadd_s(_f3, _c); + _f4 = __lsx_vfadd_s(_f4, _c); + _f5 = __lsx_vfadd_s(_f5, _c); + _f6 = __lsx_vfadd_s(_f6, _c); + _f7 = __lsx_vfadd_s(_f7, _c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = __lsx_vfadd_s(_f0, __lsx_vreplfr2vr_s(c0)); + _f1 = __lsx_vfadd_s(_f1, __lsx_vreplfr2vr_s(c1)); + _f2 = __lsx_vfadd_s(_f2, __lsx_vreplfr2vr_s(c2)); + _f3 = __lsx_vfadd_s(_f3, __lsx_vreplfr2vr_s(c3)); + _f4 = __lsx_vfadd_s(_f4, __lsx_vreplfr2vr_s(c4)); + _f5 = __lsx_vfadd_s(_f5, __lsx_vreplfr2vr_s(c5)); + _f6 = __lsx_vfadd_s(_f6, __lsx_vreplfr2vr_s(c6)); + _f7 = __lsx_vfadd_s(_f7, __lsx_vreplfr2vr_s(c7)); + } + if (broadcast_type_C == 3) + { + const __m128 _c0 = (__m128)__lsx_vldrepl_d(pC, 0); + const __m128 _c1 = (__m128)__lsx_vldrepl_d(pC + c_hstep, 0); + const __m128 _c2 = (__m128)__lsx_vldrepl_d(pC + c_hstep * 2, 0); + const __m128 _c3 = (__m128)__lsx_vldrepl_d(pC + c_hstep * 3, 0); + const __m128 _c4 = (__m128)__lsx_vldrepl_d(pC + c_hstep * 4, 0); + const __m128 _c5 = (__m128)__lsx_vldrepl_d(pC + c_hstep * 5, 0); + const __m128 _c6 = (__m128)__lsx_vldrepl_d(pC + c_hstep * 6, 0); + const __m128 _c7 = (__m128)__lsx_vldrepl_d(pC + c_hstep * 7, 0); + if (beta == 1.f) + { + _f0 = __lsx_vfadd_s(_f0, _c0); + _f1 = __lsx_vfadd_s(_f1, _c1); + _f2 = __lsx_vfadd_s(_f2, _c2); + _f3 = __lsx_vfadd_s(_f3, _c3); + _f4 = __lsx_vfadd_s(_f4, _c4); + _f5 = __lsx_vfadd_s(_f5, _c5); + _f6 = __lsx_vfadd_s(_f6, _c6); + _f7 = __lsx_vfadd_s(_f7, _c7); + } + else + { + _f0 = __lsx_vfmadd_s(_c0, _beta128, _f0); + _f1 = __lsx_vfmadd_s(_c1, _beta128, _f1); + _f2 = __lsx_vfmadd_s(_c2, _beta128, _f2); + _f3 = __lsx_vfmadd_s(_c3, _beta128, _f3); + _f4 = __lsx_vfmadd_s(_c4, _beta128, _f4); + _f5 = __lsx_vfmadd_s(_c5, _beta128, _f5); + _f6 = __lsx_vfmadd_s(_c6, _beta128, _f6); + _f7 = __lsx_vfmadd_s(_c7, _beta128, _f7); + } + } + if (broadcast_type_C == 4) + { + __m128 _c = (__m128)__lsx_vldrepl_d(pC, 0); + if (beta != 1.f) + _c = __lsx_vfmul_s(_c, _beta128); + _f0 = __lsx_vfadd_s(_f0, _c); + _f1 = __lsx_vfadd_s(_f1, _c); + _f2 = __lsx_vfadd_s(_f2, _c); + _f3 = __lsx_vfadd_s(_f3, _c); + _f4 = __lsx_vfadd_s(_f4, _c); + _f5 = __lsx_vfadd_s(_f5, _c); + _f6 = __lsx_vfadd_s(_f6, _c); + _f7 = __lsx_vfadd_s(_f7, _c); + } + } + if (alpha != 1.f) + { + _f0 = __lsx_vfmul_s(_f0, _alpha128); + _f1 = __lsx_vfmul_s(_f1, _alpha128); + _f2 = __lsx_vfmul_s(_f2, _alpha128); + _f3 = __lsx_vfmul_s(_f3, _alpha128); + _f4 = __lsx_vfmul_s(_f4, _alpha128); + _f5 = __lsx_vfmul_s(_f5, _alpha128); + _f6 = __lsx_vfmul_s(_f6, _alpha128); + _f7 = __lsx_vfmul_s(_f7, _alpha128); + } + __lsx_vstelm_d((__m128i)_f0, p0, 0, 0); + __lsx_vstelm_d((__m128i)_f1, p1, 0, 0); + __lsx_vstelm_d((__m128i)_f2, p2, 0, 0); + __lsx_vstelm_d((__m128i)_f3, p3, 0, 0); + __lsx_vstelm_d((__m128i)_f4, p4, 0, 0); + __lsx_vstelm_d((__m128i)_f5, p5, 0, 0); + __lsx_vstelm_d((__m128i)_f6, p6, 0, 0); + __lsx_vstelm_d((__m128i)_f7, p7, 0, 0); + p0 += 2; + p1 += 2; + p2 += 2; + p3 += 2; + p4 += 2; + p5 += 2; + p6 += 2; + p7 += 2; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 2; + } + for (; jj < max_jj; jj++) + { + __m128 _f0 = (__m128)__lsx_vld(pp, 0); + __m128 _f4 = (__m128)__lsx_vld(pp + 4, 0); + pp += 8; + if (pC) + { + if (broadcast_type_C == 0) + { + const __m128 _c = __lsx_vreplfr2vr_s(c0); + _f0 = __lsx_vfadd_s(_f0, _c); + _f4 = __lsx_vfadd_s(_f4, _c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = __lsx_vfadd_s(_f0, __lsx_vfmul_s((__m128)__lsx_vld(pC_base + i + ii, 0), _beta128)); + _f4 = __lsx_vfadd_s(_f4, __lsx_vfmul_s((__m128)__lsx_vld(pC_base + i + ii + 4, 0), _beta128)); + } + if (broadcast_type_C == 3) + { + __m128i _c0 = __lsx_vreplgr2vr_w(((const int*)pC)[0]); + _c0 = __lsx_vinsgr2vr_w(_c0, ((const int*)(pC + c_hstep))[0], 1); + _c0 = __lsx_vinsgr2vr_w(_c0, ((const int*)(pC + c_hstep * 2))[0], 2); + _c0 = __lsx_vinsgr2vr_w(_c0, ((const int*)(pC + c_hstep * 3))[0], 3); + __m128i _c4 = __lsx_vreplgr2vr_w(((const int*)(pC + c_hstep * 4))[0]); + _c4 = __lsx_vinsgr2vr_w(_c4, ((const int*)(pC + c_hstep * 5))[0], 1); + _c4 = __lsx_vinsgr2vr_w(_c4, ((const int*)(pC + c_hstep * 6))[0], 2); + _c4 = __lsx_vinsgr2vr_w(_c4, ((const int*)(pC + c_hstep * 7))[0], 3); + if (beta == 1.f) + { + _f0 = __lsx_vfadd_s(_f0, (__m128)_c0); + _f4 = __lsx_vfadd_s(_f4, (__m128)_c4); + } + else + { + _f0 = __lsx_vfmadd_s((__m128)_c0, _beta128, _f0); + _f4 = __lsx_vfmadd_s((__m128)_c4, _beta128, _f4); + } + } + if (broadcast_type_C == 4) + { + const __m128 _c = __lsx_vreplfr2vr_s(beta == 1.f ? pC[0] : pC[0] * beta); + _f0 = __lsx_vfadd_s(_f0, _c); + _f4 = __lsx_vfadd_s(_f4, _c); + } + } + if (alpha != 1.f) + { + _f0 = __lsx_vfmul_s(_f0, _alpha128); + _f4 = __lsx_vfmul_s(_f4, _alpha128); + } + __lsx_vstelm_w((__m128i)_f0, p0, 0, 0); + __lsx_vstelm_w((__m128i)_f0, p1, 0, 1); + __lsx_vstelm_w((__m128i)_f0, p2, 0, 2); + __lsx_vstelm_w((__m128i)_f0, p3, 0, 3); + __lsx_vstelm_w((__m128i)_f4, p4, 0, 0); + __lsx_vstelm_w((__m128i)_f4, p5, 0, 1); + __lsx_vstelm_w((__m128i)_f4, p6, 0, 2); + __lsx_vstelm_w((__m128i)_f4, p7, 0, 3); + p0++; + p1++; + p2++; + p3++; + p4++; + p5++; + p6++; + p7++; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC++; + } + } +#endif // __loongarch_sx + for (; ii + 3 < max_ii; ii += 4) + { + float* p0 = outptr + (size_t)(i + ii) * N + j; + float* p1 = p0 + N; + float* p2 = p1 + N; + float* p3 = p2 + N; + + float c0 = 0.f; + float c1 = 0.f; + float c2 = 0.f; + float c3 = 0.f; + const float* pC = pC_base; + if (pC) + { + if (broadcast_type_C == 0) + { + c0 = pC[0] * beta; + c1 = c0; + c2 = c0; + c3 = c0; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + c0 = pC[i + ii] * beta; + c1 = pC[i + ii + 1] * beta; + c2 = pC[i + ii + 2] * beta; + c3 = pC[i + ii + 3] * beta; + } + if (broadcast_type_C == 3) + { + pC += (size_t)(i + ii) * c_hstep + j; + } + if (broadcast_type_C == 4) + { + pC += j; + } + } + + int jj = 0; +#if __loongarch_asx + const __m256 _alpha256 = (__m256)__lasx_xvreplfr2vr_s(alpha); + const __m256 _beta256 = (__m256)__lasx_xvreplfr2vr_s(beta); + for (; jj + 15 < max_jj; jj += 16) + { + __m256 _f00 = (__m256)__lasx_xvld(pp + (ii + 0) * max_jj + jj, 0); + __m256 _f01 = (__m256)__lasx_xvld(pp + (ii + 0) * max_jj + jj + 8, 0); + __m256 _f10 = (__m256)__lasx_xvld(pp + (ii + 1) * max_jj + jj, 0); + __m256 _f11 = (__m256)__lasx_xvld(pp + (ii + 1) * max_jj + jj + 8, 0); + __m256 _f20 = (__m256)__lasx_xvld(pp + (ii + 2) * max_jj + jj, 0); + __m256 _f21 = (__m256)__lasx_xvld(pp + (ii + 2) * max_jj + jj + 8, 0); + __m256 _f30 = (__m256)__lasx_xvld(pp + (ii + 3) * max_jj + jj, 0); + __m256 _f31 = (__m256)__lasx_xvld(pp + (ii + 3) * max_jj + jj + 8, 0); + if (pC) + { + if (broadcast_type_C == 0) + { + const __m256 _c = (__m256)__lasx_xvreplfr2vr_s(c0); + _f00 = __lasx_xvfadd_s(_f00, _c); + _f01 = __lasx_xvfadd_s(_f01, _c); + _f10 = __lasx_xvfadd_s(_f10, _c); + _f11 = __lasx_xvfadd_s(_f11, _c); + _f20 = __lasx_xvfadd_s(_f20, _c); + _f21 = __lasx_xvfadd_s(_f21, _c); + _f30 = __lasx_xvfadd_s(_f30, _c); + _f31 = __lasx_xvfadd_s(_f31, _c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + const __m256 _c0 = (__m256)__lasx_xvreplfr2vr_s(c0); + const __m256 _c1 = (__m256)__lasx_xvreplfr2vr_s(c1); + const __m256 _c2 = (__m256)__lasx_xvreplfr2vr_s(c2); + const __m256 _c3 = (__m256)__lasx_xvreplfr2vr_s(c3); + _f00 = __lasx_xvfadd_s(_f00, _c0); + _f01 = __lasx_xvfadd_s(_f01, _c0); + _f10 = __lasx_xvfadd_s(_f10, _c1); + _f11 = __lasx_xvfadd_s(_f11, _c1); + _f20 = __lasx_xvfadd_s(_f20, _c2); + _f21 = __lasx_xvfadd_s(_f21, _c2); + _f30 = __lasx_xvfadd_s(_f30, _c3); + _f31 = __lasx_xvfadd_s(_f31, _c3); + } + if (broadcast_type_C == 3) + { + const __m256 _c00 = (__m256)__lasx_xvld(pC, 0); + const __m256 _c01 = (__m256)__lasx_xvld(pC + 8, 0); + const __m256 _c10 = (__m256)__lasx_xvld(pC + c_hstep, 0); + const __m256 _c11 = (__m256)__lasx_xvld(pC + c_hstep + 8, 0); + const __m256 _c20 = (__m256)__lasx_xvld(pC + c_hstep * 2, 0); + const __m256 _c21 = (__m256)__lasx_xvld(pC + c_hstep * 2 + 8, 0); + const __m256 _c30 = (__m256)__lasx_xvld(pC + c_hstep * 3, 0); + const __m256 _c31 = (__m256)__lasx_xvld(pC + c_hstep * 3 + 8, 0); + if (beta == 1.f) + { + _f00 = __lasx_xvfadd_s(_f00, _c00); + _f01 = __lasx_xvfadd_s(_f01, _c01); + _f10 = __lasx_xvfadd_s(_f10, _c10); + _f11 = __lasx_xvfadd_s(_f11, _c11); + _f20 = __lasx_xvfadd_s(_f20, _c20); + _f21 = __lasx_xvfadd_s(_f21, _c21); + _f30 = __lasx_xvfadd_s(_f30, _c30); + _f31 = __lasx_xvfadd_s(_f31, _c31); + } + else + { + _f00 = __lasx_xvfmadd_s(_c00, _beta256, _f00); + _f01 = __lasx_xvfmadd_s(_c01, _beta256, _f01); + _f10 = __lasx_xvfmadd_s(_c10, _beta256, _f10); + _f11 = __lasx_xvfmadd_s(_c11, _beta256, _f11); + _f20 = __lasx_xvfmadd_s(_c20, _beta256, _f20); + _f21 = __lasx_xvfmadd_s(_c21, _beta256, _f21); + _f30 = __lasx_xvfmadd_s(_c30, _beta256, _f30); + _f31 = __lasx_xvfmadd_s(_c31, _beta256, _f31); + } + } + if (broadcast_type_C == 4) + { + __m256 _c0 = (__m256)__lasx_xvld(pC, 0); + __m256 _c1 = (__m256)__lasx_xvld(pC + 8, 0); + if (beta != 1.f) + { + _c0 = __lasx_xvfmul_s(_c0, _beta256); + _c1 = __lasx_xvfmul_s(_c1, _beta256); + } + _f00 = __lasx_xvfadd_s(_f00, _c0); + _f01 = __lasx_xvfadd_s(_f01, _c1); + _f10 = __lasx_xvfadd_s(_f10, _c0); + _f11 = __lasx_xvfadd_s(_f11, _c1); + _f20 = __lasx_xvfadd_s(_f20, _c0); + _f21 = __lasx_xvfadd_s(_f21, _c1); + _f30 = __lasx_xvfadd_s(_f30, _c0); + _f31 = __lasx_xvfadd_s(_f31, _c1); + } + } + if (alpha != 1.f) + { + _f00 = __lasx_xvfmul_s(_f00, _alpha256); + _f01 = __lasx_xvfmul_s(_f01, _alpha256); + _f10 = __lasx_xvfmul_s(_f10, _alpha256); + _f11 = __lasx_xvfmul_s(_f11, _alpha256); + _f20 = __lasx_xvfmul_s(_f20, _alpha256); + _f21 = __lasx_xvfmul_s(_f21, _alpha256); + _f30 = __lasx_xvfmul_s(_f30, _alpha256); + _f31 = __lasx_xvfmul_s(_f31, _alpha256); + } + __lasx_xvst(_f00, p0, 0); + __lasx_xvst(_f01, p0 + 8, 0); + __lasx_xvst(_f10, p1, 0); + __lasx_xvst(_f11, p1 + 8, 0); + __lasx_xvst(_f20, p2, 0); + __lasx_xvst(_f21, p2 + 8, 0); + __lasx_xvst(_f30, p3, 0); + __lasx_xvst(_f31, p3 + 8, 0); + p0 += 16; + p1 += 16; + p2 += 16; + p3 += 16; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 16; + } + for (; jj + 7 < max_jj; jj += 8) + { + __m256 _f0 = (__m256)__lasx_xvld(pp + (ii + 0) * max_jj + jj, 0); + __m256 _f1 = (__m256)__lasx_xvld(pp + (ii + 1) * max_jj + jj, 0); + __m256 _f2 = (__m256)__lasx_xvld(pp + (ii + 2) * max_jj + jj, 0); + __m256 _f3 = (__m256)__lasx_xvld(pp + (ii + 3) * max_jj + jj, 0); + if (pC) + { + if (broadcast_type_C == 0) + { + _f0 = __lasx_xvfadd_s(_f0, (__m256)__lasx_xvreplfr2vr_s(c0)); + _f1 = __lasx_xvfadd_s(_f1, (__m256)__lasx_xvreplfr2vr_s(c0)); + _f2 = __lasx_xvfadd_s(_f2, (__m256)__lasx_xvreplfr2vr_s(c0)); + _f3 = __lasx_xvfadd_s(_f3, (__m256)__lasx_xvreplfr2vr_s(c0)); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = __lasx_xvfadd_s(_f0, (__m256)__lasx_xvreplfr2vr_s(c0)); + _f1 = __lasx_xvfadd_s(_f1, (__m256)__lasx_xvreplfr2vr_s(c1)); + _f2 = __lasx_xvfadd_s(_f2, (__m256)__lasx_xvreplfr2vr_s(c2)); + _f3 = __lasx_xvfadd_s(_f3, (__m256)__lasx_xvreplfr2vr_s(c3)); + } + if (broadcast_type_C == 3) + { + __m256 _c0 = (__m256)__lasx_xvld(pC, 0); + __m256 _c1 = (__m256)__lasx_xvld(pC + c_hstep, 0); + __m256 _c2 = (__m256)__lasx_xvld(pC + c_hstep * 2, 0); + __m256 _c3 = (__m256)__lasx_xvld(pC + c_hstep * 3, 0); + if (beta == 1.f) + { + _f0 = __lasx_xvfadd_s(_f0, _c0); + _f1 = __lasx_xvfadd_s(_f1, _c1); + _f2 = __lasx_xvfadd_s(_f2, _c2); + _f3 = __lasx_xvfadd_s(_f3, _c3); + } + else + { + _f0 = __lasx_xvfmadd_s(_c0, _beta256, _f0); + _f1 = __lasx_xvfmadd_s(_c1, _beta256, _f1); + _f2 = __lasx_xvfmadd_s(_c2, _beta256, _f2); + _f3 = __lasx_xvfmadd_s(_c3, _beta256, _f3); + } + } + if (broadcast_type_C == 4) + { + __m256 _c = (__m256)__lasx_xvld(pC, 0); + if (beta != 1.f) + _c = __lasx_xvfmul_s(_c, _beta256); + _f0 = __lasx_xvfadd_s(_f0, _c); + _f1 = __lasx_xvfadd_s(_f1, _c); + _f2 = __lasx_xvfadd_s(_f2, _c); + _f3 = __lasx_xvfadd_s(_f3, _c); + } + } + if (alpha != 1.f) + { + _f0 = __lasx_xvfmul_s(_f0, _alpha256); + _f1 = __lasx_xvfmul_s(_f1, _alpha256); + _f2 = __lasx_xvfmul_s(_f2, _alpha256); + _f3 = __lasx_xvfmul_s(_f3, _alpha256); + } + __lasx_xvst(_f0, p0, 0); + __lasx_xvst(_f1, p1, 0); + __lasx_xvst(_f2, p2, 0); + __lasx_xvst(_f3, p3, 0); + p0 += 8; + p1 += 8; + p2 += 8; + p3 += 8; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 8; + } +#endif +#if __loongarch_sx + const __m128 _alpha128 = __lsx_vreplfr2vr_s(alpha); + const __m128 _beta128 = __lsx_vreplfr2vr_s(beta); + for (; jj + 7 < max_jj; jj += 8) + { + __m128 _f00 = (__m128)__lsx_vld(pp + (ii + 0) * max_jj + jj, 0); + __m128 _f01 = (__m128)__lsx_vld(pp + (ii + 0) * max_jj + jj + 4, 0); + __m128 _f10 = (__m128)__lsx_vld(pp + (ii + 1) * max_jj + jj, 0); + __m128 _f11 = (__m128)__lsx_vld(pp + (ii + 1) * max_jj + jj + 4, 0); + __m128 _f20 = (__m128)__lsx_vld(pp + (ii + 2) * max_jj + jj, 0); + __m128 _f21 = (__m128)__lsx_vld(pp + (ii + 2) * max_jj + jj + 4, 0); + __m128 _f30 = (__m128)__lsx_vld(pp + (ii + 3) * max_jj + jj, 0); + __m128 _f31 = (__m128)__lsx_vld(pp + (ii + 3) * max_jj + jj + 4, 0); + if (pC) + { + if (broadcast_type_C == 0) + { + const __m128 _c = __lsx_vreplfr2vr_s(c0); + _f00 = __lsx_vfadd_s(_f00, _c); + _f01 = __lsx_vfadd_s(_f01, _c); + _f10 = __lsx_vfadd_s(_f10, _c); + _f11 = __lsx_vfadd_s(_f11, _c); + _f20 = __lsx_vfadd_s(_f20, _c); + _f21 = __lsx_vfadd_s(_f21, _c); + _f30 = __lsx_vfadd_s(_f30, _c); + _f31 = __lsx_vfadd_s(_f31, _c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + const __m128 _c0 = __lsx_vreplfr2vr_s(c0); + const __m128 _c1 = __lsx_vreplfr2vr_s(c1); + const __m128 _c2 = __lsx_vreplfr2vr_s(c2); + const __m128 _c3 = __lsx_vreplfr2vr_s(c3); + _f00 = __lsx_vfadd_s(_f00, _c0); + _f01 = __lsx_vfadd_s(_f01, _c0); + _f10 = __lsx_vfadd_s(_f10, _c1); + _f11 = __lsx_vfadd_s(_f11, _c1); + _f20 = __lsx_vfadd_s(_f20, _c2); + _f21 = __lsx_vfadd_s(_f21, _c2); + _f30 = __lsx_vfadd_s(_f30, _c3); + _f31 = __lsx_vfadd_s(_f31, _c3); + } + if (broadcast_type_C == 3) + { + const __m128 _c00 = (__m128)__lsx_vld(pC, 0); + const __m128 _c01 = (__m128)__lsx_vld(pC + 4, 0); + const __m128 _c10 = (__m128)__lsx_vld(pC + c_hstep, 0); + const __m128 _c11 = (__m128)__lsx_vld(pC + c_hstep + 4, 0); + const __m128 _c20 = (__m128)__lsx_vld(pC + c_hstep * 2, 0); + const __m128 _c21 = (__m128)__lsx_vld(pC + c_hstep * 2 + 4, 0); + const __m128 _c30 = (__m128)__lsx_vld(pC + c_hstep * 3, 0); + const __m128 _c31 = (__m128)__lsx_vld(pC + c_hstep * 3 + 4, 0); + if (beta == 1.f) + { + _f00 = __lsx_vfadd_s(_f00, _c00); + _f01 = __lsx_vfadd_s(_f01, _c01); + _f10 = __lsx_vfadd_s(_f10, _c10); + _f11 = __lsx_vfadd_s(_f11, _c11); + _f20 = __lsx_vfadd_s(_f20, _c20); + _f21 = __lsx_vfadd_s(_f21, _c21); + _f30 = __lsx_vfadd_s(_f30, _c30); + _f31 = __lsx_vfadd_s(_f31, _c31); + } + else + { + _f00 = __lsx_vfmadd_s(_c00, _beta128, _f00); + _f01 = __lsx_vfmadd_s(_c01, _beta128, _f01); + _f10 = __lsx_vfmadd_s(_c10, _beta128, _f10); + _f11 = __lsx_vfmadd_s(_c11, _beta128, _f11); + _f20 = __lsx_vfmadd_s(_c20, _beta128, _f20); + _f21 = __lsx_vfmadd_s(_c21, _beta128, _f21); + _f30 = __lsx_vfmadd_s(_c30, _beta128, _f30); + _f31 = __lsx_vfmadd_s(_c31, _beta128, _f31); + } + } + if (broadcast_type_C == 4) + { + __m128 _c0 = (__m128)__lsx_vld(pC, 0); + __m128 _c1 = (__m128)__lsx_vld(pC + 4, 0); + if (beta != 1.f) + { + _c0 = __lsx_vfmul_s(_c0, _beta128); + _c1 = __lsx_vfmul_s(_c1, _beta128); + } + _f00 = __lsx_vfadd_s(_f00, _c0); + _f01 = __lsx_vfadd_s(_f01, _c1); + _f10 = __lsx_vfadd_s(_f10, _c0); + _f11 = __lsx_vfadd_s(_f11, _c1); + _f20 = __lsx_vfadd_s(_f20, _c0); + _f21 = __lsx_vfadd_s(_f21, _c1); + _f30 = __lsx_vfadd_s(_f30, _c0); + _f31 = __lsx_vfadd_s(_f31, _c1); + } + } + if (alpha != 1.f) + { + _f00 = __lsx_vfmul_s(_f00, _alpha128); + _f01 = __lsx_vfmul_s(_f01, _alpha128); + _f10 = __lsx_vfmul_s(_f10, _alpha128); + _f11 = __lsx_vfmul_s(_f11, _alpha128); + _f20 = __lsx_vfmul_s(_f20, _alpha128); + _f21 = __lsx_vfmul_s(_f21, _alpha128); + _f30 = __lsx_vfmul_s(_f30, _alpha128); + _f31 = __lsx_vfmul_s(_f31, _alpha128); + } + __lsx_vst((__m128i)_f00, p0, 0); + __lsx_vst((__m128i)_f01, p0 + 4, 0); + __lsx_vst((__m128i)_f10, p1, 0); + __lsx_vst((__m128i)_f11, p1 + 4, 0); + __lsx_vst((__m128i)_f20, p2, 0); + __lsx_vst((__m128i)_f21, p2 + 4, 0); + __lsx_vst((__m128i)_f30, p3, 0); + __lsx_vst((__m128i)_f31, p3 + 4, 0); + p0 += 8; + p1 += 8; + p2 += 8; + p3 += 8; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 8; + } + for (; jj + 3 < max_jj; jj += 4) + { + __m128 _f0 = (__m128)__lsx_vld(pp + (ii + 0) * max_jj + jj, 0); + __m128 _f1 = (__m128)__lsx_vld(pp + (ii + 1) * max_jj + jj, 0); + __m128 _f2 = (__m128)__lsx_vld(pp + (ii + 2) * max_jj + jj, 0); + __m128 _f3 = (__m128)__lsx_vld(pp + (ii + 3) * max_jj + jj, 0); + if (pC) + { + if (broadcast_type_C == 0) + { + _f0 = __lsx_vfadd_s(_f0, __lsx_vreplfr2vr_s(c0)); + _f1 = __lsx_vfadd_s(_f1, __lsx_vreplfr2vr_s(c0)); + _f2 = __lsx_vfadd_s(_f2, __lsx_vreplfr2vr_s(c0)); + _f3 = __lsx_vfadd_s(_f3, __lsx_vreplfr2vr_s(c0)); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = __lsx_vfadd_s(_f0, __lsx_vreplfr2vr_s(c0)); + _f1 = __lsx_vfadd_s(_f1, __lsx_vreplfr2vr_s(c1)); + _f2 = __lsx_vfadd_s(_f2, __lsx_vreplfr2vr_s(c2)); + _f3 = __lsx_vfadd_s(_f3, __lsx_vreplfr2vr_s(c3)); + } + if (broadcast_type_C == 3) + { + __m128 _c0 = (__m128)__lsx_vld(pC, 0); + __m128 _c1 = (__m128)__lsx_vld(pC + c_hstep, 0); + __m128 _c2 = (__m128)__lsx_vld(pC + c_hstep * 2, 0); + __m128 _c3 = (__m128)__lsx_vld(pC + c_hstep * 3, 0); + if (beta == 1.f) + { + _f0 = __lsx_vfadd_s(_f0, _c0); + _f1 = __lsx_vfadd_s(_f1, _c1); + _f2 = __lsx_vfadd_s(_f2, _c2); + _f3 = __lsx_vfadd_s(_f3, _c3); + } + else + { + _f0 = __lsx_vfmadd_s(_c0, _beta128, _f0); + _f1 = __lsx_vfmadd_s(_c1, _beta128, _f1); + _f2 = __lsx_vfmadd_s(_c2, _beta128, _f2); + _f3 = __lsx_vfmadd_s(_c3, _beta128, _f3); + } + } + if (broadcast_type_C == 4) + { + __m128 _c = (__m128)__lsx_vld(pC, 0); + if (beta != 1.f) + _c = __lsx_vfmul_s(_c, _beta128); + _f0 = __lsx_vfadd_s(_f0, _c); + _f1 = __lsx_vfadd_s(_f1, _c); + _f2 = __lsx_vfadd_s(_f2, _c); + _f3 = __lsx_vfadd_s(_f3, _c); + } + } + if (alpha != 1.f) + { + _f0 = __lsx_vfmul_s(_f0, _alpha128); + _f1 = __lsx_vfmul_s(_f1, _alpha128); + _f2 = __lsx_vfmul_s(_f2, _alpha128); + _f3 = __lsx_vfmul_s(_f3, _alpha128); + } + __lsx_vst((__m128i)_f0, p0, 0); + __lsx_vst((__m128i)_f1, p1, 0); + __lsx_vst((__m128i)_f2, p2, 0); + __lsx_vst((__m128i)_f3, p3, 0); + p0 += 4; + p1 += 4; + p2 += 4; + p3 += 4; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 4; + } + for (; jj + 1 < max_jj; jj += 2) + { + __m128 _f0 = (__m128)__lsx_vldrepl_d(pp + (ii + 0) * max_jj + jj, 0); + __m128 _f1 = (__m128)__lsx_vldrepl_d(pp + (ii + 1) * max_jj + jj, 0); + __m128 _f2 = (__m128)__lsx_vldrepl_d(pp + (ii + 2) * max_jj + jj, 0); + __m128 _f3 = (__m128)__lsx_vldrepl_d(pp + (ii + 3) * max_jj + jj, 0); + if (pC) + { + if (broadcast_type_C == 0) + { + __m128 _c = __lsx_vreplfr2vr_s(c0); + _f0 = __lsx_vfadd_s(_f0, _c); + _f1 = __lsx_vfadd_s(_f1, _c); + _f2 = __lsx_vfadd_s(_f2, _c); + _f3 = __lsx_vfadd_s(_f3, _c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = __lsx_vfadd_s(_f0, __lsx_vreplfr2vr_s(c0)); + _f1 = __lsx_vfadd_s(_f1, __lsx_vreplfr2vr_s(c1)); + _f2 = __lsx_vfadd_s(_f2, __lsx_vreplfr2vr_s(c2)); + _f3 = __lsx_vfadd_s(_f3, __lsx_vreplfr2vr_s(c3)); + } + if (broadcast_type_C == 3) + { + __m128 _c0 = (__m128)__lsx_vldrepl_d(pC, 0); + __m128 _c1 = (__m128)__lsx_vldrepl_d(pC + c_hstep, 0); + __m128 _c2 = (__m128)__lsx_vldrepl_d(pC + c_hstep * 2, 0); + __m128 _c3 = (__m128)__lsx_vldrepl_d(pC + c_hstep * 3, 0); + if (beta == 1.f) + { + _f0 = __lsx_vfadd_s(_f0, _c0); + _f1 = __lsx_vfadd_s(_f1, _c1); + _f2 = __lsx_vfadd_s(_f2, _c2); + _f3 = __lsx_vfadd_s(_f3, _c3); + } + else + { + _f0 = __lsx_vfmadd_s(_c0, _beta128, _f0); + _f1 = __lsx_vfmadd_s(_c1, _beta128, _f1); + _f2 = __lsx_vfmadd_s(_c2, _beta128, _f2); + _f3 = __lsx_vfmadd_s(_c3, _beta128, _f3); + } + } + if (broadcast_type_C == 4) + { + __m128 _c = (__m128)__lsx_vldrepl_d(pC, 0); + if (beta != 1.f) + _c = __lsx_vfmul_s(_c, _beta128); + _f0 = __lsx_vfadd_s(_f0, _c); + _f1 = __lsx_vfadd_s(_f1, _c); + _f2 = __lsx_vfadd_s(_f2, _c); + _f3 = __lsx_vfadd_s(_f3, _c); + } + } + if (alpha != 1.f) + { + _f0 = __lsx_vfmul_s(_f0, _alpha128); + _f1 = __lsx_vfmul_s(_f1, _alpha128); + _f2 = __lsx_vfmul_s(_f2, _alpha128); + _f3 = __lsx_vfmul_s(_f3, _alpha128); + } + __lsx_vstelm_d((__m128i)_f0, p0, 0, 0); + __lsx_vstelm_d((__m128i)_f1, p1, 0, 0); + __lsx_vstelm_d((__m128i)_f2, p2, 0, 0); + __lsx_vstelm_d((__m128i)_f3, p3, 0, 0); + p0 += 2; + p1 += 2; + p2 += 2; + p3 += 2; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 2; + } +#endif +#if __loongarch_sx + for (; jj < max_jj; jj++) + { + __m128i _fi = __lsx_vreplgr2vr_w(((const int*)(pp + (ii + 0) * max_jj + jj))[0]); + _fi = __lsx_vinsgr2vr_w(_fi, ((const int*)(pp + (ii + 1) * max_jj + jj))[0], 1); + _fi = __lsx_vinsgr2vr_w(_fi, ((const int*)(pp + (ii + 2) * max_jj + jj))[0], 2); + _fi = __lsx_vinsgr2vr_w(_fi, ((const int*)(pp + (ii + 3) * max_jj + jj))[0], 3); + __m128 _f0 = (__m128)_fi; + if (pC) + { + if (broadcast_type_C == 0) + _f0 = __lsx_vfadd_s(_f0, __lsx_vreplfr2vr_s(c0)); + if (broadcast_type_C == 1 || broadcast_type_C == 2) + _f0 = __lsx_vfadd_s(_f0, __lsx_vfmul_s((__m128)__lsx_vld(pC_base + i + ii, 0), _beta128)); + if (broadcast_type_C == 3) + { + __m128i _c0 = __lsx_vreplgr2vr_w(((const int*)pC)[0]); + _c0 = __lsx_vinsgr2vr_w(_c0, ((const int*)(pC + c_hstep))[0], 1); + _c0 = __lsx_vinsgr2vr_w(_c0, ((const int*)(pC + c_hstep * 2))[0], 2); + _c0 = __lsx_vinsgr2vr_w(_c0, ((const int*)(pC + c_hstep * 3))[0], 3); + if (beta == 1.f) + _f0 = __lsx_vfadd_s(_f0, (__m128)_c0); + else + _f0 = __lsx_vfmadd_s((__m128)_c0, _beta128, _f0); + } + if (broadcast_type_C == 4) + _f0 = __lsx_vfadd_s(_f0, __lsx_vreplfr2vr_s(beta == 1.f ? pC[0] : pC[0] * beta)); + } + if (alpha != 1.f) + _f0 = __lsx_vfmul_s(_f0, _alpha128); + __lsx_vstelm_w((__m128i)_f0, p0, 0, 0); + __lsx_vstelm_w((__m128i)_f0, p1, 0, 1); + __lsx_vstelm_w((__m128i)_f0, p2, 0, 2); + __lsx_vstelm_w((__m128i)_f0, p3, 0, 3); + p0++; + p1++; + p2++; + p3++; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC++; + } +#else + for (; jj < max_jj; jj++) + { + float f0 = pp[(ii + 0) * max_jj + jj]; + float f1 = pp[(ii + 1) * max_jj + jj]; + float f2 = pp[(ii + 2) * max_jj + jj]; + float f3 = pp[(ii + 3) * max_jj + jj]; + if (pC) + { + if (broadcast_type_C == 0) + { + f0 += c0; + f1 += c0; + f2 += c0; + f3 += c0; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + f0 += c0; + f1 += c1; + f2 += c2; + f3 += c3; + } + if (broadcast_type_C == 3) + { + if (beta == 1.f) + { + f0 += pC[0]; + f1 += pC[c_hstep]; + f2 += pC[c_hstep * 2]; + f3 += pC[c_hstep * 3]; + } + else + { + f0 += pC[0] * beta; + f1 += pC[c_hstep] * beta; + f2 += pC[c_hstep * 2] * beta; + f3 += pC[c_hstep * 3] * beta; + } + } + if (broadcast_type_C == 4) + { + float c = beta == 1.f ? pC[0] : pC[0] * beta; + f0 += c; + f1 += c; + f2 += c; + f3 += c; + } + } + if (alpha != 1.f) + { + f0 *= alpha; + f1 *= alpha; + f2 *= alpha; + f3 *= alpha; + } + p0[0] = f0; + p1[0] = f1; + p2[0] = f2; + p3[0] = f3; + p0++; + p1++; + p2++; + p3++; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC++; + } +#endif + } + for (; ii + 1 < max_ii; ii += 2) + { + float* p0 = outptr + (size_t)(i + ii) * N + j; + float* p1 = p0 + N; + + float c0 = 0.f; + float c1 = 0.f; + const float* pC = pC_base; + if (pC) + { + if (broadcast_type_C == 0) + { + c0 = pC[0] * beta; + c1 = c0; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + c0 = pC[i + ii] * beta; + c1 = pC[i + ii + 1] * beta; + } + if (broadcast_type_C == 3) + { + pC += (size_t)(i + ii) * c_hstep + j; + } + if (broadcast_type_C == 4) + { + pC += j; + } + } + + int jj = 0; +#if __loongarch_asx + const __m256 _alpha256 = (__m256)__lasx_xvreplfr2vr_s(alpha); + const __m256 _beta256 = (__m256)__lasx_xvreplfr2vr_s(beta); + for (; jj + 15 < max_jj; jj += 16) + { + __m256 _f00 = (__m256)__lasx_xvld(pp + (ii + 0) * max_jj + jj, 0); + __m256 _f01 = (__m256)__lasx_xvld(pp + (ii + 0) * max_jj + jj + 8, 0); + __m256 _f10 = (__m256)__lasx_xvld(pp + (ii + 1) * max_jj + jj, 0); + __m256 _f11 = (__m256)__lasx_xvld(pp + (ii + 1) * max_jj + jj + 8, 0); + if (pC) + { + if (broadcast_type_C == 0) + { + const __m256 _c = (__m256)__lasx_xvreplfr2vr_s(c0); + _f00 = __lasx_xvfadd_s(_f00, _c); + _f01 = __lasx_xvfadd_s(_f01, _c); + _f10 = __lasx_xvfadd_s(_f10, _c); + _f11 = __lasx_xvfadd_s(_f11, _c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + const __m256 _c0 = (__m256)__lasx_xvreplfr2vr_s(c0); + const __m256 _c1 = (__m256)__lasx_xvreplfr2vr_s(c1); + _f00 = __lasx_xvfadd_s(_f00, _c0); + _f01 = __lasx_xvfadd_s(_f01, _c0); + _f10 = __lasx_xvfadd_s(_f10, _c1); + _f11 = __lasx_xvfadd_s(_f11, _c1); + } + if (broadcast_type_C == 3) + { + const __m256 _c00 = (__m256)__lasx_xvld(pC, 0); + const __m256 _c01 = (__m256)__lasx_xvld(pC + 8, 0); + const __m256 _c10 = (__m256)__lasx_xvld(pC + c_hstep, 0); + const __m256 _c11 = (__m256)__lasx_xvld(pC + c_hstep + 8, 0); + if (beta == 1.f) + { + _f00 = __lasx_xvfadd_s(_f00, _c00); + _f01 = __lasx_xvfadd_s(_f01, _c01); + _f10 = __lasx_xvfadd_s(_f10, _c10); + _f11 = __lasx_xvfadd_s(_f11, _c11); + } + else + { + _f00 = __lasx_xvfmadd_s(_c00, _beta256, _f00); + _f01 = __lasx_xvfmadd_s(_c01, _beta256, _f01); + _f10 = __lasx_xvfmadd_s(_c10, _beta256, _f10); + _f11 = __lasx_xvfmadd_s(_c11, _beta256, _f11); + } + } + if (broadcast_type_C == 4) + { + __m256 _c0 = (__m256)__lasx_xvld(pC, 0); + __m256 _c1 = (__m256)__lasx_xvld(pC + 8, 0); + if (beta != 1.f) + { + _c0 = __lasx_xvfmul_s(_c0, _beta256); + _c1 = __lasx_xvfmul_s(_c1, _beta256); + } + _f00 = __lasx_xvfadd_s(_f00, _c0); + _f01 = __lasx_xvfadd_s(_f01, _c1); + _f10 = __lasx_xvfadd_s(_f10, _c0); + _f11 = __lasx_xvfadd_s(_f11, _c1); + } + } + if (alpha != 1.f) + { + _f00 = __lasx_xvfmul_s(_f00, _alpha256); + _f01 = __lasx_xvfmul_s(_f01, _alpha256); + _f10 = __lasx_xvfmul_s(_f10, _alpha256); + _f11 = __lasx_xvfmul_s(_f11, _alpha256); + } + __lasx_xvst(_f00, p0, 0); + __lasx_xvst(_f01, p0 + 8, 0); + __lasx_xvst(_f10, p1, 0); + __lasx_xvst(_f11, p1 + 8, 0); + p0 += 16; + p1 += 16; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 16; + } + for (; jj + 7 < max_jj; jj += 8) + { + __m256 _f0 = (__m256)__lasx_xvld(pp + (ii + 0) * max_jj + jj, 0); + __m256 _f1 = (__m256)__lasx_xvld(pp + (ii + 1) * max_jj + jj, 0); + if (pC) + { + if (broadcast_type_C == 0) + { + _f0 = __lasx_xvfadd_s(_f0, (__m256)__lasx_xvreplfr2vr_s(c0)); + _f1 = __lasx_xvfadd_s(_f1, (__m256)__lasx_xvreplfr2vr_s(c0)); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = __lasx_xvfadd_s(_f0, (__m256)__lasx_xvreplfr2vr_s(c0)); + _f1 = __lasx_xvfadd_s(_f1, (__m256)__lasx_xvreplfr2vr_s(c1)); + } + if (broadcast_type_C == 3) + { + __m256 _c0 = (__m256)__lasx_xvld(pC, 0); + __m256 _c1 = (__m256)__lasx_xvld(pC + c_hstep, 0); + if (beta == 1.f) + { + _f0 = __lasx_xvfadd_s(_f0, _c0); + _f1 = __lasx_xvfadd_s(_f1, _c1); + } + else + { + _f0 = __lasx_xvfmadd_s(_c0, _beta256, _f0); + _f1 = __lasx_xvfmadd_s(_c1, _beta256, _f1); + } + } + if (broadcast_type_C == 4) + { + __m256 _c = (__m256)__lasx_xvld(pC, 0); + if (beta != 1.f) + _c = __lasx_xvfmul_s(_c, _beta256); + _f0 = __lasx_xvfadd_s(_f0, _c); + _f1 = __lasx_xvfadd_s(_f1, _c); + } + } + if (alpha != 1.f) + { + _f0 = __lasx_xvfmul_s(_f0, _alpha256); + _f1 = __lasx_xvfmul_s(_f1, _alpha256); + } + __lasx_xvst(_f0, p0, 0); + __lasx_xvst(_f1, p1, 0); + p0 += 8; + p1 += 8; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 8; + } +#endif +#if __loongarch_sx + const __m128 _alpha128 = __lsx_vreplfr2vr_s(alpha); + const __m128 _beta128 = __lsx_vreplfr2vr_s(beta); + for (; jj + 7 < max_jj; jj += 8) + { + __m128 _f00 = (__m128)__lsx_vld(pp + (ii + 0) * max_jj + jj, 0); + __m128 _f01 = (__m128)__lsx_vld(pp + (ii + 0) * max_jj + jj + 4, 0); + __m128 _f10 = (__m128)__lsx_vld(pp + (ii + 1) * max_jj + jj, 0); + __m128 _f11 = (__m128)__lsx_vld(pp + (ii + 1) * max_jj + jj + 4, 0); + if (pC) + { + if (broadcast_type_C == 0) + { + const __m128 _c = __lsx_vreplfr2vr_s(c0); + _f00 = __lsx_vfadd_s(_f00, _c); + _f01 = __lsx_vfadd_s(_f01, _c); + _f10 = __lsx_vfadd_s(_f10, _c); + _f11 = __lsx_vfadd_s(_f11, _c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + const __m128 _c0 = __lsx_vreplfr2vr_s(c0); + const __m128 _c1 = __lsx_vreplfr2vr_s(c1); + _f00 = __lsx_vfadd_s(_f00, _c0); + _f01 = __lsx_vfadd_s(_f01, _c0); + _f10 = __lsx_vfadd_s(_f10, _c1); + _f11 = __lsx_vfadd_s(_f11, _c1); + } + if (broadcast_type_C == 3) + { + const __m128 _c00 = (__m128)__lsx_vld(pC, 0); + const __m128 _c01 = (__m128)__lsx_vld(pC + 4, 0); + const __m128 _c10 = (__m128)__lsx_vld(pC + c_hstep, 0); + const __m128 _c11 = (__m128)__lsx_vld(pC + c_hstep + 4, 0); + if (beta == 1.f) + { + _f00 = __lsx_vfadd_s(_f00, _c00); + _f01 = __lsx_vfadd_s(_f01, _c01); + _f10 = __lsx_vfadd_s(_f10, _c10); + _f11 = __lsx_vfadd_s(_f11, _c11); + } + else + { + _f00 = __lsx_vfmadd_s(_c00, _beta128, _f00); + _f01 = __lsx_vfmadd_s(_c01, _beta128, _f01); + _f10 = __lsx_vfmadd_s(_c10, _beta128, _f10); + _f11 = __lsx_vfmadd_s(_c11, _beta128, _f11); + } + } + if (broadcast_type_C == 4) + { + __m128 _c0 = (__m128)__lsx_vld(pC, 0); + __m128 _c1 = (__m128)__lsx_vld(pC + 4, 0); + if (beta != 1.f) + { + _c0 = __lsx_vfmul_s(_c0, _beta128); + _c1 = __lsx_vfmul_s(_c1, _beta128); + } + _f00 = __lsx_vfadd_s(_f00, _c0); + _f01 = __lsx_vfadd_s(_f01, _c1); + _f10 = __lsx_vfadd_s(_f10, _c0); + _f11 = __lsx_vfadd_s(_f11, _c1); + } + } + if (alpha != 1.f) + { + _f00 = __lsx_vfmul_s(_f00, _alpha128); + _f01 = __lsx_vfmul_s(_f01, _alpha128); + _f10 = __lsx_vfmul_s(_f10, _alpha128); + _f11 = __lsx_vfmul_s(_f11, _alpha128); + } + __lsx_vst((__m128i)_f00, p0, 0); + __lsx_vst((__m128i)_f01, p0 + 4, 0); + __lsx_vst((__m128i)_f10, p1, 0); + __lsx_vst((__m128i)_f11, p1 + 4, 0); + p0 += 8; + p1 += 8; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 8; + } + for (; jj + 3 < max_jj; jj += 4) + { + __m128 _f0 = (__m128)__lsx_vld(pp + (ii + 0) * max_jj + jj, 0); + __m128 _f1 = (__m128)__lsx_vld(pp + (ii + 1) * max_jj + jj, 0); + if (pC) + { + if (broadcast_type_C == 0) + { + _f0 = __lsx_vfadd_s(_f0, __lsx_vreplfr2vr_s(c0)); + _f1 = __lsx_vfadd_s(_f1, __lsx_vreplfr2vr_s(c0)); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = __lsx_vfadd_s(_f0, __lsx_vreplfr2vr_s(c0)); + _f1 = __lsx_vfadd_s(_f1, __lsx_vreplfr2vr_s(c1)); + } + if (broadcast_type_C == 3) + { + __m128 _c0 = (__m128)__lsx_vld(pC, 0); + __m128 _c1 = (__m128)__lsx_vld(pC + c_hstep, 0); + if (beta == 1.f) + { + _f0 = __lsx_vfadd_s(_f0, _c0); + _f1 = __lsx_vfadd_s(_f1, _c1); + } + else + { + _f0 = __lsx_vfmadd_s(_c0, _beta128, _f0); + _f1 = __lsx_vfmadd_s(_c1, _beta128, _f1); + } + } + if (broadcast_type_C == 4) + { + __m128 _c = (__m128)__lsx_vld(pC, 0); + if (beta != 1.f) + _c = __lsx_vfmul_s(_c, _beta128); + _f0 = __lsx_vfadd_s(_f0, _c); + _f1 = __lsx_vfadd_s(_f1, _c); + } + } + if (alpha != 1.f) + { + _f0 = __lsx_vfmul_s(_f0, _alpha128); + _f1 = __lsx_vfmul_s(_f1, _alpha128); + } + __lsx_vst((__m128i)_f0, p0, 0); + __lsx_vst((__m128i)_f1, p1, 0); + p0 += 4; + p1 += 4; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 4; + } + for (; jj + 1 < max_jj; jj += 2) + { + __m128 _f0 = (__m128)__lsx_vldrepl_d(pp + (ii + 0) * max_jj + jj, 0); + __m128 _f1 = (__m128)__lsx_vldrepl_d(pp + (ii + 1) * max_jj + jj, 0); + if (pC) + { + if (broadcast_type_C == 0) + { + __m128 _c = __lsx_vreplfr2vr_s(c0); + _f0 = __lsx_vfadd_s(_f0, _c); + _f1 = __lsx_vfadd_s(_f1, _c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = __lsx_vfadd_s(_f0, __lsx_vreplfr2vr_s(c0)); + _f1 = __lsx_vfadd_s(_f1, __lsx_vreplfr2vr_s(c1)); + } + if (broadcast_type_C == 3) + { + __m128 _c0 = (__m128)__lsx_vldrepl_d(pC, 0); + __m128 _c1 = (__m128)__lsx_vldrepl_d(pC + c_hstep, 0); + if (beta == 1.f) + { + _f0 = __lsx_vfadd_s(_f0, _c0); + _f1 = __lsx_vfadd_s(_f1, _c1); + } + else + { + _f0 = __lsx_vfmadd_s(_c0, _beta128, _f0); + _f1 = __lsx_vfmadd_s(_c1, _beta128, _f1); + } + } + if (broadcast_type_C == 4) + { + __m128 _c = (__m128)__lsx_vldrepl_d(pC, 0); + if (beta != 1.f) + _c = __lsx_vfmul_s(_c, _beta128); + _f0 = __lsx_vfadd_s(_f0, _c); + _f1 = __lsx_vfadd_s(_f1, _c); + } + } + if (alpha != 1.f) + { + _f0 = __lsx_vfmul_s(_f0, _alpha128); + _f1 = __lsx_vfmul_s(_f1, _alpha128); + } + __lsx_vstelm_d((__m128i)_f0, p0, 0, 0); + __lsx_vstelm_d((__m128i)_f1, p1, 0, 0); + p0 += 2; + p1 += 2; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 2; + } +#endif + for (; jj < max_jj; jj++) + { + float f0 = pp[(ii + 0) * max_jj + jj]; + float f1 = pp[(ii + 1) * max_jj + jj]; + if (pC) + { + if (broadcast_type_C == 0) + { + f0 += c0; + f1 += c0; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + f0 += c0; + f1 += c1; + } + if (broadcast_type_C == 3) + { + if (beta == 1.f) + { + f0 += pC[0]; + f1 += pC[c_hstep]; + } + else + { + f0 += pC[0] * beta; + f1 += pC[c_hstep] * beta; + } + } + if (broadcast_type_C == 4) + { + float c = beta == 1.f ? pC[0] : pC[0] * beta; + f0 += c; + f1 += c; + } + } + if (alpha != 1.f) + { + f0 *= alpha; + f1 *= alpha; + } + p0[0] = f0; + p1[0] = f1; + p0++; + p1++; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC++; + } + } + for (; ii < max_ii; ii++) + { + float* p0 = outptr + (size_t)(i + ii) * N + j; + + float c0 = 0.f; + const float* pC = pC_base; + if (pC) + { + if (broadcast_type_C == 0) + c0 = pC[0] * beta; + if (broadcast_type_C == 1 || broadcast_type_C == 2) + c0 = pC[i + ii] * beta; + if (broadcast_type_C == 3) + pC += (size_t)(i + ii) * c_hstep + j; + if (broadcast_type_C == 4) + pC += j; + } + + int jj = 0; +#if __loongarch_asx + const __m256 _alpha256 = (__m256)__lasx_xvreplfr2vr_s(alpha); + const __m256 _beta256 = (__m256)__lasx_xvreplfr2vr_s(beta); + for (; jj + 15 < max_jj; jj += 16) + { + __m256 _f0 = (__m256)__lasx_xvld(pp + ii * max_jj + jj, 0); + __m256 _f1 = (__m256)__lasx_xvld(pp + ii * max_jj + jj + 8, 0); + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + const __m256 _c = (__m256)__lasx_xvreplfr2vr_s(c0); + _f0 = __lasx_xvfadd_s(_f0, _c); + _f1 = __lasx_xvfadd_s(_f1, _c); + } + if (broadcast_type_C == 3 || broadcast_type_C == 4) + { + const __m256 _c0 = (__m256)__lasx_xvld(pC, 0); + const __m256 _c1 = (__m256)__lasx_xvld(pC + 8, 0); + if (beta == 1.f) + { + _f0 = __lasx_xvfadd_s(_f0, _c0); + _f1 = __lasx_xvfadd_s(_f1, _c1); + } + else + { + _f0 = __lasx_xvfmadd_s(_c0, _beta256, _f0); + _f1 = __lasx_xvfmadd_s(_c1, _beta256, _f1); + } + } + } + if (alpha != 1.f) + { + _f0 = __lasx_xvfmul_s(_f0, _alpha256); + _f1 = __lasx_xvfmul_s(_f1, _alpha256); + } + __lasx_xvst(_f0, p0, 0); + __lasx_xvst(_f1, p0 + 8, 0); + p0 += 16; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 16; + } + for (; jj + 7 < max_jj; jj += 8) + { + __m256 _f0 = (__m256)__lasx_xvld(pp + ii * max_jj + jj, 0); + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + _f0 = __lasx_xvfadd_s(_f0, (__m256)__lasx_xvreplfr2vr_s(c0)); + if (broadcast_type_C == 3 || broadcast_type_C == 4) + { + __m256 _c0 = (__m256)__lasx_xvld(pC, 0); + if (beta == 1.f) + _f0 = __lasx_xvfadd_s(_f0, _c0); + else + _f0 = __lasx_xvfmadd_s(_c0, _beta256, _f0); + } + } + if (alpha != 1.f) _f0 = __lasx_xvfmul_s(_f0, _alpha256); + __lasx_xvst(_f0, p0, 0); + p0 += 8; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 8; + } +#endif +#if __loongarch_sx + const __m128 _alpha128 = __lsx_vreplfr2vr_s(alpha); + const __m128 _beta128 = __lsx_vreplfr2vr_s(beta); + for (; jj + 7 < max_jj; jj += 8) + { + __m128 _f0 = (__m128)__lsx_vld(pp + ii * max_jj + jj, 0); + __m128 _f1 = (__m128)__lsx_vld(pp + ii * max_jj + jj + 4, 0); + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + const __m128 _c = __lsx_vreplfr2vr_s(c0); + _f0 = __lsx_vfadd_s(_f0, _c); + _f1 = __lsx_vfadd_s(_f1, _c); + } + if (broadcast_type_C == 3 || broadcast_type_C == 4) + { + const __m128 _c0 = (__m128)__lsx_vld(pC, 0); + const __m128 _c1 = (__m128)__lsx_vld(pC + 4, 0); + if (beta == 1.f) + { + _f0 = __lsx_vfadd_s(_f0, _c0); + _f1 = __lsx_vfadd_s(_f1, _c1); + } + else + { + _f0 = __lsx_vfmadd_s(_c0, _beta128, _f0); + _f1 = __lsx_vfmadd_s(_c1, _beta128, _f1); + } + } + } + if (alpha != 1.f) + { + _f0 = __lsx_vfmul_s(_f0, _alpha128); + _f1 = __lsx_vfmul_s(_f1, _alpha128); + } + __lsx_vst((__m128i)_f0, p0, 0); + __lsx_vst((__m128i)_f1, p0 + 4, 0); + p0 += 8; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 8; + } + for (; jj + 3 < max_jj; jj += 4) + { + __m128 _f0 = (__m128)__lsx_vld(pp + ii * max_jj + jj, 0); + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + _f0 = __lsx_vfadd_s(_f0, __lsx_vreplfr2vr_s(c0)); + if (broadcast_type_C == 3 || broadcast_type_C == 4) + { + __m128 _c0 = (__m128)__lsx_vld(pC, 0); + if (beta == 1.f) + _f0 = __lsx_vfadd_s(_f0, _c0); + else + _f0 = __lsx_vfmadd_s(_c0, _beta128, _f0); + } + } + if (alpha != 1.f) _f0 = __lsx_vfmul_s(_f0, _alpha128); + __lsx_vst((__m128i)_f0, p0, 0); + p0 += 4; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 4; + } + for (; jj + 1 < max_jj; jj += 2) + { + __m128 _f0 = (__m128)__lsx_vldrepl_d(pp + ii * max_jj + jj, 0); + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + _f0 = __lsx_vfadd_s(_f0, __lsx_vreplfr2vr_s(c0)); + if (broadcast_type_C == 3 || broadcast_type_C == 4) + { + __m128 _c0 = (__m128)__lsx_vldrepl_d(pC, 0); + if (beta == 1.f) + _f0 = __lsx_vfadd_s(_f0, _c0); + else + _f0 = __lsx_vfmadd_s(_c0, _beta128, _f0); + } + } + if (alpha != 1.f) _f0 = __lsx_vfmul_s(_f0, _alpha128); + __lsx_vstelm_d((__m128i)_f0, p0, 0, 0); + p0 += 2; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 2; + } +#endif + for (; jj < max_jj; jj++) + { + float f0 = pp[ii * max_jj + jj]; + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + f0 += c0; + if (broadcast_type_C == 3 || broadcast_type_C == 4) + f0 += beta == 1.f ? pC[0] : pC[0] * beta; + } + if (alpha != 1.f) + f0 *= alpha; + p0[0] = f0; + p0++; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC++; + } + } +} + +static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int M, float alpha, float beta) +{ + const float* pp = topT; + const size_t c_hstep = C.dims == 3 ? C.cstep : (size_t)C.w; + const float* pC_base = C; + float* outptr = top_blob; + int ii = 0; +#if __loongarch_sx + for (; ii + 7 < max_ii; ii += 8) + { + float* p0 = outptr + (size_t)j * M + i + ii; + const float* pC = pC_base; + + __m128 _c0 = __lsx_vreplfr2vr_s(0.f); + __m128 _c1 = __lsx_vreplfr2vr_s(0.f); + if (pC && broadcast_type_C == 0) + { + _c0 = __lsx_vreplfr2vr_s(pC[0] * beta); + _c1 = _c0; + } + if (pC && (broadcast_type_C == 1 || broadcast_type_C == 2)) + { + _c0 = (__m128)__lsx_vld(pC + i + ii, 0); + _c1 = (__m128)__lsx_vld(pC + i + ii + 4, 0); + if (beta != 1.f) + { + const __m128 _beta = __lsx_vreplfr2vr_s(beta); + _c0 = __lsx_vfmul_s(_c0, _beta); + _c1 = __lsx_vfmul_s(_c1, _beta); + } + } + + if (pC && broadcast_type_C == 3) + pC += (size_t)(i + ii) * c_hstep + j; + if (pC && broadcast_type_C == 4) + pC += j; + + const __m128 _alpha = __lsx_vreplfr2vr_s(alpha); + const __m128 _beta = __lsx_vreplfr2vr_s(beta); + int jj = 0; +#if __loongarch_asx + const __m256 _alpha256 = (__m256)__lasx_xvreplfr2vr_s(alpha); + const __m256 _beta256 = (__m256)__lasx_xvreplfr2vr_s(beta); + const __m256 _c256 = __lasx_concat_128_s(_c0, _c1); + for (; jj + 7 < max_jj; jj += 8) + { + __m256i _sum0 = __lasx_xvld(pp, 0); + __m256i _sum1 = __lasx_xvld(pp + 8, 0); + __m256i _sum2 = __lasx_xvld(pp + 16, 0); + __m256i _sum3 = __lasx_xvld(pp + 24, 0); + __m256i _sum4 = __lasx_xvld(pp + 32, 0); + __m256i _sum5 = __lasx_xvld(pp + 40, 0); + __m256i _sum6 = __lasx_xvld(pp + 48, 0); + __m256i _sum7 = __lasx_xvld(pp + 56, 0); + __m256i _tmp0 = _sum0; + __m256i _tmp1 = __lasx_xvshuf4i_w(_sum1, _LSX_SHUFFLE(2, 1, 0, 3)); + __m256i _tmp2 = _sum2; + __m256i _tmp3 = __lasx_xvshuf4i_w(_sum3, _LSX_SHUFFLE(2, 1, 0, 3)); + __m256i _tmp4 = _sum4; + __m256i _tmp5 = __lasx_xvshuf4i_w(_sum5, _LSX_SHUFFLE(2, 1, 0, 3)); + __m256i _tmp6 = _sum6; + __m256i _tmp7 = __lasx_xvshuf4i_w(_sum7, _LSX_SHUFFLE(2, 1, 0, 3)); + _sum0 = __lasx_xvilvl_w(_tmp3, _tmp0); + _sum1 = __lasx_xvilvh_w(_tmp3, _tmp0); + _sum2 = __lasx_xvilvl_w(_tmp1, _tmp2); + _sum3 = __lasx_xvilvh_w(_tmp1, _tmp2); + _sum4 = __lasx_xvilvl_w(_tmp7, _tmp4); + _sum5 = __lasx_xvilvh_w(_tmp7, _tmp4); + _sum6 = __lasx_xvilvl_w(_tmp5, _tmp6); + _sum7 = __lasx_xvilvh_w(_tmp5, _tmp6); + _tmp0 = __lasx_xvilvl_d(_sum2, _sum0); + _tmp1 = __lasx_xvilvh_d(_sum2, _sum0); + _tmp2 = __lasx_xvilvl_d(_sum1, _sum3); + _tmp3 = __lasx_xvilvh_d(_sum1, _sum3); + _tmp4 = __lasx_xvilvl_d(_sum6, _sum4); + _tmp5 = __lasx_xvilvh_d(_sum6, _sum4); + _tmp6 = __lasx_xvilvl_d(_sum5, _sum7); + _tmp7 = __lasx_xvilvh_d(_sum5, _sum7); + _tmp1 = __lasx_xvshuf4i_w(_tmp1, _LSX_SHUFFLE(2, 1, 0, 3)); + _tmp3 = __lasx_xvshuf4i_w(_tmp3, _LSX_SHUFFLE(2, 1, 0, 3)); + _tmp5 = __lasx_xvshuf4i_w(_tmp5, _LSX_SHUFFLE(2, 1, 0, 3)); + _tmp7 = __lasx_xvshuf4i_w(_tmp7, _LSX_SHUFFLE(2, 1, 0, 3)); + __m256 _f0 = (__m256)__lasx_xvpermi_q(_tmp4, _tmp0, _LSX_SHUFFLE(0, 3, 0, 0)); + __m256 _f1 = (__m256)__lasx_xvpermi_q(_tmp5, _tmp1, _LSX_SHUFFLE(0, 3, 0, 0)); + __m256 _f2 = (__m256)__lasx_xvpermi_q(_tmp6, _tmp2, _LSX_SHUFFLE(0, 3, 0, 0)); + __m256 _f3 = (__m256)__lasx_xvpermi_q(_tmp7, _tmp3, _LSX_SHUFFLE(0, 3, 0, 0)); + __m256 _f4 = (__m256)__lasx_xvpermi_q(_tmp0, _tmp4, _LSX_SHUFFLE(0, 3, 0, 0)); + __m256 _f5 = (__m256)__lasx_xvpermi_q(_tmp1, _tmp5, _LSX_SHUFFLE(0, 3, 0, 0)); + __m256 _f6 = (__m256)__lasx_xvpermi_q(_tmp2, _tmp6, _LSX_SHUFFLE(0, 3, 0, 0)); + __m256 _f7 = (__m256)__lasx_xvpermi_q(_tmp3, _tmp7, _LSX_SHUFFLE(0, 3, 0, 0)); + pp += 64; + transpose8x8_ps(_f0, _f1, _f2, _f3, _f4, _f5, _f6, _f7); + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = __lasx_xvfadd_s(_f0, _c256); + _f1 = __lasx_xvfadd_s(_f1, _c256); + _f2 = __lasx_xvfadd_s(_f2, _c256); + _f3 = __lasx_xvfadd_s(_f3, _c256); + _f4 = __lasx_xvfadd_s(_f4, _c256); + _f5 = __lasx_xvfadd_s(_f5, _c256); + _f6 = __lasx_xvfadd_s(_f6, _c256); + _f7 = __lasx_xvfadd_s(_f7, _c256); + } + if (broadcast_type_C == 3) + { + __m256 _cc0 = (__m256)__lasx_xvld(pC, 0); + __m256 _cc1 = (__m256)__lasx_xvld(pC + c_hstep, 0); + __m256 _cc2 = (__m256)__lasx_xvld(pC + c_hstep * 2, 0); + __m256 _cc3 = (__m256)__lasx_xvld(pC + c_hstep * 3, 0); + __m256 _cc4 = (__m256)__lasx_xvld(pC + c_hstep * 4, 0); + __m256 _cc5 = (__m256)__lasx_xvld(pC + c_hstep * 5, 0); + __m256 _cc6 = (__m256)__lasx_xvld(pC + c_hstep * 6, 0); + __m256 _cc7 = (__m256)__lasx_xvld(pC + c_hstep * 7, 0); + transpose8x8_ps(_cc0, _cc1, _cc2, _cc3, _cc4, _cc5, _cc6, _cc7); + if (beta == 1.f) + { + _f0 = __lasx_xvfadd_s(_f0, _cc0); + _f1 = __lasx_xvfadd_s(_f1, _cc1); + _f2 = __lasx_xvfadd_s(_f2, _cc2); + _f3 = __lasx_xvfadd_s(_f3, _cc3); + _f4 = __lasx_xvfadd_s(_f4, _cc4); + _f5 = __lasx_xvfadd_s(_f5, _cc5); + _f6 = __lasx_xvfadd_s(_f6, _cc6); + _f7 = __lasx_xvfadd_s(_f7, _cc7); + } + else + { + _f0 = __lasx_xvfmadd_s(_cc0, _beta256, _f0); + _f1 = __lasx_xvfmadd_s(_cc1, _beta256, _f1); + _f2 = __lasx_xvfmadd_s(_cc2, _beta256, _f2); + _f3 = __lasx_xvfmadd_s(_cc3, _beta256, _f3); + _f4 = __lasx_xvfmadd_s(_cc4, _beta256, _f4); + _f5 = __lasx_xvfmadd_s(_cc5, _beta256, _f5); + _f6 = __lasx_xvfmadd_s(_cc6, _beta256, _f6); + _f7 = __lasx_xvfmadd_s(_cc7, _beta256, _f7); + } + } + if (broadcast_type_C == 4) + { + _f0 = __lasx_xvfadd_s(_f0, (__m256)__lasx_xvreplfr2vr_s(pC[0] * beta)); + _f1 = __lasx_xvfadd_s(_f1, (__m256)__lasx_xvreplfr2vr_s(pC[1] * beta)); + _f2 = __lasx_xvfadd_s(_f2, (__m256)__lasx_xvreplfr2vr_s(pC[2] * beta)); + _f3 = __lasx_xvfadd_s(_f3, (__m256)__lasx_xvreplfr2vr_s(pC[3] * beta)); + _f4 = __lasx_xvfadd_s(_f4, (__m256)__lasx_xvreplfr2vr_s(pC[4] * beta)); + _f5 = __lasx_xvfadd_s(_f5, (__m256)__lasx_xvreplfr2vr_s(pC[5] * beta)); + _f6 = __lasx_xvfadd_s(_f6, (__m256)__lasx_xvreplfr2vr_s(pC[6] * beta)); + _f7 = __lasx_xvfadd_s(_f7, (__m256)__lasx_xvreplfr2vr_s(pC[7] * beta)); + } + } + if (alpha != 1.f) + { + _f0 = __lasx_xvfmul_s(_f0, _alpha256); + _f1 = __lasx_xvfmul_s(_f1, _alpha256); + _f2 = __lasx_xvfmul_s(_f2, _alpha256); + _f3 = __lasx_xvfmul_s(_f3, _alpha256); + _f4 = __lasx_xvfmul_s(_f4, _alpha256); + _f5 = __lasx_xvfmul_s(_f5, _alpha256); + _f6 = __lasx_xvfmul_s(_f6, _alpha256); + _f7 = __lasx_xvfmul_s(_f7, _alpha256); + } + __lasx_xvst(_f0, p0, 0); + __lasx_xvst(_f1, p0 + M, 0); + __lasx_xvst(_f2, p0 + M * 2, 0); + __lasx_xvst(_f3, p0 + M * 3, 0); + __lasx_xvst(_f4, p0 + M * 4, 0); + __lasx_xvst(_f5, p0 + M * 5, 0); + __lasx_xvst(_f6, p0 + M * 6, 0); + __lasx_xvst(_f7, p0 + M * 7, 0); + p0 += M * 8; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 8; + } +#endif + for (; jj + 3 < max_jj; jj += 4) + { + __m128i _sum0 = __lsx_vld(pp, 0); + __m128i _sum1 = __lsx_vld(pp + 8, 0); + __m128i _sum2 = __lsx_vld(pp + 16, 0); + __m128i _sum3 = __lsx_vld(pp + 24, 0); + _sum2 = __lsx_vshuf4i_w(_sum2, _LSX_SHUFFLE(1, 0, 3, 2)); + _sum3 = __lsx_vshuf4i_w(_sum3, _LSX_SHUFFLE(1, 0, 3, 2)); + transpose4x4_epi32(_sum0, _sum1, _sum2, _sum3); + _sum1 = __lsx_vshuf4i_w(_sum1, _LSX_SHUFFLE(2, 1, 0, 3)); + _sum2 = __lsx_vshuf4i_w(_sum2, _LSX_SHUFFLE(1, 0, 3, 2)); + _sum3 = __lsx_vshuf4i_w(_sum3, _LSX_SHUFFLE(0, 3, 2, 1)); + __m128i _sum4 = __lsx_vld(pp + 4, 0); + __m128i _sum5 = __lsx_vld(pp + 12, 0); + __m128i _sum6 = __lsx_vld(pp + 20, 0); + __m128i _sum7 = __lsx_vld(pp + 28, 0); + _sum6 = __lsx_vshuf4i_w(_sum6, _LSX_SHUFFLE(1, 0, 3, 2)); + _sum7 = __lsx_vshuf4i_w(_sum7, _LSX_SHUFFLE(1, 0, 3, 2)); + transpose4x4_epi32(_sum4, _sum5, _sum6, _sum7); + _sum5 = __lsx_vshuf4i_w(_sum5, _LSX_SHUFFLE(2, 1, 0, 3)); + _sum6 = __lsx_vshuf4i_w(_sum6, _LSX_SHUFFLE(1, 0, 3, 2)); + _sum7 = __lsx_vshuf4i_w(_sum7, _LSX_SHUFFLE(0, 3, 2, 1)); + __m128 _f0 = (__m128)_sum0; + __m128 _f1 = (__m128)_sum1; + __m128 _f2 = (__m128)_sum2; + __m128 _f3 = (__m128)_sum3; + __m128 _f4 = (__m128)_sum4; + __m128 _f5 = (__m128)_sum5; + __m128 _f6 = (__m128)_sum6; + __m128 _f7 = (__m128)_sum7; + pp += 32; + transpose4x4_ps(_f0, _f1, _f2, _f3); + transpose4x4_ps(_f4, _f5, _f6, _f7); + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = __lsx_vfadd_s(_f0, _c0); + _f1 = __lsx_vfadd_s(_f1, _c0); + _f2 = __lsx_vfadd_s(_f2, _c0); + _f3 = __lsx_vfadd_s(_f3, _c0); + _f4 = __lsx_vfadd_s(_f4, _c1); + _f5 = __lsx_vfadd_s(_f5, _c1); + _f6 = __lsx_vfadd_s(_f6, _c1); + _f7 = __lsx_vfadd_s(_f7, _c1); + } + if (broadcast_type_C == 3) + { + __m128 _cc0 = (__m128)__lsx_vld(pC, 0); + __m128 _cc1 = (__m128)__lsx_vld(pC + c_hstep, 0); + __m128 _cc2 = (__m128)__lsx_vld(pC + c_hstep * 2, 0); + __m128 _cc3 = (__m128)__lsx_vld(pC + c_hstep * 3, 0); + __m128 _cc4 = (__m128)__lsx_vld(pC + c_hstep * 4, 0); + __m128 _cc5 = (__m128)__lsx_vld(pC + c_hstep * 5, 0); + __m128 _cc6 = (__m128)__lsx_vld(pC + c_hstep * 6, 0); + __m128 _cc7 = (__m128)__lsx_vld(pC + c_hstep * 7, 0); + transpose4x4_ps(_cc0, _cc1, _cc2, _cc3); + transpose4x4_ps(_cc4, _cc5, _cc6, _cc7); + if (beta == 1.f) + { + _f0 = __lsx_vfadd_s(_f0, _cc0); + _f1 = __lsx_vfadd_s(_f1, _cc1); + _f2 = __lsx_vfadd_s(_f2, _cc2); + _f3 = __lsx_vfadd_s(_f3, _cc3); + _f4 = __lsx_vfadd_s(_f4, _cc4); + _f5 = __lsx_vfadd_s(_f5, _cc5); + _f6 = __lsx_vfadd_s(_f6, _cc6); + _f7 = __lsx_vfadd_s(_f7, _cc7); + } + else + { + _f0 = __lsx_vfmadd_s(_cc0, _beta, _f0); + _f1 = __lsx_vfmadd_s(_cc1, _beta, _f1); + _f2 = __lsx_vfmadd_s(_cc2, _beta, _f2); + _f3 = __lsx_vfmadd_s(_cc3, _beta, _f3); + _f4 = __lsx_vfmadd_s(_cc4, _beta, _f4); + _f5 = __lsx_vfmadd_s(_cc5, _beta, _f5); + _f6 = __lsx_vfmadd_s(_cc6, _beta, _f6); + _f7 = __lsx_vfmadd_s(_cc7, _beta, _f7); + } + } + if (broadcast_type_C == 4) + { + __m128 _cc = (__m128)__lsx_vld(pC, 0); + if (beta != 1.f) + _cc = __lsx_vfmul_s(_cc, _beta); + _f0 = __lsx_vfadd_s(_f0, (__m128)__lsx_vreplvei_w((__m128i)_cc, 0)); + _f1 = __lsx_vfadd_s(_f1, (__m128)__lsx_vreplvei_w((__m128i)_cc, 1)); + _f2 = __lsx_vfadd_s(_f2, (__m128)__lsx_vreplvei_w((__m128i)_cc, 2)); + _f3 = __lsx_vfadd_s(_f3, (__m128)__lsx_vreplvei_w((__m128i)_cc, 3)); + _f4 = __lsx_vfadd_s(_f4, (__m128)__lsx_vreplvei_w((__m128i)_cc, 0)); + _f5 = __lsx_vfadd_s(_f5, (__m128)__lsx_vreplvei_w((__m128i)_cc, 1)); + _f6 = __lsx_vfadd_s(_f6, (__m128)__lsx_vreplvei_w((__m128i)_cc, 2)); + _f7 = __lsx_vfadd_s(_f7, (__m128)__lsx_vreplvei_w((__m128i)_cc, 3)); + } + } + if (alpha != 1.f) + { + _f0 = __lsx_vfmul_s(_f0, _alpha); + _f1 = __lsx_vfmul_s(_f1, _alpha); + _f2 = __lsx_vfmul_s(_f2, _alpha); + _f3 = __lsx_vfmul_s(_f3, _alpha); + _f4 = __lsx_vfmul_s(_f4, _alpha); + _f5 = __lsx_vfmul_s(_f5, _alpha); + _f6 = __lsx_vfmul_s(_f6, _alpha); + _f7 = __lsx_vfmul_s(_f7, _alpha); + } + __lsx_vst((__m128i)_f0, p0, 0); + __lsx_vst((__m128i)_f4, p0 + 4, 0); + __lsx_vst((__m128i)_f1, p0 + M, 0); + __lsx_vst((__m128i)_f5, p0 + M + 4, 0); + __lsx_vst((__m128i)_f2, p0 + M * 2, 0); + __lsx_vst((__m128i)_f6, p0 + M * 2 + 4, 0); + __lsx_vst((__m128i)_f3, p0 + M * 3, 0); + __lsx_vst((__m128i)_f7, p0 + M * 3 + 4, 0); + p0 += M * 4; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 4; + } + for (; jj + 1 < max_jj; jj += 2) + { + __m128 _f0 = (__m128)__lsx_vld(pp, 0); + __m128 _f1 = (__m128)__lsx_vld(pp + 4, 0); + __m128 _f2 = (__m128)__lsx_vld(pp + 8, 0); + __m128 _f3 = (__m128)__lsx_vld(pp + 12, 0); + pp += 16; + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = __lsx_vfadd_s(_f0, _c0); + _f1 = __lsx_vfadd_s(_f1, _c1); + _f2 = __lsx_vfadd_s(_f2, _c0); + _f3 = __lsx_vfadd_s(_f3, _c1); + } + if (broadcast_type_C == 3) + { + __m128i _ci0 = __lsx_vreplgr2vr_w(((const int*)pC)[0]); + _ci0 = __lsx_vinsgr2vr_w(_ci0, ((const int*)(pC + c_hstep))[0], 1); + _ci0 = __lsx_vinsgr2vr_w(_ci0, ((const int*)(pC + c_hstep * 2))[0], 2); + _ci0 = __lsx_vinsgr2vr_w(_ci0, ((const int*)(pC + c_hstep * 3))[0], 3); + __m128i _ci1 = __lsx_vreplgr2vr_w(((const int*)(pC + c_hstep * 4))[0]); + _ci1 = __lsx_vinsgr2vr_w(_ci1, ((const int*)(pC + c_hstep * 5))[0], 1); + _ci1 = __lsx_vinsgr2vr_w(_ci1, ((const int*)(pC + c_hstep * 6))[0], 2); + _ci1 = __lsx_vinsgr2vr_w(_ci1, ((const int*)(pC + c_hstep * 7))[0], 3); + __m128i _ci2 = __lsx_vreplgr2vr_w(((const int*)pC)[1]); + _ci2 = __lsx_vinsgr2vr_w(_ci2, ((const int*)(pC + c_hstep))[1], 1); + _ci2 = __lsx_vinsgr2vr_w(_ci2, ((const int*)(pC + c_hstep * 2))[1], 2); + _ci2 = __lsx_vinsgr2vr_w(_ci2, ((const int*)(pC + c_hstep * 3))[1], 3); + __m128i _ci3 = __lsx_vreplgr2vr_w(((const int*)(pC + c_hstep * 4))[1]); + _ci3 = __lsx_vinsgr2vr_w(_ci3, ((const int*)(pC + c_hstep * 5))[1], 1); + _ci3 = __lsx_vinsgr2vr_w(_ci3, ((const int*)(pC + c_hstep * 6))[1], 2); + _ci3 = __lsx_vinsgr2vr_w(_ci3, ((const int*)(pC + c_hstep * 7))[1], 3); + if (beta == 1.f) + { + _f0 = __lsx_vfadd_s(_f0, (__m128)_ci0); + _f1 = __lsx_vfadd_s(_f1, (__m128)_ci1); + _f2 = __lsx_vfadd_s(_f2, (__m128)_ci2); + _f3 = __lsx_vfadd_s(_f3, (__m128)_ci3); + } + else + { + _f0 = __lsx_vfmadd_s((__m128)_ci0, _beta, _f0); + _f1 = __lsx_vfmadd_s((__m128)_ci1, _beta, _f1); + _f2 = __lsx_vfmadd_s((__m128)_ci2, _beta, _f2); + _f3 = __lsx_vfmadd_s((__m128)_ci3, _beta, _f3); + } + } + if (broadcast_type_C == 4) + { + const __m128 _cc0 = __lsx_vreplfr2vr_s(pC[0] * beta); + const __m128 _cc1 = __lsx_vreplfr2vr_s(pC[1] * beta); + _f0 = __lsx_vfadd_s(_f0, _cc0); + _f1 = __lsx_vfadd_s(_f1, _cc0); + _f2 = __lsx_vfadd_s(_f2, _cc1); + _f3 = __lsx_vfadd_s(_f3, _cc1); + } + } + if (alpha != 1.f) + { + _f0 = __lsx_vfmul_s(_f0, _alpha); + _f1 = __lsx_vfmul_s(_f1, _alpha); + _f2 = __lsx_vfmul_s(_f2, _alpha); + _f3 = __lsx_vfmul_s(_f3, _alpha); + } + __lsx_vst((__m128i)_f0, p0, 0); + __lsx_vst((__m128i)_f1, p0 + 4, 0); + __lsx_vst((__m128i)_f2, p0 + M, 0); + __lsx_vst((__m128i)_f3, p0 + M + 4, 0); + p0 += M * 2; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 2; + } + for (; jj < max_jj; jj++) + { + __m128i _fi0 = __lsx_vld(pp, 0); + __m128i _fi1 = __lsx_vld(pp + 4, 0); + pp += 8; + __m128 _f0 = (__m128)_fi0; + __m128 _f1 = (__m128)_fi1; + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = __lsx_vfadd_s(_f0, _c0); + _f1 = __lsx_vfadd_s(_f1, _c1); + } + if (broadcast_type_C == 3) + { + __m128i _ci0 = __lsx_vldrepl_w(pC, 0); + _ci0 = __lsx_vinsgr2vr_w(_ci0, ((const int*)(pC + c_hstep))[0], 1); + _ci0 = __lsx_vinsgr2vr_w(_ci0, ((const int*)(pC + c_hstep * 2))[0], 2); + _ci0 = __lsx_vinsgr2vr_w(_ci0, ((const int*)(pC + c_hstep * 3))[0], 3); + __m128i _ci1 = __lsx_vldrepl_w(pC + c_hstep * 4, 0); + _ci1 = __lsx_vinsgr2vr_w(_ci1, ((const int*)(pC + c_hstep * 5))[0], 1); + _ci1 = __lsx_vinsgr2vr_w(_ci1, ((const int*)(pC + c_hstep * 6))[0], 2); + _ci1 = __lsx_vinsgr2vr_w(_ci1, ((const int*)(pC + c_hstep * 7))[0], 3); + if (beta == 1.f) + { + _f0 = __lsx_vfadd_s(_f0, (__m128)_ci0); + _f1 = __lsx_vfadd_s(_f1, (__m128)_ci1); + } + else + { + _f0 = __lsx_vfmadd_s((__m128)_ci0, _beta, _f0); + _f1 = __lsx_vfmadd_s((__m128)_ci1, _beta, _f1); + } + } + if (broadcast_type_C == 4) + { + const __m128 _cc = __lsx_vreplfr2vr_s(pC[0] * beta); + _f0 = __lsx_vfadd_s(_f0, _cc); + _f1 = __lsx_vfadd_s(_f1, _cc); + } + } + if (alpha != 1.f) + { + _f0 = __lsx_vfmul_s(_f0, _alpha); + _f1 = __lsx_vfmul_s(_f1, _alpha); + } + __lsx_vst((__m128i)_f0, p0, 0); + __lsx_vst((__m128i)_f1, p0 + 4, 0); + p0 += M; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC++; + } + } + for (; ii + 3 < max_ii; ii += 4) + { + const float* pp0 = pp + (ii + 0) * max_jj; + const float* pp1 = pp + (ii + 1) * max_jj; + const float* pp2 = pp + (ii + 2) * max_jj; + const float* pp3 = pp + (ii + 3) * max_jj; + float* p0 = outptr + (size_t)j * M + i + ii; + const float* pC = pC_base; + + __m128 _c = __lsx_vreplfr2vr_s(0.f); + if (pC && broadcast_type_C == 0) + _c = __lsx_vreplfr2vr_s(pC[0] * beta); + if (pC && (broadcast_type_C == 1 || broadcast_type_C == 2)) + { + _c = (__m128)__lsx_vld(pC + i + ii, 0); + if (beta != 1.f) + _c = __lsx_vfmul_s(_c, __lsx_vreplfr2vr_s(beta)); + } + + if (pC && broadcast_type_C == 3) + pC += (size_t)(i + ii) * c_hstep + j; + if (pC && broadcast_type_C == 4) + pC += j; + + const __m128 _alpha = __lsx_vreplfr2vr_s(alpha); + const __m128 _beta = __lsx_vreplfr2vr_s(beta); + int jj = 0; +#if __loongarch_asx + const __m256 _alpha256 = (__m256)__lasx_xvreplfr2vr_s(alpha); + const __m256 _beta256 = (__m256)__lasx_xvreplfr2vr_s(beta); + for (; jj + 15 < max_jj; jj += 16) + { + __m256 _f00 = (__m256)__lasx_xvld(pp0 + jj, 0); + __m256 _f01 = (__m256)__lasx_xvld(pp0 + jj + 8, 0); + __m256 _f10 = (__m256)__lasx_xvld(pp1 + jj, 0); + __m256 _f11 = (__m256)__lasx_xvld(pp1 + jj + 8, 0); + __m256 _f20 = (__m256)__lasx_xvld(pp2 + jj, 0); + __m256 _f21 = (__m256)__lasx_xvld(pp2 + jj + 8, 0); + __m256 _f30 = (__m256)__lasx_xvld(pp3 + jj, 0); + __m256 _f31 = (__m256)__lasx_xvld(pp3 + jj + 8, 0); + if (pC) + { + if (broadcast_type_C == 0) + { + const __m256 _cc = (__m256)__lasx_xvreplfr2vr_s(pC[0] * beta); + _f00 = __lasx_xvfadd_s(_f00, _cc); + _f01 = __lasx_xvfadd_s(_f01, _cc); + _f10 = __lasx_xvfadd_s(_f10, _cc); + _f11 = __lasx_xvfadd_s(_f11, _cc); + _f20 = __lasx_xvfadd_s(_f20, _cc); + _f21 = __lasx_xvfadd_s(_f21, _cc); + _f30 = __lasx_xvfadd_s(_f30, _cc); + _f31 = __lasx_xvfadd_s(_f31, _cc); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + const __m256 _c0 = (__m256)__lasx_xvreplfr2vr_s(pC[i + ii] * beta); + const __m256 _c1 = (__m256)__lasx_xvreplfr2vr_s(pC[i + ii + 1] * beta); + const __m256 _c2 = (__m256)__lasx_xvreplfr2vr_s(pC[i + ii + 2] * beta); + const __m256 _c3 = (__m256)__lasx_xvreplfr2vr_s(pC[i + ii + 3] * beta); + _f00 = __lasx_xvfadd_s(_f00, _c0); + _f01 = __lasx_xvfadd_s(_f01, _c0); + _f10 = __lasx_xvfadd_s(_f10, _c1); + _f11 = __lasx_xvfadd_s(_f11, _c1); + _f20 = __lasx_xvfadd_s(_f20, _c2); + _f21 = __lasx_xvfadd_s(_f21, _c2); + _f30 = __lasx_xvfadd_s(_f30, _c3); + _f31 = __lasx_xvfadd_s(_f31, _c3); + } + if (broadcast_type_C == 3) + { + const __m256 _c00 = (__m256)__lasx_xvld(pC, 0); + const __m256 _c01 = (__m256)__lasx_xvld(pC + 8, 0); + const __m256 _c10 = (__m256)__lasx_xvld(pC + c_hstep, 0); + const __m256 _c11 = (__m256)__lasx_xvld(pC + c_hstep + 8, 0); + const __m256 _c20 = (__m256)__lasx_xvld(pC + c_hstep * 2, 0); + const __m256 _c21 = (__m256)__lasx_xvld(pC + c_hstep * 2 + 8, 0); + const __m256 _c30 = (__m256)__lasx_xvld(pC + c_hstep * 3, 0); + const __m256 _c31 = (__m256)__lasx_xvld(pC + c_hstep * 3 + 8, 0); + if (beta == 1.f) + { + _f00 = __lasx_xvfadd_s(_f00, _c00); + _f01 = __lasx_xvfadd_s(_f01, _c01); + _f10 = __lasx_xvfadd_s(_f10, _c10); + _f11 = __lasx_xvfadd_s(_f11, _c11); + _f20 = __lasx_xvfadd_s(_f20, _c20); + _f21 = __lasx_xvfadd_s(_f21, _c21); + _f30 = __lasx_xvfadd_s(_f30, _c30); + _f31 = __lasx_xvfadd_s(_f31, _c31); + } + else + { + _f00 = __lasx_xvfmadd_s(_c00, _beta256, _f00); + _f01 = __lasx_xvfmadd_s(_c01, _beta256, _f01); + _f10 = __lasx_xvfmadd_s(_c10, _beta256, _f10); + _f11 = __lasx_xvfmadd_s(_c11, _beta256, _f11); + _f20 = __lasx_xvfmadd_s(_c20, _beta256, _f20); + _f21 = __lasx_xvfmadd_s(_c21, _beta256, _f21); + _f30 = __lasx_xvfmadd_s(_c30, _beta256, _f30); + _f31 = __lasx_xvfmadd_s(_c31, _beta256, _f31); + } + } + if (broadcast_type_C == 4) + { + __m256 _c0 = (__m256)__lasx_xvld(pC, 0); + __m256 _c1 = (__m256)__lasx_xvld(pC + 8, 0); + if (beta != 1.f) + { + _c0 = __lasx_xvfmul_s(_c0, _beta256); + _c1 = __lasx_xvfmul_s(_c1, _beta256); + } + _f00 = __lasx_xvfadd_s(_f00, _c0); + _f01 = __lasx_xvfadd_s(_f01, _c1); + _f10 = __lasx_xvfadd_s(_f10, _c0); + _f11 = __lasx_xvfadd_s(_f11, _c1); + _f20 = __lasx_xvfadd_s(_f20, _c0); + _f21 = __lasx_xvfadd_s(_f21, _c1); + _f30 = __lasx_xvfadd_s(_f30, _c0); + _f31 = __lasx_xvfadd_s(_f31, _c1); + } + } + if (alpha != 1.f) + { + _f00 = __lasx_xvfmul_s(_f00, _alpha256); + _f01 = __lasx_xvfmul_s(_f01, _alpha256); + _f10 = __lasx_xvfmul_s(_f10, _alpha256); + _f11 = __lasx_xvfmul_s(_f11, _alpha256); + _f20 = __lasx_xvfmul_s(_f20, _alpha256); + _f21 = __lasx_xvfmul_s(_f21, _alpha256); + _f30 = __lasx_xvfmul_s(_f30, _alpha256); + _f31 = __lasx_xvfmul_s(_f31, _alpha256); + } + + __m256i _tmp0 = __lasx_xvilvl_w((__m256i)_f10, (__m256i)_f00); + __m256i _tmp1 = __lasx_xvilvh_w((__m256i)_f10, (__m256i)_f00); + __m256i _tmp2 = __lasx_xvilvl_w((__m256i)_f30, (__m256i)_f20); + __m256i _tmp3 = __lasx_xvilvh_w((__m256i)_f30, (__m256i)_f20); + __m256i _r0 = __lasx_xvilvl_d(_tmp2, _tmp0); + __m256i _r1 = __lasx_xvilvh_d(_tmp2, _tmp0); + __m256i _r2 = __lasx_xvilvl_d(_tmp3, _tmp1); + __m256i _r3 = __lasx_xvilvh_d(_tmp3, _tmp1); + __lsx_vst(__lasx_extract_128_lo(_r0), p0, 0); + __lsx_vst(__lasx_extract_128_lo(_r1), p0 + M, 0); + __lsx_vst(__lasx_extract_128_lo(_r2), p0 + M * 2, 0); + __lsx_vst(__lasx_extract_128_lo(_r3), p0 + M * 3, 0); + __lsx_vst(__lasx_extract_128_hi(_r0), p0 + M * 4, 0); + __lsx_vst(__lasx_extract_128_hi(_r1), p0 + M * 5, 0); + __lsx_vst(__lasx_extract_128_hi(_r2), p0 + M * 6, 0); + __lsx_vst(__lasx_extract_128_hi(_r3), p0 + M * 7, 0); + + _tmp0 = __lasx_xvilvl_w((__m256i)_f11, (__m256i)_f01); + _tmp1 = __lasx_xvilvh_w((__m256i)_f11, (__m256i)_f01); + _tmp2 = __lasx_xvilvl_w((__m256i)_f31, (__m256i)_f21); + _tmp3 = __lasx_xvilvh_w((__m256i)_f31, (__m256i)_f21); + _r0 = __lasx_xvilvl_d(_tmp2, _tmp0); + _r1 = __lasx_xvilvh_d(_tmp2, _tmp0); + _r2 = __lasx_xvilvl_d(_tmp3, _tmp1); + _r3 = __lasx_xvilvh_d(_tmp3, _tmp1); + __lsx_vst(__lasx_extract_128_lo(_r0), p0 + M * 8, 0); + __lsx_vst(__lasx_extract_128_lo(_r1), p0 + M * 9, 0); + __lsx_vst(__lasx_extract_128_lo(_r2), p0 + M * 10, 0); + __lsx_vst(__lasx_extract_128_lo(_r3), p0 + M * 11, 0); + __lsx_vst(__lasx_extract_128_hi(_r0), p0 + M * 12, 0); + __lsx_vst(__lasx_extract_128_hi(_r1), p0 + M * 13, 0); + __lsx_vst(__lasx_extract_128_hi(_r2), p0 + M * 14, 0); + __lsx_vst(__lasx_extract_128_hi(_r3), p0 + M * 15, 0); + p0 += M * 16; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 16; + } + for (; jj + 7 < max_jj; jj += 8) + { + __m256 _f0 = (__m256)__lasx_xvld(pp0 + jj, 0); + __m256 _f1 = (__m256)__lasx_xvld(pp1 + jj, 0); + __m256 _f2 = (__m256)__lasx_xvld(pp2 + jj, 0); + __m256 _f3 = (__m256)__lasx_xvld(pp3 + jj, 0); + if (pC) + { + if (broadcast_type_C == 0) + { + __m256 _c0 = (__m256)__lasx_xvreplfr2vr_s(pC[0] * beta); + _f0 = __lasx_xvfadd_s(_f0, _c0); + _f1 = __lasx_xvfadd_s(_f1, _c0); + _f2 = __lasx_xvfadd_s(_f2, _c0); + _f3 = __lasx_xvfadd_s(_f3, _c0); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = __lasx_xvfadd_s(_f0, (__m256)__lasx_xvreplfr2vr_s(pC[i + ii] * beta)); + _f1 = __lasx_xvfadd_s(_f1, (__m256)__lasx_xvreplfr2vr_s(pC[i + ii + 1] * beta)); + _f2 = __lasx_xvfadd_s(_f2, (__m256)__lasx_xvreplfr2vr_s(pC[i + ii + 2] * beta)); + _f3 = __lasx_xvfadd_s(_f3, (__m256)__lasx_xvreplfr2vr_s(pC[i + ii + 3] * beta)); + } + if (broadcast_type_C == 3) + { + if (beta == 1.f) + { + _f0 = __lasx_xvfadd_s(_f0, (__m256)__lasx_xvld(pC, 0)); + _f1 = __lasx_xvfadd_s(_f1, (__m256)__lasx_xvld(pC + c_hstep, 0)); + _f2 = __lasx_xvfadd_s(_f2, (__m256)__lasx_xvld(pC + c_hstep * 2, 0)); + _f3 = __lasx_xvfadd_s(_f3, (__m256)__lasx_xvld(pC + c_hstep * 3, 0)); + } + else + { + _f0 = __lasx_xvfmadd_s((__m256)__lasx_xvld(pC, 0), _beta256, _f0); + _f1 = __lasx_xvfmadd_s((__m256)__lasx_xvld(pC + c_hstep, 0), _beta256, _f1); + _f2 = __lasx_xvfmadd_s((__m256)__lasx_xvld(pC + c_hstep * 2, 0), _beta256, _f2); + _f3 = __lasx_xvfmadd_s((__m256)__lasx_xvld(pC + c_hstep * 3, 0), _beta256, _f3); + } + } + if (broadcast_type_C == 4) + { + __m256 _c4 = (__m256)__lasx_xvld(pC, 0); + if (beta != 1.f) + _c4 = __lasx_xvfmul_s(_c4, _beta256); + _f0 = __lasx_xvfadd_s(_f0, _c4); + _f1 = __lasx_xvfadd_s(_f1, _c4); + _f2 = __lasx_xvfadd_s(_f2, _c4); + _f3 = __lasx_xvfadd_s(_f3, _c4); + } + } + if (alpha != 1.f) + { + _f0 = __lasx_xvfmul_s(_f0, _alpha256); + _f1 = __lasx_xvfmul_s(_f1, _alpha256); + _f2 = __lasx_xvfmul_s(_f2, _alpha256); + _f3 = __lasx_xvfmul_s(_f3, _alpha256); + } + + __m256i _tmp0 = __lasx_xvilvl_w((__m256i)_f1, (__m256i)_f0); + __m256i _tmp1 = __lasx_xvilvh_w((__m256i)_f1, (__m256i)_f0); + __m256i _tmp2 = __lasx_xvilvl_w((__m256i)_f3, (__m256i)_f2); + __m256i _tmp3 = __lasx_xvilvh_w((__m256i)_f3, (__m256i)_f2); + __m256i _r0 = __lasx_xvilvl_d(_tmp2, _tmp0); + __m256i _r1 = __lasx_xvilvh_d(_tmp2, _tmp0); + __m256i _r2 = __lasx_xvilvl_d(_tmp3, _tmp1); + __m256i _r3 = __lasx_xvilvh_d(_tmp3, _tmp1); + + __lsx_vst(__lasx_extract_128_lo(_r0), p0, 0); + __lsx_vst(__lasx_extract_128_lo(_r1), p0 + M, 0); + __lsx_vst(__lasx_extract_128_lo(_r2), p0 + M * 2, 0); + __lsx_vst(__lasx_extract_128_lo(_r3), p0 + M * 3, 0); + __lsx_vst(__lasx_extract_128_hi(_r0), p0 + M * 4, 0); + __lsx_vst(__lasx_extract_128_hi(_r1), p0 + M * 5, 0); + __lsx_vst(__lasx_extract_128_hi(_r2), p0 + M * 6, 0); + __lsx_vst(__lasx_extract_128_hi(_r3), p0 + M * 7, 0); + p0 += M * 8; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 8; + } +#endif + for (; jj + 7 < max_jj; jj += 8) + { + __m128 _f00 = (__m128)__lsx_vld(pp0 + jj, 0); + __m128 _f01 = (__m128)__lsx_vld(pp0 + jj + 4, 0); + __m128 _f10 = (__m128)__lsx_vld(pp1 + jj, 0); + __m128 _f11 = (__m128)__lsx_vld(pp1 + jj + 4, 0); + __m128 _f20 = (__m128)__lsx_vld(pp2 + jj, 0); + __m128 _f21 = (__m128)__lsx_vld(pp2 + jj + 4, 0); + __m128 _f30 = (__m128)__lsx_vld(pp3 + jj, 0); + __m128 _f31 = (__m128)__lsx_vld(pp3 + jj + 4, 0); + transpose4x4_ps(_f00, _f10, _f20, _f30); + transpose4x4_ps(_f01, _f11, _f21, _f31); + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f00 = __lsx_vfadd_s(_f00, _c); + _f10 = __lsx_vfadd_s(_f10, _c); + _f20 = __lsx_vfadd_s(_f20, _c); + _f30 = __lsx_vfadd_s(_f30, _c); + _f01 = __lsx_vfadd_s(_f01, _c); + _f11 = __lsx_vfadd_s(_f11, _c); + _f21 = __lsx_vfadd_s(_f21, _c); + _f31 = __lsx_vfadd_s(_f31, _c); + } + if (broadcast_type_C == 3) + { + __m128 _c00 = (__m128)__lsx_vld(pC, 0); + __m128 _c01 = (__m128)__lsx_vld(pC + 4, 0); + __m128 _c10 = (__m128)__lsx_vld(pC + c_hstep, 0); + __m128 _c11 = (__m128)__lsx_vld(pC + c_hstep + 4, 0); + __m128 _c20 = (__m128)__lsx_vld(pC + c_hstep * 2, 0); + __m128 _c21 = (__m128)__lsx_vld(pC + c_hstep * 2 + 4, 0); + __m128 _c30 = (__m128)__lsx_vld(pC + c_hstep * 3, 0); + __m128 _c31 = (__m128)__lsx_vld(pC + c_hstep * 3 + 4, 0); + transpose4x4_ps(_c00, _c10, _c20, _c30); + transpose4x4_ps(_c01, _c11, _c21, _c31); + if (beta == 1.f) + { + _f00 = __lsx_vfadd_s(_f00, _c00); + _f10 = __lsx_vfadd_s(_f10, _c10); + _f20 = __lsx_vfadd_s(_f20, _c20); + _f30 = __lsx_vfadd_s(_f30, _c30); + _f01 = __lsx_vfadd_s(_f01, _c01); + _f11 = __lsx_vfadd_s(_f11, _c11); + _f21 = __lsx_vfadd_s(_f21, _c21); + _f31 = __lsx_vfadd_s(_f31, _c31); + } + else + { + _f00 = __lsx_vfmadd_s(_c00, _beta, _f00); + _f10 = __lsx_vfmadd_s(_c10, _beta, _f10); + _f20 = __lsx_vfmadd_s(_c20, _beta, _f20); + _f30 = __lsx_vfmadd_s(_c30, _beta, _f30); + _f01 = __lsx_vfmadd_s(_c01, _beta, _f01); + _f11 = __lsx_vfmadd_s(_c11, _beta, _f11); + _f21 = __lsx_vfmadd_s(_c21, _beta, _f21); + _f31 = __lsx_vfmadd_s(_c31, _beta, _f31); + } + } + if (broadcast_type_C == 4) + { + __m128 _cc0 = (__m128)__lsx_vld(pC, 0); + __m128 _cc1 = (__m128)__lsx_vld(pC + 4, 0); + if (beta != 1.f) + { + _cc0 = __lsx_vfmul_s(_cc0, _beta); + _cc1 = __lsx_vfmul_s(_cc1, _beta); + } + _f00 = __lsx_vfadd_s(_f00, (__m128)__lsx_vreplvei_w((__m128i)_cc0, 0)); + _f10 = __lsx_vfadd_s(_f10, (__m128)__lsx_vreplvei_w((__m128i)_cc0, 1)); + _f20 = __lsx_vfadd_s(_f20, (__m128)__lsx_vreplvei_w((__m128i)_cc0, 2)); + _f30 = __lsx_vfadd_s(_f30, (__m128)__lsx_vreplvei_w((__m128i)_cc0, 3)); + _f01 = __lsx_vfadd_s(_f01, (__m128)__lsx_vreplvei_w((__m128i)_cc1, 0)); + _f11 = __lsx_vfadd_s(_f11, (__m128)__lsx_vreplvei_w((__m128i)_cc1, 1)); + _f21 = __lsx_vfadd_s(_f21, (__m128)__lsx_vreplvei_w((__m128i)_cc1, 2)); + _f31 = __lsx_vfadd_s(_f31, (__m128)__lsx_vreplvei_w((__m128i)_cc1, 3)); + } + } + if (alpha != 1.f) + { + _f00 = __lsx_vfmul_s(_f00, _alpha); + _f10 = __lsx_vfmul_s(_f10, _alpha); + _f20 = __lsx_vfmul_s(_f20, _alpha); + _f30 = __lsx_vfmul_s(_f30, _alpha); + _f01 = __lsx_vfmul_s(_f01, _alpha); + _f11 = __lsx_vfmul_s(_f11, _alpha); + _f21 = __lsx_vfmul_s(_f21, _alpha); + _f31 = __lsx_vfmul_s(_f31, _alpha); + } + __lsx_vst((__m128i)_f00, p0, 0); + __lsx_vst((__m128i)_f10, p0 + M, 0); + __lsx_vst((__m128i)_f20, p0 + M * 2, 0); + __lsx_vst((__m128i)_f30, p0 + M * 3, 0); + __lsx_vst((__m128i)_f01, p0 + M * 4, 0); + __lsx_vst((__m128i)_f11, p0 + M * 5, 0); + __lsx_vst((__m128i)_f21, p0 + M * 6, 0); + __lsx_vst((__m128i)_f31, p0 + M * 7, 0); + p0 += M * 8; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 8; + } + for (; jj + 3 < max_jj; jj += 4) + { + __m128 _f0 = (__m128)__lsx_vld(pp0 + jj, 0); + __m128 _f1 = (__m128)__lsx_vld(pp1 + jj, 0); + __m128 _f2 = (__m128)__lsx_vld(pp2 + jj, 0); + __m128 _f3 = (__m128)__lsx_vld(pp3 + jj, 0); + transpose4x4_ps(_f0, _f1, _f2, _f3); + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = __lsx_vfadd_s(_f0, _c); + _f1 = __lsx_vfadd_s(_f1, _c); + _f2 = __lsx_vfadd_s(_f2, _c); + _f3 = __lsx_vfadd_s(_f3, _c); + } + if (broadcast_type_C == 3) + { + __m128 _c0 = (__m128)__lsx_vld(pC, 0); + __m128 _c1 = (__m128)__lsx_vld(pC + c_hstep, 0); + __m128 _c2 = (__m128)__lsx_vld(pC + c_hstep * 2, 0); + __m128 _c3 = (__m128)__lsx_vld(pC + c_hstep * 3, 0); + transpose4x4_ps(_c0, _c1, _c2, _c3); + if (beta == 1.f) + { + _f0 = __lsx_vfadd_s(_f0, _c0); + _f1 = __lsx_vfadd_s(_f1, _c1); + _f2 = __lsx_vfadd_s(_f2, _c2); + _f3 = __lsx_vfadd_s(_f3, _c3); + } + else + { + _f0 = __lsx_vfmadd_s(_c0, _beta, _f0); + _f1 = __lsx_vfmadd_s(_c1, _beta, _f1); + _f2 = __lsx_vfmadd_s(_c2, _beta, _f2); + _f3 = __lsx_vfmadd_s(_c3, _beta, _f3); + } + } + if (broadcast_type_C == 4) + { + __m128 _c4 = (__m128)__lsx_vld(pC, 0); + if (beta != 1.f) + _c4 = __lsx_vfmul_s(_c4, _beta); + _f0 = __lsx_vfadd_s(_f0, (__m128)__lsx_vreplvei_w((__m128i)_c4, 0)); + _f1 = __lsx_vfadd_s(_f1, (__m128)__lsx_vreplvei_w((__m128i)_c4, 1)); + _f2 = __lsx_vfadd_s(_f2, (__m128)__lsx_vreplvei_w((__m128i)_c4, 2)); + _f3 = __lsx_vfadd_s(_f3, (__m128)__lsx_vreplvei_w((__m128i)_c4, 3)); + } + } + if (alpha != 1.f) + { + _f0 = __lsx_vfmul_s(_f0, _alpha); + _f1 = __lsx_vfmul_s(_f1, _alpha); + _f2 = __lsx_vfmul_s(_f2, _alpha); + _f3 = __lsx_vfmul_s(_f3, _alpha); + } + __lsx_vst((__m128i)_f0, p0, 0); + __lsx_vst((__m128i)_f1, p0 + M, 0); + __lsx_vst((__m128i)_f2, p0 + M * 2, 0); + __lsx_vst((__m128i)_f3, p0 + M * 3, 0); + p0 += M * 4; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 4; + } + for (; jj + 1 < max_jj; jj += 2) + { + __m128i _r0 = __lsx_vldrepl_d(pp0 + jj, 0); + __m128i _r1 = __lsx_vldrepl_d(pp1 + jj, 0); + __m128i _r2 = __lsx_vldrepl_d(pp2 + jj, 0); + __m128i _r3 = __lsx_vldrepl_d(pp3 + jj, 0); + __m128i _t0 = __lsx_vilvl_w(_r1, _r0); + __m128i _t1 = __lsx_vilvl_w(_r3, _r2); + __m128 _f0 = (__m128)__lsx_vilvl_d(_t1, _t0); + __m128 _f1 = (__m128)__lsx_vilvh_d(_t1, _t0); + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = __lsx_vfadd_s(_f0, _c); + _f1 = __lsx_vfadd_s(_f1, _c); + } + if (broadcast_type_C == 3) + { + _r0 = __lsx_vldrepl_d(pC, 0); + _r1 = __lsx_vldrepl_d(pC + c_hstep, 0); + _r2 = __lsx_vldrepl_d(pC + c_hstep * 2, 0); + _r3 = __lsx_vldrepl_d(pC + c_hstep * 3, 0); + _t0 = __lsx_vilvl_w(_r1, _r0); + _t1 = __lsx_vilvl_w(_r3, _r2); + const __m128 _cc0 = (__m128)__lsx_vilvl_d(_t1, _t0); + const __m128 _cc1 = (__m128)__lsx_vilvh_d(_t1, _t0); + if (beta == 1.f) + { + _f0 = __lsx_vfadd_s(_f0, _cc0); + _f1 = __lsx_vfadd_s(_f1, _cc1); + } + else + { + _f0 = __lsx_vfmadd_s(_cc0, _beta, _f0); + _f1 = __lsx_vfmadd_s(_cc1, _beta, _f1); + } + } + if (broadcast_type_C == 4) + { + _f0 = __lsx_vfadd_s(_f0, __lsx_vreplfr2vr_s(pC[0] * beta)); + _f1 = __lsx_vfadd_s(_f1, __lsx_vreplfr2vr_s(pC[1] * beta)); + } + } + if (alpha != 1.f) + { + _f0 = __lsx_vfmul_s(_f0, _alpha); + _f1 = __lsx_vfmul_s(_f1, _alpha); + } + __lsx_vst((__m128i)_f0, p0, 0); + __lsx_vst((__m128i)_f1, p0 + M, 0); + p0 += M * 2; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 2; + } + for (; jj < max_jj; jj++) + { + __m128i _fi = __lsx_vldrepl_w(pp0 + jj, 0); + _fi = __lsx_vinsgr2vr_w(_fi, ((const int*)(pp1 + jj))[0], 1); + _fi = __lsx_vinsgr2vr_w(_fi, ((const int*)(pp2 + jj))[0], 2); + _fi = __lsx_vinsgr2vr_w(_fi, ((const int*)(pp3 + jj))[0], 3); + __m128 _f = (__m128)_fi; + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + _f = __lsx_vfadd_s(_f, _c); + if (broadcast_type_C == 3) + { + __m128i _ci = __lsx_vldrepl_w(pC, 0); + _ci = __lsx_vinsgr2vr_w(_ci, ((const int*)(pC + c_hstep))[0], 1); + _ci = __lsx_vinsgr2vr_w(_ci, ((const int*)(pC + c_hstep * 2))[0], 2); + _ci = __lsx_vinsgr2vr_w(_ci, ((const int*)(pC + c_hstep * 3))[0], 3); + if (beta == 1.f) + _f = __lsx_vfadd_s(_f, (__m128)_ci); + else + _f = __lsx_vfmadd_s((__m128)_ci, _beta, _f); + } + if (broadcast_type_C == 4) + { + __m128 _c4 = __lsx_vreplfr2vr_s(pC[0]); + if (beta == 1.f) + _f = __lsx_vfadd_s(_f, _c4); + else + _f = __lsx_vfmadd_s(_c4, _beta, _f); + } + } + if (alpha != 1.f) + _f = __lsx_vfmul_s(_f, _alpha); + __lsx_vst((__m128i)_f, p0, 0); + p0 += M; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC++; + } + } +#endif + for (; ii + 1 < max_ii; ii += 2) + { + const float* pp0 = pp + (ii + 0) * max_jj; + const float* pp1 = pp + (ii + 1) * max_jj; + float* p0 = outptr + (size_t)j * M + i + ii; + const float* pC = pC_base; + + float c0 = 0.f; + float c1 = 0.f; + if (pC) + { + if (broadcast_type_C == 0) + { + c0 = pC[0] * beta; + c1 = c0; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + c0 = pC[i + ii] * beta; + c1 = pC[i + ii + 1] * beta; + } + if (broadcast_type_C == 3) + { + pC += (size_t)(i + ii) * c_hstep + j; + } + if (broadcast_type_C == 4) + pC += j; + } + + int jj = 0; +#if __loongarch_sx + __m128 _c = __lsx_vreplfr2vr_s(0.f); + if (pC && broadcast_type_C == 0) + _c = __lsx_vreplfr2vr_s(c0); + if (pC && (broadcast_type_C == 1 || broadcast_type_C == 2)) + { + _c = (__m128)__lsx_vldrepl_d(pC + i + ii, 0); + if (beta != 1.f) + _c = __lsx_vfmul_s(_c, __lsx_vreplfr2vr_s(beta)); + } + const __m128 _alpha = __lsx_vreplfr2vr_s(alpha); + const __m128 _beta = __lsx_vreplfr2vr_s(beta); +#if __loongarch_asx + const __m256 _alpha256 = (__m256)__lasx_xvreplfr2vr_s(alpha); + const __m256 _beta256 = (__m256)__lasx_xvreplfr2vr_s(beta); + for (; jj + 15 < max_jj; jj += 16) + { + __m256 _f00 = (__m256)__lasx_xvld(pp0 + jj, 0); + __m256 _f01 = (__m256)__lasx_xvld(pp0 + jj + 8, 0); + __m256 _f10 = (__m256)__lasx_xvld(pp1 + jj, 0); + __m256 _f11 = (__m256)__lasx_xvld(pp1 + jj + 8, 0); + if (pC) + { + if (broadcast_type_C == 0) + { + const __m256 _cc = (__m256)__lasx_xvreplfr2vr_s(c0); + _f00 = __lasx_xvfadd_s(_f00, _cc); + _f01 = __lasx_xvfadd_s(_f01, _cc); + _f10 = __lasx_xvfadd_s(_f10, _cc); + _f11 = __lasx_xvfadd_s(_f11, _cc); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + const __m256 _c0 = (__m256)__lasx_xvreplfr2vr_s(c0); + const __m256 _c1 = (__m256)__lasx_xvreplfr2vr_s(c1); + _f00 = __lasx_xvfadd_s(_f00, _c0); + _f01 = __lasx_xvfadd_s(_f01, _c0); + _f10 = __lasx_xvfadd_s(_f10, _c1); + _f11 = __lasx_xvfadd_s(_f11, _c1); + } + if (broadcast_type_C == 3) + { + const __m256 _c00 = (__m256)__lasx_xvld(pC, 0); + const __m256 _c01 = (__m256)__lasx_xvld(pC + 8, 0); + const __m256 _c10 = (__m256)__lasx_xvld(pC + c_hstep, 0); + const __m256 _c11 = (__m256)__lasx_xvld(pC + c_hstep + 8, 0); + if (beta == 1.f) + { + _f00 = __lasx_xvfadd_s(_f00, _c00); + _f01 = __lasx_xvfadd_s(_f01, _c01); + _f10 = __lasx_xvfadd_s(_f10, _c10); + _f11 = __lasx_xvfadd_s(_f11, _c11); + } + else + { + _f00 = __lasx_xvfmadd_s(_c00, _beta256, _f00); + _f01 = __lasx_xvfmadd_s(_c01, _beta256, _f01); + _f10 = __lasx_xvfmadd_s(_c10, _beta256, _f10); + _f11 = __lasx_xvfmadd_s(_c11, _beta256, _f11); + } + } + if (broadcast_type_C == 4) + { + __m256 _c0 = (__m256)__lasx_xvld(pC, 0); + __m256 _c1 = (__m256)__lasx_xvld(pC + 8, 0); + if (beta != 1.f) + { + _c0 = __lasx_xvfmul_s(_c0, _beta256); + _c1 = __lasx_xvfmul_s(_c1, _beta256); + } + _f00 = __lasx_xvfadd_s(_f00, _c0); + _f01 = __lasx_xvfadd_s(_f01, _c1); + _f10 = __lasx_xvfadd_s(_f10, _c0); + _f11 = __lasx_xvfadd_s(_f11, _c1); + } + } + if (alpha != 1.f) + { + _f00 = __lasx_xvfmul_s(_f00, _alpha256); + _f01 = __lasx_xvfmul_s(_f01, _alpha256); + _f10 = __lasx_xvfmul_s(_f10, _alpha256); + _f11 = __lasx_xvfmul_s(_f11, _alpha256); + } + __m256i _tmp0 = __lasx_xvilvl_w((__m256i)_f10, (__m256i)_f00); + __m256i _tmp1 = __lasx_xvilvh_w((__m256i)_f10, (__m256i)_f00); + __lasx_xvstelm_d(_tmp0, p0, 0, 0); + __lasx_xvstelm_d(_tmp0, p0 + M, 0, 1); + __lasx_xvstelm_d(_tmp1, p0 + M * 2, 0, 0); + __lasx_xvstelm_d(_tmp1, p0 + M * 3, 0, 1); + __lasx_xvstelm_d(_tmp0, p0 + M * 4, 0, 2); + __lasx_xvstelm_d(_tmp0, p0 + M * 5, 0, 3); + __lasx_xvstelm_d(_tmp1, p0 + M * 6, 0, 2); + __lasx_xvstelm_d(_tmp1, p0 + M * 7, 0, 3); + _tmp0 = __lasx_xvilvl_w((__m256i)_f11, (__m256i)_f01); + _tmp1 = __lasx_xvilvh_w((__m256i)_f11, (__m256i)_f01); + __lasx_xvstelm_d(_tmp0, p0 + M * 8, 0, 0); + __lasx_xvstelm_d(_tmp0, p0 + M * 9, 0, 1); + __lasx_xvstelm_d(_tmp1, p0 + M * 10, 0, 0); + __lasx_xvstelm_d(_tmp1, p0 + M * 11, 0, 1); + __lasx_xvstelm_d(_tmp0, p0 + M * 12, 0, 2); + __lasx_xvstelm_d(_tmp0, p0 + M * 13, 0, 3); + __lasx_xvstelm_d(_tmp1, p0 + M * 14, 0, 2); + __lasx_xvstelm_d(_tmp1, p0 + M * 15, 0, 3); + p0 += M * 16; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 16; + } + for (; jj + 7 < max_jj; jj += 8) + { + __m256 _f0 = (__m256)__lasx_xvld(pp0 + jj, 0); + __m256 _f1 = (__m256)__lasx_xvld(pp1 + jj, 0); + if (pC) + { + if (broadcast_type_C == 0) + { + __m256 _c0 = (__m256)__lasx_xvreplfr2vr_s(c0); + _f0 = __lasx_xvfadd_s(_f0, _c0); + _f1 = __lasx_xvfadd_s(_f1, _c0); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = __lasx_xvfadd_s(_f0, (__m256)__lasx_xvreplfr2vr_s(c0)); + _f1 = __lasx_xvfadd_s(_f1, (__m256)__lasx_xvreplfr2vr_s(c1)); + } + if (broadcast_type_C == 3) + { + if (beta == 1.f) + { + _f0 = __lasx_xvfadd_s(_f0, (__m256)__lasx_xvld(pC, 0)); + _f1 = __lasx_xvfadd_s(_f1, (__m256)__lasx_xvld(pC + c_hstep, 0)); + } + else + { + _f0 = __lasx_xvfmadd_s((__m256)__lasx_xvld(pC, 0), _beta256, _f0); + _f1 = __lasx_xvfmadd_s((__m256)__lasx_xvld(pC + c_hstep, 0), _beta256, _f1); + } + } + if (broadcast_type_C == 4) + { + __m256 _c4 = (__m256)__lasx_xvld(pC, 0); + if (beta != 1.f) + _c4 = __lasx_xvfmul_s(_c4, _beta256); + _f0 = __lasx_xvfadd_s(_f0, _c4); + _f1 = __lasx_xvfadd_s(_f1, _c4); + } + } + if (alpha != 1.f) + { + _f0 = __lasx_xvfmul_s(_f0, _alpha256); + _f1 = __lasx_xvfmul_s(_f1, _alpha256); + } + + __m256i _tmp0 = __lasx_xvilvl_w((__m256i)_f1, (__m256i)_f0); + __m256i _tmp1 = __lasx_xvilvh_w((__m256i)_f1, (__m256i)_f0); + __lasx_xvstelm_d(_tmp0, p0, 0, 0); + __lasx_xvstelm_d(_tmp0, p0 + M, 0, 1); + __lasx_xvstelm_d(_tmp1, p0 + M * 2, 0, 0); + __lasx_xvstelm_d(_tmp1, p0 + M * 3, 0, 1); + __lasx_xvstelm_d(_tmp0, p0 + M * 4, 0, 2); + __lasx_xvstelm_d(_tmp0, p0 + M * 5, 0, 3); + __lasx_xvstelm_d(_tmp1, p0 + M * 6, 0, 2); + __lasx_xvstelm_d(_tmp1, p0 + M * 7, 0, 3); + p0 += M * 8; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 8; + } +#endif + for (; jj + 7 < max_jj; jj += 8) + { + __m128 _f00 = (__m128)__lsx_vld(pp0 + jj, 0); + __m128 _f01 = (__m128)__lsx_vld(pp0 + jj + 4, 0); + __m128 _f10 = (__m128)__lsx_vld(pp1 + jj, 0); + __m128 _f11 = (__m128)__lsx_vld(pp1 + jj + 4, 0); + if (pC) + { + if (broadcast_type_C == 0) + { + const __m128 _cc = __lsx_vreplfr2vr_s(c0); + _f00 = __lsx_vfadd_s(_f00, _cc); + _f01 = __lsx_vfadd_s(_f01, _cc); + _f10 = __lsx_vfadd_s(_f10, _cc); + _f11 = __lsx_vfadd_s(_f11, _cc); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + const __m128 _c0 = __lsx_vreplfr2vr_s(c0); + const __m128 _c1 = __lsx_vreplfr2vr_s(c1); + _f00 = __lsx_vfadd_s(_f00, _c0); + _f01 = __lsx_vfadd_s(_f01, _c0); + _f10 = __lsx_vfadd_s(_f10, _c1); + _f11 = __lsx_vfadd_s(_f11, _c1); + } + if (broadcast_type_C == 3) + { + const __m128 _c00 = (__m128)__lsx_vld(pC, 0); + const __m128 _c01 = (__m128)__lsx_vld(pC + 4, 0); + const __m128 _c10 = (__m128)__lsx_vld(pC + c_hstep, 0); + const __m128 _c11 = (__m128)__lsx_vld(pC + c_hstep + 4, 0); + if (beta == 1.f) + { + _f00 = __lsx_vfadd_s(_f00, _c00); + _f01 = __lsx_vfadd_s(_f01, _c01); + _f10 = __lsx_vfadd_s(_f10, _c10); + _f11 = __lsx_vfadd_s(_f11, _c11); + } + else + { + _f00 = __lsx_vfmadd_s(_c00, _beta, _f00); + _f01 = __lsx_vfmadd_s(_c01, _beta, _f01); + _f10 = __lsx_vfmadd_s(_c10, _beta, _f10); + _f11 = __lsx_vfmadd_s(_c11, _beta, _f11); + } + } + if (broadcast_type_C == 4) + { + __m128 _c0 = (__m128)__lsx_vld(pC, 0); + __m128 _c1 = (__m128)__lsx_vld(pC + 4, 0); + if (beta != 1.f) + { + _c0 = __lsx_vfmul_s(_c0, _beta); + _c1 = __lsx_vfmul_s(_c1, _beta); + } + _f00 = __lsx_vfadd_s(_f00, _c0); + _f01 = __lsx_vfadd_s(_f01, _c1); + _f10 = __lsx_vfadd_s(_f10, _c0); + _f11 = __lsx_vfadd_s(_f11, _c1); + } + } + if (alpha != 1.f) + { + _f00 = __lsx_vfmul_s(_f00, _alpha); + _f01 = __lsx_vfmul_s(_f01, _alpha); + _f10 = __lsx_vfmul_s(_f10, _alpha); + _f11 = __lsx_vfmul_s(_f11, _alpha); + } + __m128i _tmp0 = __lsx_vilvl_w((__m128i)_f10, (__m128i)_f00); + __m128i _tmp1 = __lsx_vilvh_w((__m128i)_f10, (__m128i)_f00); + __lsx_vstelm_d(_tmp0, p0, 0, 0); + __lsx_vstelm_d(_tmp0, p0 + M, 0, 1); + __lsx_vstelm_d(_tmp1, p0 + M * 2, 0, 0); + __lsx_vstelm_d(_tmp1, p0 + M * 3, 0, 1); + _tmp0 = __lsx_vilvl_w((__m128i)_f11, (__m128i)_f01); + _tmp1 = __lsx_vilvh_w((__m128i)_f11, (__m128i)_f01); + __lsx_vstelm_d(_tmp0, p0 + M * 4, 0, 0); + __lsx_vstelm_d(_tmp0, p0 + M * 5, 0, 1); + __lsx_vstelm_d(_tmp1, p0 + M * 6, 0, 0); + __lsx_vstelm_d(_tmp1, p0 + M * 7, 0, 1); + p0 += M * 8; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 8; + } + for (; jj + 3 < max_jj; jj += 4) + { + __m128 _f0 = (__m128)__lsx_vld(pp0 + jj, 0); + __m128 _f1 = (__m128)__lsx_vld(pp1 + jj, 0); + if (pC) + { + if (broadcast_type_C == 0) + { + __m128 _cc = __lsx_vreplfr2vr_s(c0); + _f0 = __lsx_vfadd_s(_f0, _cc); + _f1 = __lsx_vfadd_s(_f1, _cc); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = __lsx_vfadd_s(_f0, __lsx_vreplfr2vr_s(c0)); + _f1 = __lsx_vfadd_s(_f1, __lsx_vreplfr2vr_s(c1)); + } + if (broadcast_type_C == 3) + { + __m128 _c0 = (__m128)__lsx_vld(pC, 0); + __m128 _c1 = (__m128)__lsx_vld(pC + c_hstep, 0); + if (beta == 1.f) + { + _f0 = __lsx_vfadd_s(_f0, _c0); + _f1 = __lsx_vfadd_s(_f1, _c1); + } + else + { + _f0 = __lsx_vfmadd_s(_c0, _beta, _f0); + _f1 = __lsx_vfmadd_s(_c1, _beta, _f1); + } + } + if (broadcast_type_C == 4) + { + __m128 _c4 = (__m128)__lsx_vld(pC, 0); + if (beta != 1.f) + _c4 = __lsx_vfmul_s(_c4, _beta); + _f0 = __lsx_vfadd_s(_f0, _c4); + _f1 = __lsx_vfadd_s(_f1, _c4); + } + } + if (alpha != 1.f) + { + _f0 = __lsx_vfmul_s(_f0, _alpha); + _f1 = __lsx_vfmul_s(_f1, _alpha); + } + __m128i _tmp0 = __lsx_vilvl_w((__m128i)_f1, (__m128i)_f0); + __m128i _tmp1 = __lsx_vilvh_w((__m128i)_f1, (__m128i)_f0); + __lsx_vstelm_d(_tmp0, p0, 0, 0); + __lsx_vstelm_d(_tmp0, p0 + M, 0, 1); + __lsx_vstelm_d(_tmp1, p0 + M * 2, 0, 0); + __lsx_vstelm_d(_tmp1, p0 + M * 3, 0, 1); + p0 += M * 4; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 4; + } + for (; jj + 1 < max_jj; jj += 2) + { + __m128i _r0 = __lsx_vldrepl_d(pp0 + jj, 0); + __m128i _r1 = __lsx_vldrepl_d(pp1 + jj, 0); + __m128 _f = (__m128)__lsx_vilvl_w(_r1, _r0); + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + _f = __lsx_vfadd_s(_f, _c); + if (broadcast_type_C == 3) + { + _r0 = __lsx_vldrepl_d(pC, 0); + _r1 = __lsx_vldrepl_d(pC + c_hstep, 0); + const __m128 _cc = (__m128)__lsx_vilvl_w(_r1, _r0); + if (beta == 1.f) + _f = __lsx_vfadd_s(_f, _cc); + else + _f = __lsx_vfmadd_s(_cc, _beta, _f); + } + if (broadcast_type_C == 4) + { + __m128i _cc = __lsx_vldrepl_d(pC, 0); + _cc = __lsx_vilvl_w(_cc, _cc); + if (beta == 1.f) + _f = __lsx_vfadd_s(_f, (__m128)_cc); + else + _f = __lsx_vfmadd_s((__m128)_cc, _beta, _f); + } + } + if (alpha != 1.f) + _f = __lsx_vfmul_s(_f, _alpha); + __lsx_vstelm_d((__m128i)_f, p0, 0, 0); + __lsx_vstelm_d((__m128i)_f, p0 + M, 0, 1); + p0 += M * 2; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 2; + } + for (; jj < max_jj; jj++) + { + __m128i _fi = __lsx_vldrepl_w(pp0 + jj, 0); + _fi = __lsx_vinsgr2vr_w(_fi, ((const int*)(pp1 + jj))[0], 1); + __m128 _f = (__m128)_fi; + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + _f = __lsx_vfadd_s(_f, _c); + if (broadcast_type_C == 3) + { + __m128i _ci = __lsx_vldrepl_w(pC, 0); + _ci = __lsx_vinsgr2vr_w(_ci, ((const int*)(pC + c_hstep))[0], 1); + if (beta == 1.f) + _f = __lsx_vfadd_s(_f, (__m128)_ci); + else + _f = __lsx_vfmadd_s((__m128)_ci, _beta, _f); + } + if (broadcast_type_C == 4) + { + __m128 _c4 = __lsx_vreplfr2vr_s(pC[0]); + if (beta == 1.f) + _f = __lsx_vfadd_s(_f, _c4); + else + _f = __lsx_vfmadd_s(_c4, _beta, _f); + } + } + if (alpha != 1.f) + _f = __lsx_vfmul_s(_f, _alpha); + __lsx_vstelm_d((__m128i)_f, p0, 0, 0); + p0 += M; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC++; + } +#endif + for (; jj < max_jj; jj++) + { + float f0 = pp0[jj]; + float f1 = pp1[jj]; + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + f0 += c0; + f1 += c1; + } + if (broadcast_type_C == 3) + { + if (beta == 1.f) + { + f0 += pC[0]; + f1 += pC[c_hstep]; + } + else + { + f0 += pC[0] * beta; + f1 += pC[c_hstep] * beta; + } + } + if (broadcast_type_C == 4) + { + float c = beta == 1.f ? pC[0] : pC[0] * beta; + f0 += c; + f1 += c; + } + } + if (alpha != 1.f) + { + f0 *= alpha; + f1 *= alpha; + } + p0[0] = f0; + p0[1] = f1; + p0 += M; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC++; + } + } + for (; ii < max_ii; ii++) + { + const float* pp0 = pp + ii * max_jj; + float* p0 = outptr + (size_t)j * M + i + ii; + const float* pC = pC_base; + + float c0 = 0.f; + if (pC) + { + if (broadcast_type_C == 0) + c0 = pC[0] * beta; + if (broadcast_type_C == 1 || broadcast_type_C == 2) + c0 = pC[i + ii] * beta; + if (broadcast_type_C == 3) + pC += (size_t)(i + ii) * c_hstep + j; + if (broadcast_type_C == 4) + pC += j; + } + + int jj = 0; +#if __loongarch_asx + const __m256 _alpha256 = (__m256)__lasx_xvreplfr2vr_s(alpha); + const __m256 _beta256 = (__m256)__lasx_xvreplfr2vr_s(beta); + const __m256 _c256 = (__m256)__lasx_xvreplfr2vr_s(c0); + for (; jj + 15 < max_jj; jj += 16) + { + __m256 _f0 = (__m256)__lasx_xvld(pp0 + jj, 0); + __m256 _f1 = (__m256)__lasx_xvld(pp0 + jj + 8, 0); + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = __lasx_xvfadd_s(_f0, _c256); + _f1 = __lasx_xvfadd_s(_f1, _c256); + } + if (broadcast_type_C == 3 || broadcast_type_C == 4) + { + const __m256 _c0 = (__m256)__lasx_xvld(pC, 0); + const __m256 _c1 = (__m256)__lasx_xvld(pC + 8, 0); + if (beta == 1.f) + { + _f0 = __lasx_xvfadd_s(_f0, _c0); + _f1 = __lasx_xvfadd_s(_f1, _c1); + } + else + { + _f0 = __lasx_xvfmadd_s(_c0, _beta256, _f0); + _f1 = __lasx_xvfmadd_s(_c1, _beta256, _f1); + } + } + } + if (alpha != 1.f) + { + _f0 = __lasx_xvfmul_s(_f0, _alpha256); + _f1 = __lasx_xvfmul_s(_f1, _alpha256); + } + if (M == 1) + { + __lasx_xvst(_f0, p0, 0); + __lasx_xvst(_f1, p0 + 8, 0); + } + else + { + __lasx_xvstelm_w((__m256i)_f0, p0, 0, 0); + __lasx_xvstelm_w((__m256i)_f0, p0 + M, 0, 1); + __lasx_xvstelm_w((__m256i)_f0, p0 + M * 2, 0, 2); + __lasx_xvstelm_w((__m256i)_f0, p0 + M * 3, 0, 3); + __lasx_xvstelm_w((__m256i)_f0, p0 + M * 4, 0, 4); + __lasx_xvstelm_w((__m256i)_f0, p0 + M * 5, 0, 5); + __lasx_xvstelm_w((__m256i)_f0, p0 + M * 6, 0, 6); + __lasx_xvstelm_w((__m256i)_f0, p0 + M * 7, 0, 7); + __lasx_xvstelm_w((__m256i)_f1, p0 + M * 8, 0, 0); + __lasx_xvstelm_w((__m256i)_f1, p0 + M * 9, 0, 1); + __lasx_xvstelm_w((__m256i)_f1, p0 + M * 10, 0, 2); + __lasx_xvstelm_w((__m256i)_f1, p0 + M * 11, 0, 3); + __lasx_xvstelm_w((__m256i)_f1, p0 + M * 12, 0, 4); + __lasx_xvstelm_w((__m256i)_f1, p0 + M * 13, 0, 5); + __lasx_xvstelm_w((__m256i)_f1, p0 + M * 14, 0, 6); + __lasx_xvstelm_w((__m256i)_f1, p0 + M * 15, 0, 7); + } + p0 += M * 16; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 16; + } + for (; jj + 7 < max_jj; jj += 8) + { + __m256 _f0 = (__m256)__lasx_xvld(pp0 + jj, 0); + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + _f0 = __lasx_xvfadd_s(_f0, _c256); + if (broadcast_type_C == 3 || broadcast_type_C == 4) + { + __m256 _c0 = (__m256)__lasx_xvld(pC, 0); + if (beta == 1.f) + _f0 = __lasx_xvfadd_s(_f0, _c0); + else + _f0 = __lasx_xvfmadd_s(_c0, _beta256, _f0); + } + } + if (alpha != 1.f) _f0 = __lasx_xvfmul_s(_f0, _alpha256); + if (M == 1) + __lasx_xvst(_f0, p0, 0); + else + { + __lasx_xvstelm_w((__m256i)_f0, p0, 0, 0); + __lasx_xvstelm_w((__m256i)_f0, p0 + M, 0, 1); + __lasx_xvstelm_w((__m256i)_f0, p0 + M * 2, 0, 2); + __lasx_xvstelm_w((__m256i)_f0, p0 + M * 3, 0, 3); + __lasx_xvstelm_w((__m256i)_f0, p0 + M * 4, 0, 4); + __lasx_xvstelm_w((__m256i)_f0, p0 + M * 5, 0, 5); + __lasx_xvstelm_w((__m256i)_f0, p0 + M * 6, 0, 6); + __lasx_xvstelm_w((__m256i)_f0, p0 + M * 7, 0, 7); + } + p0 += M * 8; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 8; + } +#endif +#if __loongarch_sx + const __m128 _alpha128 = __lsx_vreplfr2vr_s(alpha); + const __m128 _beta128 = __lsx_vreplfr2vr_s(beta); + const __m128 _c128 = __lsx_vreplfr2vr_s(c0); + for (; jj + 7 < max_jj; jj += 8) + { + __m128 _f0 = (__m128)__lsx_vld(pp0 + jj, 0); + __m128 _f1 = (__m128)__lsx_vld(pp0 + jj + 4, 0); + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = __lsx_vfadd_s(_f0, _c128); + _f1 = __lsx_vfadd_s(_f1, _c128); + } + if (broadcast_type_C == 3 || broadcast_type_C == 4) + { + const __m128 _c0 = (__m128)__lsx_vld(pC, 0); + const __m128 _c1 = (__m128)__lsx_vld(pC + 4, 0); + if (beta == 1.f) + { + _f0 = __lsx_vfadd_s(_f0, _c0); + _f1 = __lsx_vfadd_s(_f1, _c1); + } + else + { + _f0 = __lsx_vfmadd_s(_c0, _beta128, _f0); + _f1 = __lsx_vfmadd_s(_c1, _beta128, _f1); + } + } + } + if (alpha != 1.f) + { + _f0 = __lsx_vfmul_s(_f0, _alpha128); + _f1 = __lsx_vfmul_s(_f1, _alpha128); + } + if (M == 1) + { + __lsx_vst((__m128i)_f0, p0, 0); + __lsx_vst((__m128i)_f1, p0 + 4, 0); + } + else + { + __lsx_vstelm_w((__m128i)_f0, p0, 0, 0); + __lsx_vstelm_w((__m128i)_f0, p0 + M, 0, 1); + __lsx_vstelm_w((__m128i)_f0, p0 + M * 2, 0, 2); + __lsx_vstelm_w((__m128i)_f0, p0 + M * 3, 0, 3); + __lsx_vstelm_w((__m128i)_f1, p0 + M * 4, 0, 0); + __lsx_vstelm_w((__m128i)_f1, p0 + M * 5, 0, 1); + __lsx_vstelm_w((__m128i)_f1, p0 + M * 6, 0, 2); + __lsx_vstelm_w((__m128i)_f1, p0 + M * 7, 0, 3); + } + p0 += M * 8; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 8; + } + for (; jj + 3 < max_jj; jj += 4) + { + __m128 _f0 = (__m128)__lsx_vld(pp0 + jj, 0); + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + _f0 = __lsx_vfadd_s(_f0, _c128); + if (broadcast_type_C == 3 || broadcast_type_C == 4) + { + __m128 _c0 = (__m128)__lsx_vld(pC, 0); + if (beta == 1.f) + _f0 = __lsx_vfadd_s(_f0, _c0); + else + _f0 = __lsx_vfmadd_s(_c0, _beta128, _f0); + } + } + if (alpha != 1.f) _f0 = __lsx_vfmul_s(_f0, _alpha128); + if (M == 1) + __lsx_vst((__m128i)_f0, p0, 0); + else + { + __lsx_vstelm_w((__m128i)_f0, p0, 0, 0); + __lsx_vstelm_w((__m128i)_f0, p0 + M, 0, 1); + __lsx_vstelm_w((__m128i)_f0, p0 + M * 2, 0, 2); + __lsx_vstelm_w((__m128i)_f0, p0 + M * 3, 0, 3); + } + p0 += M * 4; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 4; + } + for (; jj + 1 < max_jj; jj += 2) + { + __m128 _f0 = (__m128)__lsx_vldrepl_d(pp0 + jj, 0); + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + _f0 = __lsx_vfadd_s(_f0, _c128); + if (broadcast_type_C == 3 || broadcast_type_C == 4) + { + const __m128 _cc = (__m128)__lsx_vldrepl_d(pC, 0); + if (beta == 1.f) + _f0 = __lsx_vfadd_s(_f0, _cc); + else + _f0 = __lsx_vfmadd_s(_cc, _beta128, _f0); + } + } + if (alpha != 1.f) + _f0 = __lsx_vfmul_s(_f0, _alpha128); + if (M == 1) + __lsx_vstelm_d((__m128i)_f0, p0, 0, 0); + else + { + __lsx_vstelm_w((__m128i)_f0, p0, 0, 0); + __lsx_vstelm_w((__m128i)_f0, p0 + M, 0, 1); + } + p0 += M * 2; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 2; + } +#endif + for (; jj < max_jj; jj++) + { + float f0 = pp0[jj]; + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + f0 += c0; + if (broadcast_type_C == 3 || broadcast_type_C == 4) + f0 += beta == 1.f ? pC[0] : pC[0] * beta; + } + if (alpha != 1.f) + f0 *= alpha; + p0[0] = f0; + p0 += M; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC++; + } + } +} + +static void get_optimal_tile_mnk_wq_int8(int M, int N, int K, int constant_TILE_M, int constant_TILE_N, int constant_TILE_K, int& TILE_M, int& TILE_N, int& TILE_K, int nT) +{ + // resolve optimal tile size from cache size + const size_t l2_cache_size = get_cpu_level2_cache_size(); + + if (nT == 0) + nT = get_physical_big_cpu_count(); + + const int tile_size = std::max(1, (int)((float)l2_cache_size / 2 / sizeof(signed char) / std::max(1, K))); + +#if __loongarch_sx + const int tile_m_align = M >= nT * 8 ? 8 : M >= nT * 4 ? 4 : M >= nT * 2 ? 2 : 1; +#if __loongarch_asx + const int tile_n_align = tile_m_align == 8 ? 8 : 16; +#else + const int tile_n_align = tile_m_align == 8 ? 4 : 8; +#endif +#else + const int tile_m_align = M >= nT * 4 ? 4 : M >= nT * 2 ? 2 : 1; + const int tile_n_align = 2; +#endif + TILE_M = tile_m_align; + TILE_N = std::max(tile_n_align, tile_size / tile_n_align * tile_n_align); + TILE_K = K; + + if (N > 0) + { + const int nn_N = (N + TILE_N - 1) / TILE_N; + TILE_N = std::min(TILE_N, ((N + nn_N - 1) / nn_N + tile_n_align - 1) / tile_n_align * tile_n_align); + } + + // always take constant TILE_N value when provided + if (constant_TILE_N > 0) + TILE_N = (constant_TILE_N + tile_n_align - 1) / tile_n_align * tile_n_align; + + (void)constant_TILE_M; + (void)constant_TILE_K; +} diff --git a/src/layer/loongarch/multiheadattention_loongarch.cpp b/src/layer/loongarch/multiheadattention_loongarch.cpp index 7e749533d1c4..fa8f7ee481bf 100644 --- a/src/layer/loongarch/multiheadattention_loongarch.cpp +++ b/src/layer/loongarch/multiheadattention_loongarch.cpp @@ -28,10 +28,355 @@ MultiHeadAttention_loongarch::MultiHeadAttention_loongarch() o_gemm = 0; } +#if NCNN_WEIGHT_QUANT +int MultiHeadAttention_loongarch::create_pipeline_wq_int8(const Option& _opt) +{ + if (q_gemm) + return 0; + + Option opt = _opt; + Option opt_wq = opt; + opt_wq.use_packing_layout = false; + opt_wq.use_fp16_packed = false; + opt_wq.use_fp16_storage = false; + opt_wq.use_fp16_arithmetic = false; + opt_wq.use_bf16_packed = false; + opt_wq.use_bf16_storage = false; + + { + qk_softmax = ncnn::create_layer_cpu(ncnn::LayerType::Softmax); + if (!qk_softmax) + return -100; + ncnn::ParamDict pd; + pd.set(0, -1); + pd.set(1, 1); + int ret = qk_softmax->load_param(pd); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = qk_softmax->load_model(ModelBinFromMatArray(0)); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = qk_softmax->create_pipeline(opt); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + } + + const int qdim = weight_data_size / embed_dim; + + { + q_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!q_gemm) + { + destroy_pipeline(_opt); + return -100; + } + ncnn::ParamDict pd; + pd.set(0, scale); + pd.set(1, 1.f); + pd.set(2, 0); // transA + pd.set(3, 1); // transB + pd.set(4, 0); // constantA + pd.set(5, 1); // constantB + pd.set(6, 1); // constantC + pd.set(7, 0); // M + pd.set(8, embed_dim); // N + pd.set(9, qdim); // K + pd.set(10, 4); // constant_broadcast_type_C + pd.set(11, 0); // output_N1M + pd.set(12, 0); // output_elempack + pd.set(14, 1); // output_transpose + pd.set(18, quantize_term); + int ret = q_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + Mat weights[4]; + weights[0] = q_weight_data; + weights[1] = q_bias_data; + weights[2] = q_weight_data_quantize_scales; + weights[3] = q_weight_data_input_scales; + ret = q_gemm->load_model(ModelBinFromMatArray(weights)); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = q_gemm->create_pipeline(opt_wq); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + } + + { + k_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!k_gemm) + { + destroy_pipeline(_opt); + return -100; + } + ncnn::ParamDict pd; + pd.set(2, 0); // transA + pd.set(3, 1); // transB + pd.set(4, 0); // constantA + pd.set(5, 1); // constantB + pd.set(6, 1); // constantC + pd.set(7, 0); // M + pd.set(8, embed_dim); // N + pd.set(9, kdim); // K + pd.set(10, 4); // constant_broadcast_type_C + pd.set(11, 0); // output_N1M + pd.set(12, 0); // output_elempack + pd.set(14, 1); // output_transpose + pd.set(18, quantize_term); + int ret = k_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + Mat weights[4]; + weights[0] = k_weight_data; + weights[1] = k_bias_data; + weights[2] = k_weight_data_quantize_scales; + weights[3] = k_weight_data_input_scales; + ret = k_gemm->load_model(ModelBinFromMatArray(weights)); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = k_gemm->create_pipeline(opt_wq); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + } + + { + v_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!v_gemm) + { + destroy_pipeline(_opt); + return -100; + } + ncnn::ParamDict pd; + pd.set(2, 0); // transA + pd.set(3, 1); // transB + pd.set(4, 0); // constantA + pd.set(5, 1); // constantB + pd.set(6, 1); // constantC + pd.set(7, 0); // M + pd.set(8, embed_dim); // N + pd.set(9, vdim); // K + pd.set(10, 4); // constant_broadcast_type_C + pd.set(11, 0); // output_N1M + pd.set(12, 0); // output_elempack + pd.set(14, 1); // output_transpose + pd.set(18, quantize_term); + int ret = v_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + Mat weights[4]; + weights[0] = v_weight_data; + weights[1] = v_bias_data; + weights[2] = v_weight_data_quantize_scales; + weights[3] = v_weight_data_input_scales; + ret = v_gemm->load_model(ModelBinFromMatArray(weights)); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = v_gemm->create_pipeline(opt_wq); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + } + + { + o_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!o_gemm) + { + destroy_pipeline(_opt); + return -100; + } + ncnn::ParamDict pd; + pd.set(2, 1); // transA + pd.set(3, 1); // transB + pd.set(4, 0); // constantA + pd.set(5, 1); // constantB + pd.set(6, 1); // constantC + pd.set(7, 0); // M = outch + pd.set(8, qdim); // N = size + pd.set(9, embed_dim); // K = maxk*inch + pd.set(10, 4); // constant_broadcast_type_C + pd.set(11, 0); // output_N1M + pd.set(18, quantize_term); + int ret = o_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + Mat weights[4]; + weights[0] = out_weight_data; + weights[1] = out_bias_data; + weights[2] = out_weight_data_quantize_scales; + weights[3] = out_weight_data_input_scales; + ret = o_gemm->load_model(ModelBinFromMatArray(weights)); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = o_gemm->create_pipeline(opt_wq); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + } + + { + qk_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!qk_gemm) + { + destroy_pipeline(_opt); + return -100; + } + ncnn::ParamDict pd; + pd.set(2, 1); // transA + pd.set(3, 0); // transB + pd.set(4, 0); // constantA + pd.set(5, 0); // constantB + pd.set(6, attn_mask ? 0 : 1); // constantC + pd.set(7, 0); // M + pd.set(8, 0); // N + pd.set(9, 0); // K + pd.set(10, attn_mask ? 3 : -1); // constant_broadcast_type_C + pd.set(11, 0); // output_N1M + pd.set(12, 1); // output_elempack + pd.set(13, 1); // output_elemtype = fp32 + int ret = qk_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = qk_gemm->load_model(ModelBinFromMatArray(0)); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + Option opt1 = opt; + opt1.use_bf16_packed = false; + opt1.use_bf16_storage = false; + opt1.num_threads = 1; + ret = qk_gemm->create_pipeline(opt1); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + } + + { + qkv_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!qkv_gemm) + { + destroy_pipeline(_opt); + return -100; + } + ncnn::ParamDict pd; + pd.set(2, 0); // transA + pd.set(3, 1); // transB + pd.set(4, 0); // constantA + pd.set(5, 0); // constantB + pd.set(6, 1); // constantC + pd.set(7, 0); // M + pd.set(8, 0); // N + pd.set(9, 0); // K + pd.set(10, -1); // constant_broadcast_type_C + pd.set(11, 0); // output_N1M + pd.set(12, 1); // output_elempack + pd.set(13, 1); // output_elemtype = fp32 + pd.set(14, 1); // output_transpose + int ret = qkv_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = qkv_gemm->load_model(ModelBinFromMatArray(0)); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + Option opt1 = opt; + opt1.use_bf16_packed = false; + opt1.use_bf16_storage = false; + opt1.num_threads = 1; + ret = qkv_gemm->create_pipeline(opt1); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + } + + q_weight_data.release(); + q_bias_data.release(); + k_weight_data.release(); + k_bias_data.release(); + v_weight_data.release(); + v_bias_data.release(); + out_weight_data.release(); + out_bias_data.release(); + q_weight_data_quantize_scales.release(); + k_weight_data_quantize_scales.release(); + v_weight_data_quantize_scales.release(); + out_weight_data_quantize_scales.release(); + q_weight_data_input_scales.release(); + k_weight_data_input_scales.release(); + v_weight_data_input_scales.release(); + out_weight_data_input_scales.release(); + + return 0; +} +#endif // NCNN_WEIGHT_QUANT + int MultiHeadAttention_loongarch::create_pipeline(const Option& _opt) { +#if NCNN_WEIGHT_QUANT if (weight_block_quantize) - return 0; + { + if (quantize_term / 100 != 8) + return MultiHeadAttention::create_pipeline(_opt); + + return create_pipeline_wq_int8(_opt); + } +#endif Option opt = _opt; if (int8_scale_term) @@ -259,14 +604,24 @@ int MultiHeadAttention_loongarch::create_pipeline(const Option& _opt) int MultiHeadAttention_loongarch::destroy_pipeline(const Option& _opt) { - if (weight_block_quantize) - return 0; + if (weight_block_quantize && quantize_term / 100 != 8) + return MultiHeadAttention::destroy_pipeline(_opt); Option opt = _opt; - if (int8_scale_term) + if (int8_scale_term && !weight_block_quantize) { opt.use_packing_layout = false; // TODO enable packing } + Option opt_wq = opt; + if (weight_block_quantize) + { + opt_wq.use_packing_layout = false; + opt_wq.use_fp16_packed = false; + opt_wq.use_fp16_storage = false; + opt_wq.use_fp16_arithmetic = false; + opt_wq.use_bf16_packed = false; + opt_wq.use_bf16_storage = false; + } if (qk_softmax) { @@ -277,28 +632,28 @@ int MultiHeadAttention_loongarch::destroy_pipeline(const Option& _opt) if (q_gemm) { - q_gemm->destroy_pipeline(opt); + q_gemm->destroy_pipeline(opt_wq); delete q_gemm; q_gemm = 0; } if (k_gemm) { - k_gemm->destroy_pipeline(opt); + k_gemm->destroy_pipeline(opt_wq); delete k_gemm; k_gemm = 0; } if (v_gemm) { - v_gemm->destroy_pipeline(opt); + v_gemm->destroy_pipeline(opt_wq); delete v_gemm; v_gemm = 0; } if (o_gemm) { - o_gemm->destroy_pipeline(opt); + o_gemm->destroy_pipeline(opt_wq); delete o_gemm; o_gemm = 0; } @@ -321,7 +676,7 @@ int MultiHeadAttention_loongarch::destroy_pipeline(const Option& _opt) int MultiHeadAttention_loongarch::forward(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& _opt) const { - if (weight_block_quantize) + if (weight_block_quantize && quantize_term / 100 != 8) return MultiHeadAttention::forward(bottom_blobs, top_blobs, _opt); int q_blob_i = 0; @@ -340,10 +695,20 @@ int MultiHeadAttention_loongarch::forward(const std::vector& bottom_blobs, const Mat& cached_xv_blob = kv_cache ? bottom_blobs[cached_xv_i] : Mat(); Option opt = _opt; - if (int8_scale_term) + if (int8_scale_term && !weight_block_quantize) { opt.use_packing_layout = false; // TODO enable packing } + Option opt_wq = opt; + if (weight_block_quantize) + { + opt_wq.use_packing_layout = false; + opt_wq.use_fp16_packed = false; + opt_wq.use_fp16_storage = false; + opt_wq.use_fp16_arithmetic = false; + opt_wq.use_bf16_packed = false; + opt_wq.use_bf16_storage = false; + } Mat attn_mask_blob_unpacked; if (attn_mask && attn_mask_blob.elempack != 1) @@ -388,7 +753,7 @@ int MultiHeadAttention_loongarch::forward(const std::vector& bottom_blobs, const int dst_seqlen = past_seqlen > 0 ? (q_blob_i == k_blob_i ? (past_seqlen + cur_seqlen) : past_seqlen) : cur_seqlen; Mat q_affine; - int retq = q_gemm->forward(q_blob, q_affine, opt); + int retq = q_gemm->forward(q_blob, q_affine, opt_wq); if (retq != 0) return retq; @@ -398,7 +763,7 @@ int MultiHeadAttention_loongarch::forward(const std::vector& bottom_blobs, if (q_blob_i == k_blob_i) { Mat k_affine_q; - int retk = k_gemm->forward(q_blob, k_affine_q, opt); + int retk = k_gemm->forward(q_blob, k_affine_q, opt_wq); if (retk != 0) return retk; @@ -426,7 +791,7 @@ int MultiHeadAttention_loongarch::forward(const std::vector& bottom_blobs, } else { - int retk = k_gemm->forward(k_blob, k_affine, opt); + int retk = k_gemm->forward(k_blob, k_affine, opt_wq); if (retk != 0) return retk; } @@ -477,7 +842,7 @@ int MultiHeadAttention_loongarch::forward(const std::vector& bottom_blobs, if (q_blob_i == v_blob_i) { Mat v_affine_q; - int retk = v_gemm->forward(v_blob, v_affine_q, opt); + int retk = v_gemm->forward(v_blob, v_affine_q, opt_wq); if (retk != 0) return retk; @@ -505,7 +870,7 @@ int MultiHeadAttention_loongarch::forward(const std::vector& bottom_blobs, } else { - int retv = v_gemm->forward(v_blob, v_affine, opt); + int retv = v_gemm->forward(v_blob, v_affine, opt_wq); if (retv != 0) return retv; } @@ -552,7 +917,7 @@ int MultiHeadAttention_loongarch::forward(const std::vector& bottom_blobs, v_affine.release(); } - int reto = o_gemm->forward(qkv_cross, top_blobs[0], opt); + int reto = o_gemm->forward(qkv_cross, top_blobs[0], opt_wq); if (reto != 0) return reto; diff --git a/src/layer/loongarch/multiheadattention_loongarch.h b/src/layer/loongarch/multiheadattention_loongarch.h index fe13c097511c..f4b729edd294 100644 --- a/src/layer/loongarch/multiheadattention_loongarch.h +++ b/src/layer/loongarch/multiheadattention_loongarch.h @@ -18,6 +18,11 @@ class MultiHeadAttention_loongarch : public MultiHeadAttention virtual int forward(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& opt) const; +protected: +#if NCNN_WEIGHT_QUANT + int create_pipeline_wq_int8(const Option& opt); +#endif + public: Layer* q_gemm; Layer* k_gemm; diff --git a/src/layer/mips/gemm_mips.cpp b/src/layer/mips/gemm_mips.cpp index 81f0cb108b5d..d0a696e4c398 100644 --- a/src/layer/mips/gemm_mips.cpp +++ b/src/layer/mips/gemm_mips.cpp @@ -23,6 +23,10 @@ namespace ncnn { #include "gemm_int8.h" #endif +#if NCNN_WEIGHT_QUANT +#include "gemm_wq_int8.h" +#endif + Gemm_mips::Gemm_mips() { #if __mips_msa @@ -4470,6 +4474,200 @@ static int gemm_AT_BT_mips(const Mat& AT, const Mat& BT, const Mat& C, Mat& top_ return 0; } +#if NCNN_WEIGHT_QUANT +static int gemm_BT_mips_wq_int8(const Mat& A, const Mat& packed_B, const Mat& packed_B_descales, const Mat& input_scales, const Mat& C, Mat& top_blob, int broadcast_type_C, int N, int K, int block_size, int transA, int output_transpose, float alpha, float beta, int constant_TILE_M, int constant_TILE_N, int constant_TILE_K, int nT, const Option& opt) +{ + const int M = transA ? A.w : (A.dims == 3 ? A.c : A.h) * A.elempack; + const int block_count = (K + block_size - 1) / block_size; + int TILE_M, TILE_N, TILE_K; + get_optimal_tile_mnk_wq_int8(M, N, K, constant_TILE_M, constant_TILE_N, constant_TILE_K, TILE_M, TILE_N, TILE_K, nT); + + (void)TILE_K; + const int mr = std::min(M, TILE_M); + const int nr = std::min(N, TILE_N); + const int nn_M = (M + TILE_M - 1) / TILE_M; + const int nn_N = (N + TILE_N - 1) / TILE_N; + const float* input_scale_ptr = input_scales; + Mat BT = packed_B.reshape(K, N); + Mat BT_descales = packed_B_descales.reshape(block_count, N); + + Mat topT(mr * nr, 1, nT, (size_t)4u, 1, opt.workspace_allocator); + if (topT.empty()) + return -100; + + if (nT > nn_M) + { + Mat AT(K, mr, nn_M, (size_t)1u, 1, opt.workspace_allocator); + Mat AT_descales(block_count, mr, nn_M, (size_t)4u, 1, opt.workspace_allocator); + if (AT.empty() || AT_descales.empty()) + return -100; + + #pragma omp parallel for num_threads(nT) + for (int ppi = 0; ppi < nn_M; ppi++) + { + const int i = ppi * TILE_M; + const int max_ii = std::min(M - i, TILE_M); + + Mat AT_tile = AT.channel(i / TILE_M).row_range(0, max_ii); + Mat AT_descales_tile = AT_descales.channel(i / TILE_M).row_range(0, max_ii); + + if (transA) + transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, block_size, input_scale_ptr); + else + quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, block_size, input_scale_ptr); + } + + const int nn_MN = nn_M * nn_N; + #pragma omp parallel for num_threads(nT) + for (int ppij = 0; ppij < nn_MN; ppij++) + { + const int ppi = ppij / nn_N; + const int ppj = ppij % nn_N; + + const int i = ppi * TILE_M; + const int j = ppj * TILE_N; + const int max_ii = std::min(M - i, TILE_M); + const int max_jj = std::min(N - j, TILE_N); + + Mat AT_tile = AT.channel(i / TILE_M).row_range(0, max_ii); + Mat AT_descales_tile = AT_descales.channel(i / TILE_M).row_range(0, max_ii); + Mat BT_tile = BT.row_range(j, max_jj); + Mat BT_descales_tile = BT_descales.row_range(j, max_jj); + Mat topT_tile = topT.channel(get_omp_thread_num()); + + gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, block_size); + + if (output_transpose) + transpose_unpack_output_tile_wq_int8(topT_tile, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, alpha, beta); + else + unpack_output_tile_wq_int8(topT_tile, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, alpha, beta); + } + } + else + { + Mat ATX(K, mr, nT, (size_t)1u, 1, opt.workspace_allocator); + Mat ATX_descales(block_count, mr, nT, (size_t)4u, 1, opt.workspace_allocator); + if (ATX.empty() || ATX_descales.empty()) + return -100; + + #pragma omp parallel for num_threads(nT) + for (int ppi = 0; ppi < nn_M; ppi++) + { + const int i = ppi * TILE_M; + const int max_ii = std::min(M - i, TILE_M); + + Mat AT_tile = ATX.channel(get_omp_thread_num()).row_range(0, max_ii); + Mat AT_descales_tile = ATX_descales.channel(get_omp_thread_num()).row_range(0, max_ii); + Mat topT_tile = topT.channel(get_omp_thread_num()); + + if (transA) + transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, block_size, input_scale_ptr); + else + quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, block_size, input_scale_ptr); + + for (int j = 0; j < N; j += TILE_N) + { + const int max_jj = std::min(N - j, TILE_N); + Mat BT_tile = BT.row_range(j, max_jj); + Mat BT_descales_tile = BT_descales.row_range(j, max_jj); + gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, block_size); + + if (output_transpose) + transpose_unpack_output_tile_wq_int8(topT_tile, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, alpha, beta); + else + unpack_output_tile_wq_int8(topT_tile, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, alpha, beta); + } + } + } + + return 0; +} + +static int gemm_weight_quantize_bits_mips(int quantize_term) +{ + return quantize_term / 100; +} + +static int gemm_weight_quantize_block_size_mips(int quantize_term) +{ + const int block_size_code = quantize_term % 10; + if (block_size_code == 0) + return 32; + if (block_size_code == 1) + return 64; + if (block_size_code == 2) + return 128; + return 0; +} + +int Gemm_mips::forward_weight_block_quantize_int8(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& opt) const +{ + const Mat& A = bottom_blobs[0]; + if (A.elemsize != 4u || A.elempack != 1 || (transA && A.dims != 2)) + { + NCNN_LOGE("Gemm unsupported input"); + return -1; + } + + const int K = transA ? A.h : A.w; + if (K != constantK) + { + NCNN_LOGE("Gemm weight block quantize K mismatch"); + return -1; + } + + const int M = transA ? A.w : A.dims == 3 ? A.c : A.h; + const int N = constantN; + const int block_size = gemm_weight_quantize_block_size_mips(quantize_term); + + Mat C; + int broadcast_type_C = -1; + if (constantC) + { + C = C_data; + broadcast_type_C = constant_broadcast_type_C; + } + else + { + if (bottom_blobs.size() == 2) + C = bottom_blobs[1]; + + if (!C.empty()) + { + bool matched = false; + if (C.dims == 1 && C.w == 1) { broadcast_type_C = 0; matched = true; } + if (C.dims == 1 && C.w == M) { broadcast_type_C = 1; matched = true; } + if (C.dims == 1 && C.w == N) { broadcast_type_C = 4; matched = true; } + if (C.dims == 2 && C.w == 1 && C.h == M) { broadcast_type_C = 2; matched = true; } + if (C.dims == 2 && C.w == N && C.h == M) { broadcast_type_C = 3; matched = true; } + if (C.dims == 2 && C.w == N && C.h == 1) { broadcast_type_C = 4; matched = true; } + + if (!matched || C.elemsize != 4u || C.elempack != 1) + { + NCNN_LOGE("Gemm unsupported C"); + return -1; + } + } + } + + if (!C.empty() && (C.elemsize != 4u || C.elempack != 1)) + { + NCNN_LOGE("Gemm unsupported C"); + return -1; + } + + Mat& top_blob = top_blobs[0]; + if (output_transpose) + top_blob.create(M, N, (size_t)4u, opt.blob_allocator); + else + top_blob.create(N, M, (size_t)4u, opt.blob_allocator); + if (top_blob.empty()) + return -100; + + return gemm_BT_mips_wq_int8(A, B_data_w8a8_packed, B_data_w8a8_descales, B_data_input_scales, C, top_blob, broadcast_type_C, N, K, block_size, transA, output_transpose, alpha, beta, constant_TILE_M, constant_TILE_N, constant_TILE_K, opt.num_threads, opt); +} +#endif // NCNN_WEIGHT_QUANT + int Gemm_mips::create_pipeline(const Option& opt) { AT_data.release(); @@ -4479,6 +4677,28 @@ int Gemm_mips::create_pipeline(const Option& opt) if (weight_block_quantize) { +#if NCNN_WEIGHT_QUANT + if (gemm_weight_quantize_bits_mips(quantize_term) == 8) + { + if (!B_data_w8a8_packed.empty()) + return 0; + if (B_data.empty() || B_data_quantize_scales.empty()) + return -1; + + Mat B_data_packed; + Mat B_data_descales; + int ret = pack_B_wq_int8(B_data, B_data_quantize_scales, B_data_packed, B_data_descales, constantN, constantK, gemm_weight_quantize_block_size_mips(quantize_term), opt); + if (ret != 0) + return ret; + + B_data_w8a8_packed = B_data_packed; + B_data_w8a8_descales = B_data_descales; + + B_data.release(); + B_data_quantize_scales.release(); + } +#endif + return 0; } @@ -4616,10 +4836,25 @@ int Gemm_mips::create_pipeline(const Option& opt) return 0; } +int Gemm_mips::destroy_pipeline(const Option& opt) +{ +#if NCNN_WEIGHT_QUANT + B_data_w8a8_packed.release(); + B_data_w8a8_descales.release(); +#endif + + return Gemm::destroy_pipeline(opt); +} + int Gemm_mips::forward(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& opt) const { if (weight_block_quantize) { +#if NCNN_WEIGHT_QUANT + if (gemm_weight_quantize_bits_mips(quantize_term) == 8 && !B_data_w8a8_packed.empty()) + return forward_weight_block_quantize_int8(bottom_blobs, top_blobs, opt); +#endif + return Gemm::forward(bottom_blobs, top_blobs, opt); } diff --git a/src/layer/mips/gemm_mips.h b/src/layer/mips/gemm_mips.h index ff043262c130..521be14cd717 100644 --- a/src/layer/mips/gemm_mips.h +++ b/src/layer/mips/gemm_mips.h @@ -15,9 +15,14 @@ class Gemm_mips : public Gemm virtual int create_pipeline(const Option& opt); + virtual int destroy_pipeline(const Option& opt); + virtual int forward(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& opt) const; protected: +#if NCNN_WEIGHT_QUANT + int forward_weight_block_quantize_int8(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& opt) const; +#endif #if NCNN_INT8 int create_pipeline_int8(const Option& opt); int forward_int8(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& opt) const; @@ -33,6 +38,10 @@ class Gemm_mips : public Gemm Mat AT_data; Mat BT_data; Mat CT_data; +#if NCNN_WEIGHT_QUANT + Mat B_data_w8a8_packed; + Mat B_data_w8a8_descales; +#endif }; // expose some gemm internal routines for convolution uses diff --git a/src/layer/mips/gemm_mips_mmi.cpp b/src/layer/mips/gemm_mips_mmi.cpp index 09bae1c3979d..85fdf18c06f0 100644 --- a/src/layer/mips/gemm_mips_mmi.cpp +++ b/src/layer/mips/gemm_mips_mmi.cpp @@ -11,6 +11,9 @@ namespace ncnn { #include "gemm_int8.h" +#if NCNN_WEIGHT_QUANT +#include "gemm_wq_int8.h" +#endif void pack_A_tile_int8_loongson_mmi(const Mat& A, Mat& AT, int i, int max_ii, int k, int max_kk) { @@ -57,4 +60,26 @@ void gemm_transB_packed_tile_int8_loongson_mmi(const Mat& AT_tile, const Mat& BT gemm_transB_packed_tile_int8(AT_tile, BT_tile, topT_tile, i, max_ii, j, max_jj, k, max_kk); } +#if NCNN_WEIGHT_QUANT +int pack_B_wq_int8_loongson_mmi(const Mat& B, const Mat& B_scales, Mat& packed_B, Mat& packed_B_descales, int N, int K, int block_size, const Option& opt) +{ + return pack_B_wq_int8(B, B_scales, packed_B, packed_B_descales, N, K, block_size, opt); +} + +void quantize_A_tile_wq_int8_loongson_mmi(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int block_size, const float* input_scale_ptr) +{ + quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, block_size, input_scale_ptr); +} + +void transpose_quantize_A_tile_wq_int8_loongson_mmi(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int block_size, const float* input_scale_ptr) +{ + transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, block_size, input_scale_ptr); +} + +void gemm_transB_packed_tile_wq_int8_loongson_mmi(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int block_size) +{ + gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, block_size); +} +#endif // NCNN_WEIGHT_QUANT + } // namespace ncnn diff --git a/src/layer/mips/gemm_wq_int8.h b/src/layer/mips/gemm_wq_int8.h new file mode 100644 index 000000000000..547c1521c825 --- /dev/null +++ b/src/layer/mips/gemm_wq_int8.h @@ -0,0 +1,5273 @@ +// Copyright 2026 Tencent +// SPDX-License-Identifier: BSD-3-Clause + +#if NCNN_RUNTIME_CPU && NCNN_MMI && !__mips_msa && !__mips_loongson_mmi +int pack_B_wq_int8_loongson_mmi(const Mat& B, const Mat& B_scales, Mat& packed_B, Mat& packed_B_descales, int N, int K, int block_size, const Option& opt); +void quantize_A_tile_wq_int8_loongson_mmi(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int block_size, const float* input_scale_ptr); +void transpose_quantize_A_tile_wq_int8_loongson_mmi(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int block_size, const float* input_scale_ptr); +void gemm_transB_packed_tile_wq_int8_loongson_mmi(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int block_size); +#endif + +// group-major, output-major within each K4/K2/K1 fragment +static int pack_B_wq_int8(const Mat& B, const Mat& B_scales, Mat& packed_B, Mat& packed_B_descales, int N, int K, int block_size, const Option& opt) +{ +#if NCNN_RUNTIME_CPU && NCNN_MMI && !__mips_msa && !__mips_loongson_mmi + if (ncnn::cpu_support_loongson_mmi()) + return pack_B_wq_int8_loongson_mmi(B, B_scales, packed_B, packed_B_descales, N, K, block_size, opt); +#endif + + const int block_count = (K + block_size - 1) / block_size; + + packed_B.create(N * K, (size_t)1u, opt.blob_allocator); + if (packed_B.empty()) + return -100; + packed_B.cstep = (size_t)N * K; + + packed_B_descales.create(N * block_count, (size_t)4u, opt.blob_allocator); + if (packed_B_descales.empty()) + return -100; + packed_B_descales.cstep = (size_t)N * block_count; + + int j = 0; +#if __mips_msa + const int nn8 = N / 8; + #pragma omp parallel for num_threads(opt.num_threads) + for (int ppj = 0; ppj < nn8; ppj++) + { + const int j = ppj * 8; + signed char* pp = (signed char*)packed_B + (size_t)j * K; + float* pd = (float*)packed_B_descales + (size_t)j * block_count; + + for (int jj = 0; jj < 8; jj += 4) + { + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + const signed char* p0 = B.row(j + jj) + k0 + kk; + const signed char* p1 = B.row(j + jj + 1) + k0 + kk; + const signed char* p2 = B.row(j + jj + 2) + k0 + kk; + const signed char* p3 = B.row(j + jj + 3) + k0 + kk; + const v16i8 _p = (v16i8)__msa_set_w(__msa_load_w(p0), __msa_load_w(p1), __msa_load_w(p2), __msa_load_w(p3)); + __msa_st_b(_p, pp, 0); + pp += 16; + } + if (kk + 1 < max_kk) + { + for (int n = 0; n < 4; n++) + { + const signed char* p0 = B.row(j + jj + n) + k0 + kk; + pp[0] = p0[0]; + pp[1] = p0[1]; + pp += 2; + } + kk += 2; + } + if (kk < max_kk) + { + for (int n = 0; n < 4; n++) + *pp++ = B.row(j + jj + n)[k0 + kk]; + } + + for (int n = 0; n < 4; n++) + *pd++ = 1.f / B_scales.row(j + jj + n)[g]; + } + } + } + j += nn8 * 8; + + const int nn4 = (N - j) / 4; + #pragma omp parallel for num_threads(opt.num_threads) + for (int ppj = 0; ppj < nn4; ppj++) + { + const int j = nn8 * 8 + ppj * 4; + signed char* pp = (signed char*)packed_B + (size_t)j * K; + float* pd = (float*)packed_B_descales + (size_t)j * block_count; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + const signed char* p0 = B.row(j) + k0 + kk; + const signed char* p1 = B.row(j + 1) + k0 + kk; + const signed char* p2 = B.row(j + 2) + k0 + kk; + const signed char* p3 = B.row(j + 3) + k0 + kk; + const v16i8 _p = (v16i8)__msa_set_w(__msa_load_w(p0), __msa_load_w(p1), __msa_load_w(p2), __msa_load_w(p3)); + __msa_st_b(_p, pp, 0); + pp += 16; + } + if (kk + 1 < max_kk) + { + for (int n = 0; n < 4; n++) + { + const signed char* p0 = B.row(j + n) + k0 + kk; + pp[0] = p0[0]; + pp[1] = p0[1]; + pp += 2; + } + kk += 2; + } + if (kk < max_kk) + { + for (int n = 0; n < 4; n++) + *pp++ = B.row(j + n)[k0 + kk]; + } + + for (int n = 0; n < 4; n++) + *pd++ = 1.f / B_scales.row(j + n)[g]; + } + } + j += nn4 * 4; +#endif // __mips_msa + + const int nn2 = (N - j) / 2; + const int j2 = j; + #pragma omp parallel for num_threads(opt.num_threads) + for (int ppj = 0; ppj < nn2; ppj++) + { + const int j = j2 + ppj * 2; + signed char* pp = (signed char*)packed_B + (size_t)j * K; + float* pd = (float*)packed_B_descales + (size_t)j * block_count; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + const signed char* p0 = B.row(j) + k0 + kk; + const signed char* p1 = B.row(j + 1) + k0 + kk; + pp[0] = p0[0]; + pp[1] = p0[1]; + pp[2] = p0[2]; + pp[3] = p0[3]; + pp[4] = p1[0]; + pp[5] = p1[1]; + pp[6] = p1[2]; + pp[7] = p1[3]; + pp += 8; + } + if (kk + 1 < max_kk) + { + const signed char* p0 = B.row(j) + k0 + kk; + const signed char* p1 = B.row(j + 1) + k0 + kk; + pp[0] = p0[0]; + pp[1] = p0[1]; + pp[2] = p1[0]; + pp[3] = p1[1]; + pp += 4; + kk += 2; + } + if (kk < max_kk) + { + pp[0] = B.row(j)[k0 + kk]; + pp[1] = B.row(j + 1)[k0 + kk]; + pp += 2; + } + + *pd++ = 1.f / B_scales.row(j)[g]; + *pd++ = 1.f / B_scales.row(j + 1)[g]; + } + } + j += nn2 * 2; + + if (j < N) + { + signed char* pp = (signed char*)packed_B + (size_t)j * K; + float* pd = (float*)packed_B_descales + (size_t)j * block_count; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + const signed char* p0 = B.row(j) + k0; + + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + pp[0] = p0[kk]; + pp[1] = p0[kk + 1]; + pp[2] = p0[kk + 2]; + pp[3] = p0[kk + 3]; + pp += 4; + } + if (kk + 1 < max_kk) + { + pp[0] = p0[kk]; + pp[1] = p0[kk + 1]; + pp += 2; + kk += 2; + } + if (kk < max_kk) + *pp++ = p0[kk]; + + *pd++ = 1.f / B_scales.row(j)[g]; + } + } + + return 0; +} + +// group-major, row-major within each K4/K2/K1 fragment +static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int block_size, const float* input_scale_ptr) +{ +#if NCNN_RUNTIME_CPU && NCNN_MMI && !__mips_msa && !__mips_loongson_mmi + if (ncnn::cpu_support_loongson_mmi()) + { + quantize_A_tile_wq_int8_loongson_mmi(A, AT_tile, AT_descales_tile, i, max_ii, block_size, input_scale_ptr); + return; + } +#endif + + signed char* pp = AT_tile; + float* pd = AT_descales_tile; + const int K = AT_tile.w; + const int block_count = AT_descales_tile.w; + const size_t A_hstep = A.dims == 3 ? A.cstep : (size_t)A.w; + + int ii = 0; +#if __mips_msa + for (; ii + 7 < max_ii; ii += 8) + { + const float* p0 = (const float*)A + (size_t)(i + ii) * A_hstep; + const float* p1 = (const float*)A + (size_t)(i + ii + 1) * A_hstep; + const float* p2 = (const float*)A + (size_t)(i + ii + 2) * A_hstep; + const float* p3 = (const float*)A + (size_t)(i + ii + 3) * A_hstep; + const float* p4 = (const float*)A + (size_t)(i + ii + 4) * A_hstep; + const float* p5 = (const float*)A + (size_t)(i + ii + 5) * A_hstep; + const float* p6 = (const float*)A + (size_t)(i + ii + 6) * A_hstep; + const float* p7 = (const float*)A + (size_t)(i + ii + 7) * A_hstep; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + const v16u8 _abs_mask = (v16u8)__msa_fill_w(0x7fffffff); + v4f32 _absmax0 = (v4f32)__msa_fill_w(0); + v4f32 _absmax1 = (v4f32)__msa_fill_w(0); + v4f32 _absmax2 = (v4f32)__msa_fill_w(0); + v4f32 _absmax3 = (v4f32)__msa_fill_w(0); + v4f32 _absmax4 = (v4f32)__msa_fill_w(0); + v4f32 _absmax5 = (v4f32)__msa_fill_w(0); + v4f32 _absmax6 = (v4f32)__msa_fill_w(0); + v4f32 _absmax7 = (v4f32)__msa_fill_w(0); + + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + v4f32 _p0 = (v4f32)__msa_ld_w(p0 + k0 + kk, 0); + v4f32 _p1 = (v4f32)__msa_ld_w(p1 + k0 + kk, 0); + v4f32 _p2 = (v4f32)__msa_ld_w(p2 + k0 + kk, 0); + v4f32 _p3 = (v4f32)__msa_ld_w(p3 + k0 + kk, 0); + v4f32 _p4 = (v4f32)__msa_ld_w(p4 + k0 + kk, 0); + v4f32 _p5 = (v4f32)__msa_ld_w(p5 + k0 + kk, 0); + v4f32 _p6 = (v4f32)__msa_ld_w(p6 + k0 + kk, 0); + v4f32 _p7 = (v4f32)__msa_ld_w(p7 + k0 + kk, 0); + if (input_scale_ptr) + { + const v4f32 _s = (v4f32)__msa_ld_w(input_scale_ptr + k0 + kk, 0); + _p0 = __msa_fmul_w(_p0, _s); + _p1 = __msa_fmul_w(_p1, _s); + _p2 = __msa_fmul_w(_p2, _s); + _p3 = __msa_fmul_w(_p3, _s); + _p4 = __msa_fmul_w(_p4, _s); + _p5 = __msa_fmul_w(_p5, _s); + _p6 = __msa_fmul_w(_p6, _s); + _p7 = __msa_fmul_w(_p7, _s); + } + _absmax0 = __msa_fmax_w(_absmax0, (v4f32)__msa_and_v((v16u8)_p0, _abs_mask)); + _absmax1 = __msa_fmax_w(_absmax1, (v4f32)__msa_and_v((v16u8)_p1, _abs_mask)); + _absmax2 = __msa_fmax_w(_absmax2, (v4f32)__msa_and_v((v16u8)_p2, _abs_mask)); + _absmax3 = __msa_fmax_w(_absmax3, (v4f32)__msa_and_v((v16u8)_p3, _abs_mask)); + _absmax4 = __msa_fmax_w(_absmax4, (v4f32)__msa_and_v((v16u8)_p4, _abs_mask)); + _absmax5 = __msa_fmax_w(_absmax5, (v4f32)__msa_and_v((v16u8)_p5, _abs_mask)); + _absmax6 = __msa_fmax_w(_absmax6, (v4f32)__msa_and_v((v16u8)_p6, _abs_mask)); + _absmax7 = __msa_fmax_w(_absmax7, (v4f32)__msa_and_v((v16u8)_p7, _abs_mask)); + } + + float absmax0 = __msa_reduce_fmax_w(_absmax0); + float absmax1 = __msa_reduce_fmax_w(_absmax1); + float absmax2 = __msa_reduce_fmax_w(_absmax2); + float absmax3 = __msa_reduce_fmax_w(_absmax3); + float absmax4 = __msa_reduce_fmax_w(_absmax4); + float absmax5 = __msa_reduce_fmax_w(_absmax5); + float absmax6 = __msa_reduce_fmax_w(_absmax6); + float absmax7 = __msa_reduce_fmax_w(_absmax7); + + for (; kk < max_kk; kk++) + { + const int k = k0 + kk; + const float s = input_scale_ptr ? input_scale_ptr[k] : 1.f; + absmax0 = std::max(absmax0, fabsf(p0[k] * s)); + absmax1 = std::max(absmax1, fabsf(p1[k] * s)); + absmax2 = std::max(absmax2, fabsf(p2[k] * s)); + absmax3 = std::max(absmax3, fabsf(p3[k] * s)); + absmax4 = std::max(absmax4, fabsf(p4[k] * s)); + absmax5 = std::max(absmax5, fabsf(p5[k] * s)); + absmax6 = std::max(absmax6, fabsf(p6[k] * s)); + absmax7 = std::max(absmax7, fabsf(p7[k] * s)); + } + + volatile double scale0_fp64 = absmax0 == 0.f ? 1.0 : 127.0 / (double)absmax0; + volatile double scale1_fp64 = absmax1 == 0.f ? 1.0 : 127.0 / (double)absmax1; + volatile double scale2_fp64 = absmax2 == 0.f ? 1.0 : 127.0 / (double)absmax2; + volatile double scale3_fp64 = absmax3 == 0.f ? 1.0 : 127.0 / (double)absmax3; + volatile double scale4_fp64 = absmax4 == 0.f ? 1.0 : 127.0 / (double)absmax4; + volatile double scale5_fp64 = absmax5 == 0.f ? 1.0 : 127.0 / (double)absmax5; + volatile double scale6_fp64 = absmax6 == 0.f ? 1.0 : 127.0 / (double)absmax6; + volatile double scale7_fp64 = absmax7 == 0.f ? 1.0 : 127.0 / (double)absmax7; + const float scale0 = (float)scale0_fp64; + const float scale1 = (float)scale1_fp64; + const float scale2 = (float)scale2_fp64; + const float scale3 = (float)scale3_fp64; + const float scale4 = (float)scale4_fp64; + const float scale5 = (float)scale5_fp64; + const float scale6 = (float)scale6_fp64; + const float scale7 = (float)scale7_fp64; + pd[0] = absmax0 / 127.f; + pd[1] = absmax1 / 127.f; + pd[2] = absmax2 / 127.f; + pd[3] = absmax3 / 127.f; + pd[4] = absmax4 / 127.f; + pd[5] = absmax5 / 127.f; + pd[6] = absmax6 / 127.f; + pd[7] = absmax7 / 127.f; + pd += 8; + + kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + const v4f32 _s = input_scale_ptr ? (v4f32)__msa_ld_w(input_scale_ptr + k0 + kk, 0) : __msa_fill_w_f32(1.f); + v4f32 _p = (v4f32)__msa_ld_w(p0 + k0 + kk, 0); + _p = __msa_fmul_w(__msa_fmul_w(_p, _s), __msa_fill_w_f32(scale0)); + ((int*)pp)[0] = __msa_copy_s_w((v4i32)float2int8(_p), 0); + _p = (v4f32)__msa_ld_w(p1 + k0 + kk, 0); + _p = __msa_fmul_w(__msa_fmul_w(_p, _s), __msa_fill_w_f32(scale1)); + ((int*)pp)[1] = __msa_copy_s_w((v4i32)float2int8(_p), 0); + _p = (v4f32)__msa_ld_w(p2 + k0 + kk, 0); + _p = __msa_fmul_w(__msa_fmul_w(_p, _s), __msa_fill_w_f32(scale2)); + ((int*)pp)[2] = __msa_copy_s_w((v4i32)float2int8(_p), 0); + _p = (v4f32)__msa_ld_w(p3 + k0 + kk, 0); + _p = __msa_fmul_w(__msa_fmul_w(_p, _s), __msa_fill_w_f32(scale3)); + ((int*)pp)[3] = __msa_copy_s_w((v4i32)float2int8(_p), 0); + _p = (v4f32)__msa_ld_w(p4 + k0 + kk, 0); + _p = __msa_fmul_w(__msa_fmul_w(_p, _s), __msa_fill_w_f32(scale4)); + ((int*)pp)[4] = __msa_copy_s_w((v4i32)float2int8(_p), 0); + _p = (v4f32)__msa_ld_w(p5 + k0 + kk, 0); + _p = __msa_fmul_w(__msa_fmul_w(_p, _s), __msa_fill_w_f32(scale5)); + ((int*)pp)[5] = __msa_copy_s_w((v4i32)float2int8(_p), 0); + _p = (v4f32)__msa_ld_w(p6 + k0 + kk, 0); + _p = __msa_fmul_w(__msa_fmul_w(_p, _s), __msa_fill_w_f32(scale6)); + ((int*)pp)[6] = __msa_copy_s_w((v4i32)float2int8(_p), 0); + _p = (v4f32)__msa_ld_w(p7 + k0 + kk, 0); + _p = __msa_fmul_w(__msa_fmul_w(_p, _s), __msa_fill_w_f32(scale7)); + ((int*)pp)[7] = __msa_copy_s_w((v4i32)float2int8(_p), 0); + pp += 32; + } + if (kk + 1 < max_kk) + { + const int k = k0 + kk; + const float s0 = input_scale_ptr ? input_scale_ptr[k] : 1.f; + const float s1 = input_scale_ptr ? input_scale_ptr[k + 1] : 1.f; + pp[0] = float2int8(p0[k] * s0 * scale0); + pp[1] = float2int8(p0[k + 1] * s1 * scale0); + pp[2] = float2int8(p1[k] * s0 * scale1); + pp[3] = float2int8(p1[k + 1] * s1 * scale1); + pp[4] = float2int8(p2[k] * s0 * scale2); + pp[5] = float2int8(p2[k + 1] * s1 * scale2); + pp[6] = float2int8(p3[k] * s0 * scale3); + pp[7] = float2int8(p3[k + 1] * s1 * scale3); + pp[8] = float2int8(p4[k] * s0 * scale4); + pp[9] = float2int8(p4[k + 1] * s1 * scale4); + pp[10] = float2int8(p5[k] * s0 * scale5); + pp[11] = float2int8(p5[k + 1] * s1 * scale5); + pp[12] = float2int8(p6[k] * s0 * scale6); + pp[13] = float2int8(p6[k + 1] * s1 * scale6); + pp[14] = float2int8(p7[k] * s0 * scale7); + pp[15] = float2int8(p7[k + 1] * s1 * scale7); + pp += 16; + kk += 2; + } + if (kk < max_kk) + { + const int k = k0 + kk; + const float s = input_scale_ptr ? input_scale_ptr[k] : 1.f; + pp[0] = float2int8(p0[k] * s * scale0); + pp[1] = float2int8(p1[k] * s * scale1); + pp[2] = float2int8(p2[k] * s * scale2); + pp[3] = float2int8(p3[k] * s * scale3); + pp[4] = float2int8(p4[k] * s * scale4); + pp[5] = float2int8(p5[k] * s * scale5); + pp[6] = float2int8(p6[k] * s * scale6); + pp[7] = float2int8(p7[k] * s * scale7); + pp += 8; + } + } + } + for (; ii + 3 < max_ii; ii += 4) + { + const int i0 = i + ii; + const int i1 = i + ii + 1; + const int i2 = i + ii + 2; + const int i3 = i + ii + 3; + const float* p0 = (const float*)A + (size_t)i0 * A_hstep; + const float* p1 = (const float*)A + (size_t)i1 * A_hstep; + const float* p2 = (const float*)A + (size_t)i2 * A_hstep; + const float* p3 = (const float*)A + (size_t)i3 * A_hstep; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + float absmax0 = 0.f; + float absmax1 = 0.f; + float absmax2 = 0.f; + float absmax3 = 0.f; + + const v16u8 _abs_mask = (v16u8)__msa_fill_w(0x7fffffff); + v4f32 _absmax0 = (v4f32)__msa_fill_w(0); + v4f32 _absmax1 = (v4f32)__msa_fill_w(0); + v4f32 _absmax2 = (v4f32)__msa_fill_w(0); + v4f32 _absmax3 = (v4f32)__msa_fill_w(0); + + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + v4f32 _p0 = (v4f32)__msa_ld_w(p0 + k0 + kk, 0); + v4f32 _p1 = (v4f32)__msa_ld_w(p1 + k0 + kk, 0); + v4f32 _p2 = (v4f32)__msa_ld_w(p2 + k0 + kk, 0); + v4f32 _p3 = (v4f32)__msa_ld_w(p3 + k0 + kk, 0); + if (input_scale_ptr) + { + const v4f32 _s = (v4f32)__msa_ld_w(input_scale_ptr + k0 + kk, 0); + _p0 = __msa_fmul_w(_p0, _s); + _p1 = __msa_fmul_w(_p1, _s); + _p2 = __msa_fmul_w(_p2, _s); + _p3 = __msa_fmul_w(_p3, _s); + } + _absmax0 = __msa_fmax_w(_absmax0, (v4f32)__msa_and_v((v16u8)_p0, _abs_mask)); + _absmax1 = __msa_fmax_w(_absmax1, (v4f32)__msa_and_v((v16u8)_p1, _abs_mask)); + _absmax2 = __msa_fmax_w(_absmax2, (v4f32)__msa_and_v((v16u8)_p2, _abs_mask)); + _absmax3 = __msa_fmax_w(_absmax3, (v4f32)__msa_and_v((v16u8)_p3, _abs_mask)); + } + absmax0 = __msa_reduce_fmax_w(_absmax0); + absmax1 = __msa_reduce_fmax_w(_absmax1); + absmax2 = __msa_reduce_fmax_w(_absmax2); + absmax3 = __msa_reduce_fmax_w(_absmax3); + + for (; kk < max_kk; kk++) + { + float v0 = p0[k0 + kk]; + float v1 = p1[k0 + kk]; + float v2 = p2[k0 + kk]; + float v3 = p3[k0 + kk]; + if (input_scale_ptr) + { + const float s = input_scale_ptr[k0 + kk]; + v0 *= s; + v1 *= s; + v2 *= s; + v3 *= s; + } + absmax0 = std::max(absmax0, fabsf(v0)); + absmax1 = std::max(absmax1, fabsf(v1)); + absmax2 = std::max(absmax2, fabsf(v2)); + absmax3 = std::max(absmax3, fabsf(v3)); + } + + volatile double scale0_fp64 = absmax0 == 0.f ? 1.0 : 127.0 / (double)absmax0; + volatile double scale1_fp64 = absmax1 == 0.f ? 1.0 : 127.0 / (double)absmax1; + volatile double scale2_fp64 = absmax2 == 0.f ? 1.0 : 127.0 / (double)absmax2; + volatile double scale3_fp64 = absmax3 == 0.f ? 1.0 : 127.0 / (double)absmax3; + const float scale0 = (float)scale0_fp64; + const float scale1 = (float)scale1_fp64; + const float scale2 = (float)scale2_fp64; + const float scale3 = (float)scale3_fp64; + pd[0] = absmax0 / 127.f; + pd[1] = absmax1 / 127.f; + pd[2] = absmax2 / 127.f; + pd[3] = absmax3 / 127.f; + pd += 4; + + kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + float v00 = p0[k0 + kk]; + float v01 = p0[k0 + kk + 1]; + float v02 = p0[k0 + kk + 2]; + float v03 = p0[k0 + kk + 3]; + float v10 = p1[k0 + kk]; + float v11 = p1[k0 + kk + 1]; + float v12 = p1[k0 + kk + 2]; + float v13 = p1[k0 + kk + 3]; + float v20 = p2[k0 + kk]; + float v21 = p2[k0 + kk + 1]; + float v22 = p2[k0 + kk + 2]; + float v23 = p2[k0 + kk + 3]; + float v30 = p3[k0 + kk]; + float v31 = p3[k0 + kk + 1]; + float v32 = p3[k0 + kk + 2]; + float v33 = p3[k0 + kk + 3]; + if (input_scale_ptr) + { + const float s0 = input_scale_ptr[k0 + kk]; + const float s1 = input_scale_ptr[k0 + kk + 1]; + const float s2 = input_scale_ptr[k0 + kk + 2]; + const float s3 = input_scale_ptr[k0 + kk + 3]; + v00 *= s0; + v01 *= s1; + v02 *= s2; + v03 *= s3; + v10 *= s0; + v11 *= s1; + v12 *= s2; + v13 *= s3; + v20 *= s0; + v21 *= s1; + v22 *= s2; + v23 *= s3; + v30 *= s0; + v31 *= s1; + v32 *= s2; + v33 *= s3; + asm volatile("" : "+f"(v00), "+f"(v01), "+f"(v02), "+f"(v03)); + asm volatile("" : "+f"(v10), "+f"(v11), "+f"(v12), "+f"(v13)); + asm volatile("" : "+f"(v20), "+f"(v21), "+f"(v22), "+f"(v23)); + asm volatile("" : "+f"(v30), "+f"(v31), "+f"(v32), "+f"(v33)); + } + pp[0] = float2int8(v00 * scale0); + pp[1] = float2int8(v01 * scale0); + pp[2] = float2int8(v02 * scale0); + pp[3] = float2int8(v03 * scale0); + pp[4] = float2int8(v10 * scale1); + pp[5] = float2int8(v11 * scale1); + pp[6] = float2int8(v12 * scale1); + pp[7] = float2int8(v13 * scale1); + pp[8] = float2int8(v20 * scale2); + pp[9] = float2int8(v21 * scale2); + pp[10] = float2int8(v22 * scale2); + pp[11] = float2int8(v23 * scale2); + pp[12] = float2int8(v30 * scale3); + pp[13] = float2int8(v31 * scale3); + pp[14] = float2int8(v32 * scale3); + pp[15] = float2int8(v33 * scale3); + pp += 16; + } + if (kk + 1 < max_kk) + { + float v00 = p0[k0 + kk]; + float v01 = p0[k0 + kk + 1]; + float v10 = p1[k0 + kk]; + float v11 = p1[k0 + kk + 1]; + float v20 = p2[k0 + kk]; + float v21 = p2[k0 + kk + 1]; + float v30 = p3[k0 + kk]; + float v31 = p3[k0 + kk + 1]; + if (input_scale_ptr) + { + const float s0 = input_scale_ptr[k0 + kk]; + const float s1 = input_scale_ptr[k0 + kk + 1]; + v00 *= s0; + v01 *= s1; + v10 *= s0; + v11 *= s1; + v20 *= s0; + v21 *= s1; + v30 *= s0; + v31 *= s1; + asm volatile("" : "+f"(v00), "+f"(v01), "+f"(v10), "+f"(v11)); + asm volatile("" : "+f"(v20), "+f"(v21), "+f"(v30), "+f"(v31)); + } + pp[0] = float2int8(v00 * scale0); + pp[1] = float2int8(v01 * scale0); + pp[2] = float2int8(v10 * scale1); + pp[3] = float2int8(v11 * scale1); + pp[4] = float2int8(v20 * scale2); + pp[5] = float2int8(v21 * scale2); + pp[6] = float2int8(v30 * scale3); + pp[7] = float2int8(v31 * scale3); + pp += 8; + kk += 2; + } + for (; kk < max_kk; kk++) + { + const int k = k0 + kk; + float v0 = p0[k]; + float v1 = p1[k]; + float v2 = p2[k]; + float v3 = p3[k]; + if (input_scale_ptr) + { + const float s = input_scale_ptr[k]; + v0 *= s; + v1 *= s; + v2 *= s; + v3 *= s; + asm volatile("" : "+f"(v0), "+f"(v1), "+f"(v2), "+f"(v3)); + } + pp[0] = float2int8(v0 * scale0); + pp[1] = float2int8(v1 * scale1); + pp[2] = float2int8(v2 * scale2); + pp[3] = float2int8(v3 * scale3); + pp += 4; + } + } + } +#endif // __mips_msa + for (; ii + 1 < max_ii; ii += 2) + { + const int i0 = i + ii; + const int i1 = i + ii + 1; + const float* p0 = (const float*)A + (size_t)i0 * A_hstep; + const float* p1 = (const float*)A + (size_t)i1 * A_hstep; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + float absmax0 = 0.f; + float absmax1 = 0.f; + for (int kk = 0; kk < max_kk; kk++) + { + const int k = k0 + kk; + float v0 = p0[k]; + float v1 = p1[k]; + if (input_scale_ptr) + { + const float s = input_scale_ptr[k]; + v0 *= s; + v1 *= s; + } + absmax0 = std::max(absmax0, fabsf(v0)); + absmax1 = std::max(absmax1, fabsf(v1)); + } + + volatile double scale0_fp64 = absmax0 == 0.f ? 1.0 : 127.0 / (double)absmax0; + volatile double scale1_fp64 = absmax1 == 0.f ? 1.0 : 127.0 / (double)absmax1; + const float scale0 = (float)scale0_fp64; + const float scale1 = (float)scale1_fp64; + pd[0] = absmax0 / 127.f; + pd[1] = absmax1 / 127.f; + pd += 2; + + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + for (int r = 0; r < 4; r++) + { + const int k = k0 + kk + r; + float v0 = p0[k]; + float v1 = p1[k]; + if (input_scale_ptr) + { + v0 *= input_scale_ptr[k]; + v1 *= input_scale_ptr[k]; + asm volatile("" : "+f"(v0), "+f"(v1)); + } + pp[r] = float2int8(v0 * scale0); + pp[4 + r] = float2int8(v1 * scale1); + } + pp += 8; + } + if (kk + 1 < max_kk) + { + for (int r = 0; r < 2; r++) + { + const int k = k0 + kk + r; + float v0 = p0[k]; + float v1 = p1[k]; + if (input_scale_ptr) + { + v0 *= input_scale_ptr[k]; + v1 *= input_scale_ptr[k]; + asm volatile("" : "+f"(v0), "+f"(v1)); + } + pp[r] = float2int8(v0 * scale0); + pp[2 + r] = float2int8(v1 * scale1); + } + pp += 4; + kk += 2; + } + for (; kk < max_kk; kk++) + { + const int k = k0 + kk; + float v0 = p0[k]; + float v1 = p1[k]; + if (input_scale_ptr) + { + v0 *= input_scale_ptr[k]; + v1 *= input_scale_ptr[k]; + asm volatile("" : "+f"(v0), "+f"(v1)); + } + pp[0] = float2int8(v0 * scale0); + pp[1] = float2int8(v1 * scale1); + pp += 2; + } + } + } + for (; ii < max_ii; ii++) + { + const int i0 = i + ii; + const float* p0 = (const float*)A + (size_t)i0 * A_hstep; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + float absmax0 = 0.f; + for (int kk = 0; kk < max_kk; kk++) + { + const int k = k0 + kk; + float v0 = p0[k]; + if (input_scale_ptr) + v0 *= input_scale_ptr[k]; + absmax0 = std::max(absmax0, fabsf(v0)); + } + + volatile double scale0_fp64 = absmax0 == 0.f ? 1.0 : 127.0 / (double)absmax0; + const float scale0 = (float)scale0_fp64; + *pd++ = absmax0 / 127.f; + + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + for (int r = 0; r < 4; r++) + { + const int k = k0 + kk + r; + float v0 = p0[k]; + if (input_scale_ptr) + { + v0 *= input_scale_ptr[k]; + asm volatile("" : "+f"(v0)); + } + pp[r] = float2int8(v0 * scale0); + } + pp += 4; + } + if (kk + 1 < max_kk) + { + for (int r = 0; r < 2; r++) + { + const int k = k0 + kk + r; + float v0 = p0[k]; + if (input_scale_ptr) + { + v0 *= input_scale_ptr[k]; + asm volatile("" : "+f"(v0)); + } + pp[r] = float2int8(v0 * scale0); + } + pp += 2; + kk += 2; + } + for (; kk < max_kk; kk++) + { + const int k = k0 + kk; + float v0 = p0[k]; + if (input_scale_ptr) + { + v0 *= input_scale_ptr[k]; + asm volatile("" : "+f"(v0)); + } + *pp++ = float2int8(v0 * scale0); + } + } + } +} + +// group-major, row-major within each K4/K2/K1 fragment +static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int block_size, const float* input_scale_ptr) +{ +#if NCNN_RUNTIME_CPU && NCNN_MMI && !__mips_msa && !__mips_loongson_mmi + if (ncnn::cpu_support_loongson_mmi()) + { + transpose_quantize_A_tile_wq_int8_loongson_mmi(A, AT_tile, AT_descales_tile, i, max_ii, block_size, input_scale_ptr); + return; + } +#endif + + signed char* pp = AT_tile; + float* pd = AT_descales_tile; + const int K = AT_tile.w; + const int block_count = AT_descales_tile.w; + const size_t A_hstep = A.dims == 3 ? A.cstep : (size_t)A.w; + + int ii = 0; +#if __mips_msa + for (; ii + 7 < max_ii; ii += 8) + { + const int i0 = i + ii; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + const v16u8 _abs_mask = (v16u8)__msa_fill_w(0x7fffffff); + v4f32 _absmax0 = (v4f32)__msa_fill_w(0); + v4f32 _absmax1 = (v4f32)__msa_fill_w(0); + + int kk = 0; + for (; kk < max_kk; kk++) + { + const int k = k0 + kk; + v4f32 _p0 = (v4f32)__msa_ld_w((const float*)A + (size_t)k * A_hstep + i0, 0); + v4f32 _p1 = (v4f32)__msa_ld_w((const float*)A + (size_t)k * A_hstep + i0 + 4, 0); + if (input_scale_ptr) + { + const v4f32 _s = __msa_fill_w_f32(input_scale_ptr[k]); + _p0 = __msa_fmul_w(_p0, _s); + _p1 = __msa_fmul_w(_p1, _s); + } + _absmax0 = __msa_fmax_w(_absmax0, (v4f32)__msa_and_v((v16u8)_p0, _abs_mask)); + _absmax1 = __msa_fmax_w(_absmax1, (v4f32)__msa_and_v((v16u8)_p1, _abs_mask)); + } + + float absmax[8]; + __msa_st_w((v4i32)_absmax0, absmax, 0); + __msa_st_w((v4i32)_absmax1, absmax + 4, 0); + volatile double scale0_fp64 = absmax[0] == 0.f ? 1.0 : 127.0 / (double)absmax[0]; + volatile double scale1_fp64 = absmax[1] == 0.f ? 1.0 : 127.0 / (double)absmax[1]; + volatile double scale2_fp64 = absmax[2] == 0.f ? 1.0 : 127.0 / (double)absmax[2]; + volatile double scale3_fp64 = absmax[3] == 0.f ? 1.0 : 127.0 / (double)absmax[3]; + volatile double scale4_fp64 = absmax[4] == 0.f ? 1.0 : 127.0 / (double)absmax[4]; + volatile double scale5_fp64 = absmax[5] == 0.f ? 1.0 : 127.0 / (double)absmax[5]; + volatile double scale6_fp64 = absmax[6] == 0.f ? 1.0 : 127.0 / (double)absmax[6]; + volatile double scale7_fp64 = absmax[7] == 0.f ? 1.0 : 127.0 / (double)absmax[7]; + const float scale0 = (float)scale0_fp64; + const float scale1 = (float)scale1_fp64; + const float scale2 = (float)scale2_fp64; + const float scale3 = (float)scale3_fp64; + const float scale4 = (float)scale4_fp64; + const float scale5 = (float)scale5_fp64; + const float scale6 = (float)scale6_fp64; + const float scale7 = (float)scale7_fp64; + pd[0] = absmax[0] / 127.f; + pd[1] = absmax[1] / 127.f; + pd[2] = absmax[2] / 127.f; + pd[3] = absmax[3] / 127.f; + pd[4] = absmax[4] / 127.f; + pd[5] = absmax[5] / 127.f; + pd[6] = absmax[6] / 127.f; + pd[7] = absmax[7] / 127.f; + pd += 8; + + kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + const float* p0 = (const float*)A + (size_t)(k0 + kk) * A_hstep + i0; + const float* p1 = p0 + A_hstep; + const float* p2 = p1 + A_hstep; + const float* p3 = p2 + A_hstep; + v4f32 _p0 = (v4f32)__msa_ld_w(p0, 0); + v4f32 _p1 = (v4f32)__msa_ld_w(p1, 0); + v4f32 _p2 = (v4f32)__msa_ld_w(p2, 0); + v4f32 _p3 = (v4f32)__msa_ld_w(p3, 0); + if (input_scale_ptr) + { + _p0 = __msa_fmul_w(_p0, __msa_fill_w_f32(input_scale_ptr[k0 + kk])); + _p1 = __msa_fmul_w(_p1, __msa_fill_w_f32(input_scale_ptr[k0 + kk + 1])); + _p2 = __msa_fmul_w(_p2, __msa_fill_w_f32(input_scale_ptr[k0 + kk + 2])); + _p3 = __msa_fmul_w(_p3, __msa_fill_w_f32(input_scale_ptr[k0 + kk + 3])); + } + transpose4x4_ps(_p0, _p1, _p2, _p3); + ((int*)pp)[0] = __msa_copy_s_w((v4i32)float2int8(__msa_fmul_w(_p0, __msa_fill_w_f32(scale0))), 0); + ((int*)pp)[1] = __msa_copy_s_w((v4i32)float2int8(__msa_fmul_w(_p1, __msa_fill_w_f32(scale1))), 0); + ((int*)pp)[2] = __msa_copy_s_w((v4i32)float2int8(__msa_fmul_w(_p2, __msa_fill_w_f32(scale2))), 0); + ((int*)pp)[3] = __msa_copy_s_w((v4i32)float2int8(__msa_fmul_w(_p3, __msa_fill_w_f32(scale3))), 0); + + _p0 = (v4f32)__msa_ld_w(p0 + 4, 0); + _p1 = (v4f32)__msa_ld_w(p1 + 4, 0); + _p2 = (v4f32)__msa_ld_w(p2 + 4, 0); + _p3 = (v4f32)__msa_ld_w(p3 + 4, 0); + if (input_scale_ptr) + { + _p0 = __msa_fmul_w(_p0, __msa_fill_w_f32(input_scale_ptr[k0 + kk])); + _p1 = __msa_fmul_w(_p1, __msa_fill_w_f32(input_scale_ptr[k0 + kk + 1])); + _p2 = __msa_fmul_w(_p2, __msa_fill_w_f32(input_scale_ptr[k0 + kk + 2])); + _p3 = __msa_fmul_w(_p3, __msa_fill_w_f32(input_scale_ptr[k0 + kk + 3])); + } + transpose4x4_ps(_p0, _p1, _p2, _p3); + ((int*)pp)[4] = __msa_copy_s_w((v4i32)float2int8(__msa_fmul_w(_p0, __msa_fill_w_f32(scale4))), 0); + ((int*)pp)[5] = __msa_copy_s_w((v4i32)float2int8(__msa_fmul_w(_p1, __msa_fill_w_f32(scale5))), 0); + ((int*)pp)[6] = __msa_copy_s_w((v4i32)float2int8(__msa_fmul_w(_p2, __msa_fill_w_f32(scale6))), 0); + ((int*)pp)[7] = __msa_copy_s_w((v4i32)float2int8(__msa_fmul_w(_p3, __msa_fill_w_f32(scale7))), 0); + pp += 32; + } + if (kk + 1 < max_kk) + { + const float* p0 = (const float*)A + (size_t)(k0 + kk) * A_hstep + i0; + const float* p1 = p0 + A_hstep; + const float s0 = input_scale_ptr ? input_scale_ptr[k0 + kk] : 1.f; + const float s1 = input_scale_ptr ? input_scale_ptr[k0 + kk + 1] : 1.f; + v4f32 _p0 = __msa_fmul_w((v4f32)__msa_ld_w(p0, 0), __msa_fill_w_f32(s0)); + v4f32 _p1 = __msa_fmul_w((v4f32)__msa_ld_w(p1, 0), __msa_fill_w_f32(s1)); + v16i8 _q0 = float2int8(__msa_fmul_w(_p0, (v4f32){scale0, scale1, scale2, scale3})); + v16i8 _q1 = float2int8(__msa_fmul_w(_p1, (v4f32){scale0, scale1, scale2, scale3})); + pp[0] = __msa_copy_s_b(_q0, 0); + pp[1] = __msa_copy_s_b(_q1, 0); + pp[2] = __msa_copy_s_b(_q0, 1); + pp[3] = __msa_copy_s_b(_q1, 1); + pp[4] = __msa_copy_s_b(_q0, 2); + pp[5] = __msa_copy_s_b(_q1, 2); + pp[6] = __msa_copy_s_b(_q0, 3); + pp[7] = __msa_copy_s_b(_q1, 3); + _p0 = __msa_fmul_w((v4f32)__msa_ld_w(p0 + 4, 0), __msa_fill_w_f32(s0)); + _p1 = __msa_fmul_w((v4f32)__msa_ld_w(p1 + 4, 0), __msa_fill_w_f32(s1)); + _q0 = float2int8(__msa_fmul_w(_p0, (v4f32){scale4, scale5, scale6, scale7})); + _q1 = float2int8(__msa_fmul_w(_p1, (v4f32){scale4, scale5, scale6, scale7})); + pp[8] = __msa_copy_s_b(_q0, 0); + pp[9] = __msa_copy_s_b(_q1, 0); + pp[10] = __msa_copy_s_b(_q0, 1); + pp[11] = __msa_copy_s_b(_q1, 1); + pp[12] = __msa_copy_s_b(_q0, 2); + pp[13] = __msa_copy_s_b(_q1, 2); + pp[14] = __msa_copy_s_b(_q0, 3); + pp[15] = __msa_copy_s_b(_q1, 3); + pp += 16; + kk += 2; + } + if (kk < max_kk) + { + const int k = k0 + kk; + const float* p0 = (const float*)A + (size_t)k * A_hstep + i0; + const float s = input_scale_ptr ? input_scale_ptr[k] : 1.f; + v4f32 _p0 = __msa_fmul_w((v4f32)__msa_ld_w(p0, 0), __msa_fill_w_f32(s)); + v4f32 _p1 = __msa_fmul_w((v4f32)__msa_ld_w(p0 + 4, 0), __msa_fill_w_f32(s)); + const v16i8 _q0 = float2int8(__msa_fmul_w(_p0, (v4f32){scale0, scale1, scale2, scale3})); + const v16i8 _q1 = float2int8(__msa_fmul_w(_p1, (v4f32){scale4, scale5, scale6, scale7})); + ((int*)pp)[0] = __msa_copy_s_w((v4i32)_q0, 0); + ((int*)pp)[1] = __msa_copy_s_w((v4i32)_q1, 0); + pp += 8; + } + } + } + for (; ii + 3 < max_ii; ii += 4) + { + const int i0 = i + ii; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + const v16u8 _abs_mask = (v16u8)__msa_fill_w(0x7fffffff); + v4f32 _absmax = (v4f32)__msa_fill_w(0); + + int kk = 0; + for (; kk < max_kk; kk++) + { + const int k = k0 + kk; + v4f32 _p = (v4f32)__msa_ld_w((const float*)A + (size_t)k * A_hstep + i0, 0); + if (input_scale_ptr) + _p = __msa_fmul_w(_p, __msa_fill_w_f32(input_scale_ptr[k])); + _absmax = __msa_fmax_w(_absmax, (v4f32)__msa_and_v((v16u8)_p, _abs_mask)); + } + + float absmax[4]; + __msa_st_w((v4i32)_absmax, absmax, 0); + volatile double scale0_fp64 = absmax[0] == 0.f ? 1.0 : 127.0 / (double)absmax[0]; + volatile double scale1_fp64 = absmax[1] == 0.f ? 1.0 : 127.0 / (double)absmax[1]; + volatile double scale2_fp64 = absmax[2] == 0.f ? 1.0 : 127.0 / (double)absmax[2]; + volatile double scale3_fp64 = absmax[3] == 0.f ? 1.0 : 127.0 / (double)absmax[3]; + const float scale0 = (float)scale0_fp64; + const float scale1 = (float)scale1_fp64; + const float scale2 = (float)scale2_fp64; + const float scale3 = (float)scale3_fp64; + pd[0] = absmax[0] / 127.f; + pd[1] = absmax[1] / 127.f; + pd[2] = absmax[2] / 127.f; + pd[3] = absmax[3] / 127.f; + pd += 4; + + const v4f32 _scale0 = __msa_fill_w_f32(scale0); + const v4f32 _scale1 = __msa_fill_w_f32(scale1); + const v4f32 _scale2 = __msa_fill_w_f32(scale2); + const v4f32 _scale3 = __msa_fill_w_f32(scale3); + kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + const float* p0 = (const float*)A + (size_t)(k0 + kk) * A_hstep + i0; + const float* p1 = p0 + A_hstep; + const float* p2 = p1 + A_hstep; + const float* p3 = p2 + A_hstep; + v4f32 _p0 = (v4f32)__msa_ld_w(p0, 0); + v4f32 _p1 = (v4f32)__msa_ld_w(p1, 0); + v4f32 _p2 = (v4f32)__msa_ld_w(p2, 0); + v4f32 _p3 = (v4f32)__msa_ld_w(p3, 0); + if (input_scale_ptr) + { + _p0 = __msa_fmul_w(_p0, __msa_fill_w_f32(input_scale_ptr[k0 + kk])); + _p1 = __msa_fmul_w(_p1, __msa_fill_w_f32(input_scale_ptr[k0 + kk + 1])); + _p2 = __msa_fmul_w(_p2, __msa_fill_w_f32(input_scale_ptr[k0 + kk + 2])); + _p3 = __msa_fmul_w(_p3, __msa_fill_w_f32(input_scale_ptr[k0 + kk + 3])); + } + transpose4x4_ps(_p0, _p1, _p2, _p3); + _p0 = __msa_fmul_w(_p0, _scale0); + _p1 = __msa_fmul_w(_p1, _scale1); + _p2 = __msa_fmul_w(_p2, _scale2); + _p3 = __msa_fmul_w(_p3, _scale3); + ((int*)pp)[0] = __msa_copy_s_w((v4i32)float2int8(_p0), 0); + ((int*)pp)[1] = __msa_copy_s_w((v4i32)float2int8(_p1), 0); + ((int*)pp)[2] = __msa_copy_s_w((v4i32)float2int8(_p2), 0); + ((int*)pp)[3] = __msa_copy_s_w((v4i32)float2int8(_p3), 0); + pp += 16; + } + if (kk + 1 < max_kk) + { + const float* p0 = (const float*)A + (size_t)(k0 + kk) * A_hstep + i0; + const float* p1 = p0 + A_hstep; + float v00 = p0[0]; + float v10 = p0[1]; + float v20 = p0[2]; + float v30 = p0[3]; + float v01 = p1[0]; + float v11 = p1[1]; + float v21 = p1[2]; + float v31 = p1[3]; + if (input_scale_ptr) + { + const float s0 = input_scale_ptr[k0 + kk]; + const float s1 = input_scale_ptr[k0 + kk + 1]; + v00 *= s0; + v10 *= s0; + v20 *= s0; + v30 *= s0; + v01 *= s1; + v11 *= s1; + v21 *= s1; + v31 *= s1; + asm volatile("" : "+f"(v00), "+f"(v01), "+f"(v10), "+f"(v11)); + asm volatile("" : "+f"(v20), "+f"(v21), "+f"(v30), "+f"(v31)); + } + pp[0] = float2int8(v00 * scale0); + pp[1] = float2int8(v01 * scale0); + pp[2] = float2int8(v10 * scale1); + pp[3] = float2int8(v11 * scale1); + pp[4] = float2int8(v20 * scale2); + pp[5] = float2int8(v21 * scale2); + pp[6] = float2int8(v30 * scale3); + pp[7] = float2int8(v31 * scale3); + pp += 8; + kk += 2; + } + for (; kk < max_kk; kk++) + { + const int k = k0 + kk; + const float* p0 = (const float*)A + (size_t)k * A_hstep + i0; + float v0 = p0[0]; + float v1 = p0[1]; + float v2 = p0[2]; + float v3 = p0[3]; + if (input_scale_ptr) + { + const float s = input_scale_ptr[k]; + v0 *= s; + v1 *= s; + v2 *= s; + v3 *= s; + asm volatile("" : "+f"(v0), "+f"(v1), "+f"(v2), "+f"(v3)); + } + pp[0] = float2int8(v0 * scale0); + pp[1] = float2int8(v1 * scale1); + pp[2] = float2int8(v2 * scale2); + pp[3] = float2int8(v3 * scale3); + pp += 4; + } + } + } +#endif // __mips_msa + for (; ii + 1 < max_ii; ii += 2) + { + const int i0 = i + ii; + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + float absmax0 = 0.f; + float absmax1 = 0.f; + for (int kk = 0; kk < max_kk; kk++) + { + const int k = k0 + kk; + const float* p0 = (const float*)A + (size_t)k * A_hstep + i0; + float v0 = p0[0]; + float v1 = p0[1]; + if (input_scale_ptr) + { + const float s = input_scale_ptr[k]; + v0 *= s; + v1 *= s; + } + absmax0 = std::max(absmax0, fabsf(v0)); + absmax1 = std::max(absmax1, fabsf(v1)); + } + + volatile double scale0_fp64 = absmax0 == 0.f ? 1.0 : 127.0 / (double)absmax0; + volatile double scale1_fp64 = absmax1 == 0.f ? 1.0 : 127.0 / (double)absmax1; + const float scale0 = (float)scale0_fp64; + const float scale1 = (float)scale1_fp64; + pd[0] = absmax0 / 127.f; + pd[1] = absmax1 / 127.f; + pd += 2; + + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + for (int r = 0; r < 4; r++) + { + const int k = k0 + kk + r; + const float* p0 = (const float*)A + (size_t)k * A_hstep + i0; + float v0 = p0[0]; + float v1 = p0[1]; + if (input_scale_ptr) + { + const float s = input_scale_ptr[k]; + v0 *= s; + v1 *= s; + asm volatile("" : "+f"(v0), "+f"(v1)); + } + pp[r] = float2int8(v0 * scale0); + pp[4 + r] = float2int8(v1 * scale1); + } + pp += 8; + } + if (kk + 1 < max_kk) + { + for (int r = 0; r < 2; r++) + { + const int k = k0 + kk + r; + const float* p0 = (const float*)A + (size_t)k * A_hstep + i0; + float v0 = p0[0]; + float v1 = p0[1]; + if (input_scale_ptr) + { + const float s = input_scale_ptr[k]; + v0 *= s; + v1 *= s; + asm volatile("" : "+f"(v0), "+f"(v1)); + } + pp[r] = float2int8(v0 * scale0); + pp[2 + r] = float2int8(v1 * scale1); + } + pp += 4; + kk += 2; + } + for (; kk < max_kk; kk++) + { + const int k = k0 + kk; + const float* p0 = (const float*)A + (size_t)k * A_hstep + i0; + float v0 = p0[0]; + float v1 = p0[1]; + if (input_scale_ptr) + { + const float s = input_scale_ptr[k]; + v0 *= s; + v1 *= s; + asm volatile("" : "+f"(v0), "+f"(v1)); + } + pp[0] = float2int8(v0 * scale0); + pp[1] = float2int8(v1 * scale1); + pp += 2; + } + } + } + for (; ii < max_ii; ii++) + { + const int i0 = i + ii; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + float absmax0 = 0.f; + for (int kk = 0; kk < max_kk; kk++) + { + const int k = k0 + kk; + float v0 = ((const float*)A)[(size_t)k * A_hstep + i0]; + if (input_scale_ptr) + v0 *= input_scale_ptr[k]; + absmax0 = std::max(absmax0, fabsf(v0)); + } + + volatile double scale0_fp64 = absmax0 == 0.f ? 1.0 : 127.0 / (double)absmax0; + const float scale0 = (float)scale0_fp64; + *pd++ = absmax0 / 127.f; + + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + for (int r = 0; r < 4; r++) + { + const int k = k0 + kk + r; + float v0 = ((const float*)A)[(size_t)k * A_hstep + i0]; + if (input_scale_ptr) + { + v0 *= input_scale_ptr[k]; + asm volatile("" : "+f"(v0)); + } + pp[r] = float2int8(v0 * scale0); + } + pp += 4; + } + if (kk + 1 < max_kk) + { + for (int r = 0; r < 2; r++) + { + const int k = k0 + kk + r; + float v0 = ((const float*)A)[(size_t)k * A_hstep + i0]; + if (input_scale_ptr) + { + v0 *= input_scale_ptr[k]; + asm volatile("" : "+f"(v0)); + } + pp[r] = float2int8(v0 * scale0); + } + pp += 2; + kk += 2; + } + for (; kk < max_kk; kk++) + { + const int k = k0 + kk; + float v0 = ((const float*)A)[(size_t)k * A_hstep + i0]; + if (input_scale_ptr) + { + v0 *= input_scale_ptr[k]; + asm volatile("" : "+f"(v0)); + } + *pp++ = float2int8(v0 * scale0); + } + } + } +} + +static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int block_size) +{ +#if NCNN_RUNTIME_CPU && NCNN_MMI && !__mips_msa && !__mips_loongson_mmi + if (ncnn::cpu_support_loongson_mmi()) + { + gemm_transB_packed_tile_wq_int8_loongson_mmi(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, block_size); + return; + } +#endif + + const signed char* pAT = AT_tile; + const int A_hstep = AT_tile.w; + const float* pAT_descales = AT_descales_tile; + const int A_descales_hstep = AT_descales_tile.w; + const signed char* pBT = BT_tile; + const float* pBT_descales = BT_descales_tile; + float* outptr = topT_tile; + const int K = AT_tile.w; + const int num_blocks = (K + block_size - 1) / block_size; + + int ii = 0; +#if __mips_msa + for (; ii + 7 < max_ii; ii += 8) + { + const signed char* pB = pBT; + const float* pBD = pBT_descales; + const v8i16 _one = __msa_fill_h(1); + + int jj = 0; + for (; jj + 3 < max_jj; jj += 4) + { + const signed char* pA = pAT; + const float* pAD = pAT_descales; + v4f32 _fsum0 = (v4f32)__msa_fill_w(0); + v4f32 _fsum1 = (v4f32)__msa_fill_w(0); + v4f32 _fsum2 = (v4f32)__msa_fill_w(0); + v4f32 _fsum3 = (v4f32)__msa_fill_w(0); + v4f32 _fsum4 = (v4f32)__msa_fill_w(0); + v4f32 _fsum5 = (v4f32)__msa_fill_w(0); + v4f32 _fsum6 = (v4f32)__msa_fill_w(0); + v4f32 _fsum7 = (v4f32)__msa_fill_w(0); + + for (int k = 0; k < K; k += block_size) + { + const int max_kk = std::min(K - k, block_size); + v4i32 _sum0 = __msa_fill_w(0); + v4i32 _sum1 = __msa_fill_w(0); + v4i32 _sum2 = __msa_fill_w(0); + v4i32 _sum3 = __msa_fill_w(0); + v4i32 _sum4 = __msa_fill_w(0); + v4i32 _sum5 = __msa_fill_w(0); + v4i32 _sum6 = __msa_fill_w(0); + v4i32 _sum7 = __msa_fill_w(0); + + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __builtin_prefetch(pA + 64); + __builtin_prefetch(pB + 64); + const v16i8 _pA0 = __msa_ld_b(pA, 0); + const v16i8 _pA0r = (v16i8)__msa_shf_w((v4i32)_pA0, _MSA_SHUFFLE(1, 0, 3, 2)); + const v16i8 _pB = __msa_ld_b(pB, 0); + const v16i8 _pBr = (v16i8)__msa_shf_w((v4i32)_pB, _MSA_SHUFFLE(0, 3, 2, 1)); + _sum0 = __msa_dpadd_s_w(_sum0, __msa_dotp_s_h(_pA0, _pB), _one); + _sum1 = __msa_dpadd_s_w(_sum1, __msa_dotp_s_h(_pA0, _pBr), _one); + _sum2 = __msa_dpadd_s_w(_sum2, __msa_dotp_s_h(_pA0r, _pB), _one); + _sum3 = __msa_dpadd_s_w(_sum3, __msa_dotp_s_h(_pA0r, _pBr), _one); + + const v16i8 _pA1 = __msa_ld_b(pA + 16, 0); + const v16i8 _pA1r = (v16i8)__msa_shf_w((v4i32)_pA1, _MSA_SHUFFLE(1, 0, 3, 2)); + _sum4 = __msa_dpadd_s_w(_sum4, __msa_dotp_s_h(_pA1, _pB), _one); + _sum5 = __msa_dpadd_s_w(_sum5, __msa_dotp_s_h(_pA1, _pBr), _one); + _sum6 = __msa_dpadd_s_w(_sum6, __msa_dotp_s_h(_pA1r, _pB), _one); + _sum7 = __msa_dpadd_s_w(_sum7, __msa_dotp_s_h(_pA1r, _pBr), _one); + pA += 32; + pB += 16; + } + + _sum2 = __msa_shf_w(_sum2, _MSA_SHUFFLE(1, 0, 3, 2)); + _sum3 = __msa_shf_w(_sum3, _MSA_SHUFFLE(1, 0, 3, 2)); + transpose4x4_epi32(_sum0, _sum1, _sum2, _sum3); + _sum1 = __msa_shf_w(_sum1, _MSA_SHUFFLE(2, 1, 0, 3)); + _sum2 = __msa_shf_w(_sum2, _MSA_SHUFFLE(1, 0, 3, 2)); + _sum3 = __msa_shf_w(_sum3, _MSA_SHUFFLE(0, 3, 2, 1)); + _sum6 = __msa_shf_w(_sum6, _MSA_SHUFFLE(1, 0, 3, 2)); + _sum7 = __msa_shf_w(_sum7, _MSA_SHUFFLE(1, 0, 3, 2)); + transpose4x4_epi32(_sum4, _sum5, _sum6, _sum7); + _sum5 = __msa_shf_w(_sum5, _MSA_SHUFFLE(2, 1, 0, 3)); + _sum6 = __msa_shf_w(_sum6, _MSA_SHUFFLE(1, 0, 3, 2)); + _sum7 = __msa_shf_w(_sum7, _MSA_SHUFFLE(0, 3, 2, 1)); + + if (kk + 1 < max_kk) + { + const v16i8 _pB = (v16i8)__msa_fill_d_ptr(pB); + const v8i16 _pA = (v8i16)__msa_ld_b(pA, 0); + v8i16 _s = __msa_dotp_s_h((v16i8)__msa_splati_h(_pA, 0), _pB); + _sum0 = __msa_addv_w(_sum0, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + _s = __msa_dotp_s_h((v16i8)__msa_splati_h(_pA, 1), _pB); + _sum1 = __msa_addv_w(_sum1, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + _s = __msa_dotp_s_h((v16i8)__msa_splati_h(_pA, 2), _pB); + _sum2 = __msa_addv_w(_sum2, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + _s = __msa_dotp_s_h((v16i8)__msa_splati_h(_pA, 3), _pB); + _sum3 = __msa_addv_w(_sum3, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + _s = __msa_dotp_s_h((v16i8)__msa_splati_h(_pA, 4), _pB); + _sum4 = __msa_addv_w(_sum4, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + _s = __msa_dotp_s_h((v16i8)__msa_splati_h(_pA, 5), _pB); + _sum5 = __msa_addv_w(_sum5, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + _s = __msa_dotp_s_h((v16i8)__msa_splati_h(_pA, 6), _pB); + _sum6 = __msa_addv_w(_sum6, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + _s = __msa_dotp_s_h((v16i8)__msa_splati_h(_pA, 7), _pB); + _sum7 = __msa_addv_w(_sum7, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + pA += 16; + pB += 8; + kk += 2; + } + if (kk < max_kk) + { + const v16i8 _pA8 = (v16i8)__msa_fill_d_ptr(pA); + const v8i16 _pA = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pA8, 0), _pA8); + v8i16 _pB = (v8i16)__msa_fill_w(*(const int*)pB); + _pB = (v8i16)__msa_ilvr_b(__msa_clti_s_b((v16i8)_pB, 0), (v16i8)_pB); + v8i16 _s = __msa_mulv_h(__msa_splati_h(_pA, 0), _pB); + _sum0 = __msa_addv_w(_sum0, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + _s = __msa_mulv_h(__msa_splati_h(_pA, 1), _pB); + _sum1 = __msa_addv_w(_sum1, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + _s = __msa_mulv_h(__msa_splati_h(_pA, 2), _pB); + _sum2 = __msa_addv_w(_sum2, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + _s = __msa_mulv_h(__msa_splati_h(_pA, 3), _pB); + _sum3 = __msa_addv_w(_sum3, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + _s = __msa_mulv_h(__msa_splati_h(_pA, 4), _pB); + _sum4 = __msa_addv_w(_sum4, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + _s = __msa_mulv_h(__msa_splati_h(_pA, 5), _pB); + _sum5 = __msa_addv_w(_sum5, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + _s = __msa_mulv_h(__msa_splati_h(_pA, 6), _pB); + _sum6 = __msa_addv_w(_sum6, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + _s = __msa_mulv_h(__msa_splati_h(_pA, 7), _pB); + _sum7 = __msa_addv_w(_sum7, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + pA += 8; + pB += 4; + } + + const v4f32 _descaleB = (v4f32)__msa_ld_w(pBD, 0); + const v4f32 _descaleA0 = (v4f32)__msa_ld_w(pAD, 0); + const v4f32 _descaleA1 = (v4f32)__msa_ld_w(pAD + 4, 0); + _fsum0 = __msa_fadd_w(_fsum0, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum0), __msa_fmul_w(_descaleB, (v4f32)__msa_splati_w((v4i32)_descaleA0, 0)))); + _fsum1 = __msa_fadd_w(_fsum1, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum1), __msa_fmul_w(_descaleB, (v4f32)__msa_splati_w((v4i32)_descaleA0, 1)))); + _fsum2 = __msa_fadd_w(_fsum2, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum2), __msa_fmul_w(_descaleB, (v4f32)__msa_splati_w((v4i32)_descaleA0, 2)))); + _fsum3 = __msa_fadd_w(_fsum3, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum3), __msa_fmul_w(_descaleB, (v4f32)__msa_splati_w((v4i32)_descaleA0, 3)))); + _fsum4 = __msa_fadd_w(_fsum4, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum4), __msa_fmul_w(_descaleB, (v4f32)__msa_splati_w((v4i32)_descaleA1, 0)))); + _fsum5 = __msa_fadd_w(_fsum5, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum5), __msa_fmul_w(_descaleB, (v4f32)__msa_splati_w((v4i32)_descaleA1, 1)))); + _fsum6 = __msa_fadd_w(_fsum6, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum6), __msa_fmul_w(_descaleB, (v4f32)__msa_splati_w((v4i32)_descaleA1, 2)))); + _fsum7 = __msa_fadd_w(_fsum7, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum7), __msa_fmul_w(_descaleB, (v4f32)__msa_splati_w((v4i32)_descaleA1, 3)))); + pAD += 8; + pBD += 4; + } + + transpose4x4_ps(_fsum0, _fsum1, _fsum2, _fsum3); + transpose4x4_ps(_fsum4, _fsum5, _fsum6, _fsum7); + __msa_st_w((v4i32)_fsum0, outptr, 0); + __msa_st_w((v4i32)_fsum4, outptr + 4, 0); + __msa_st_w((v4i32)_fsum1, outptr + 8, 0); + __msa_st_w((v4i32)_fsum5, outptr + 12, 0); + __msa_st_w((v4i32)_fsum2, outptr + 16, 0); + __msa_st_w((v4i32)_fsum6, outptr + 20, 0); + __msa_st_w((v4i32)_fsum3, outptr + 24, 0); + __msa_st_w((v4i32)_fsum7, outptr + 28, 0); + outptr += 32; + } + for (; jj + 1 < max_jj; jj += 2) + { + const signed char* pA = pAT; + const float* pAD = pAT_descales; + v4f32 _fsum0 = (v4f32)__msa_fill_w(0); + v4f32 _fsum1 = (v4f32)__msa_fill_w(0); + v4f32 _fsum2 = (v4f32)__msa_fill_w(0); + v4f32 _fsum3 = (v4f32)__msa_fill_w(0); + + for (int k = 0; k < K; k += block_size) + { + const int max_kk = std::min(K - k, block_size); + v4i32 _sum0 = __msa_fill_w(0); + v4i32 _sum1 = __msa_fill_w(0); + v4i32 _sum2 = __msa_fill_w(0); + v4i32 _sum3 = __msa_fill_w(0); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + const v16i8 _pA0 = __msa_ld_b(pA, 0); + const v16i8 _pA1 = __msa_ld_b(pA + 16, 0); + const v16i8 _pB = (v16i8)__msa_fill_d_ptr(pB); + const v16i8 _pBr = (v16i8)__msa_shf_w((v4i32)_pB, _MSA_SHUFFLE(0, 3, 2, 1)); + _sum0 = __msa_dpadd_s_w(_sum0, __msa_dotp_s_h(_pA0, _pB), _one); + _sum1 = __msa_dpadd_s_w(_sum1, __msa_dotp_s_h(_pA0, _pBr), _one); + _sum2 = __msa_dpadd_s_w(_sum2, __msa_dotp_s_h(_pA1, _pB), _one); + _sum3 = __msa_dpadd_s_w(_sum3, __msa_dotp_s_h(_pA1, _pBr), _one); + pA += 32; + pB += 8; + } + + const v4i32 _sum0e = __msa_shf_w(_sum0, _MSA_SHUFFLE(3, 1, 2, 0)); + const v4i32 _sum0o = __msa_shf_w(_sum0, _MSA_SHUFFLE(2, 0, 3, 1)); + const v4i32 _sum1e = __msa_shf_w(_sum1, _MSA_SHUFFLE(3, 1, 2, 0)); + const v4i32 _sum1o = __msa_shf_w(_sum1, _MSA_SHUFFLE(2, 0, 3, 1)); + v4i32 _sum0x = (v4i32)__msa_ilvr_w(_sum1o, _sum0e); + v4i32 _sum1x = (v4i32)__msa_ilvr_w(_sum0o, _sum1e); + const v4i32 _sum2e = __msa_shf_w(_sum2, _MSA_SHUFFLE(3, 1, 2, 0)); + const v4i32 _sum2o = __msa_shf_w(_sum2, _MSA_SHUFFLE(2, 0, 3, 1)); + const v4i32 _sum3e = __msa_shf_w(_sum3, _MSA_SHUFFLE(3, 1, 2, 0)); + const v4i32 _sum3o = __msa_shf_w(_sum3, _MSA_SHUFFLE(2, 0, 3, 1)); + v4i32 _sum2x = (v4i32)__msa_ilvr_w(_sum3o, _sum2e); + v4i32 _sum3x = (v4i32)__msa_ilvr_w(_sum2o, _sum3e); + + if (kk + 1 < max_kk) + { + const v16i8 _pA = __msa_ld_b(pA, 0); + const v8i16 _pB = (v8i16)__msa_fill_w(*(const int*)pB); + const v8i16 _s0 = __msa_dotp_s_h(_pA, (v16i8)__msa_splati_h(_pB, 0)); + const v8i16 _s1 = __msa_dotp_s_h(_pA, (v16i8)__msa_splati_h(_pB, 1)); + const v8i16 _sign0 = __msa_clti_s_h(_s0, 0); + const v8i16 _sign1 = __msa_clti_s_h(_s1, 0); + _sum0x = __msa_addv_w(_sum0x, (v4i32)__msa_ilvr_h(_sign0, _s0)); + _sum1x = __msa_addv_w(_sum1x, (v4i32)__msa_ilvr_h(_sign1, _s1)); + _sum2x = __msa_addv_w(_sum2x, (v4i32)__msa_ilvl_h(_sign0, _s0)); + _sum3x = __msa_addv_w(_sum3x, (v4i32)__msa_ilvl_h(_sign1, _s1)); + pA += 16; + pB += 4; + kk += 2; + } + if (kk < max_kk) + { + const v16i8 _pA8 = (v16i8)__msa_fill_d_ptr(pA); + const v8i16 _pA = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pA8, 0), _pA8); + const v16i8 _pB8 = (v16i8)__msa_fill_h(*(const short*)pB); + const v8i16 _pB0 = (v8i16)__msa_ilvr_b(__msa_clti_s_b(__msa_splati_b(_pB8, 0), 0), __msa_splati_b(_pB8, 0)); + const v8i16 _pB1 = (v8i16)__msa_ilvr_b(__msa_clti_s_b(__msa_splati_b(_pB8, 1), 0), __msa_splati_b(_pB8, 1)); + const v8i16 _s0 = __msa_mulv_h(_pA, _pB0); + const v8i16 _s1 = __msa_mulv_h(_pA, _pB1); + const v8i16 _sign0 = __msa_clti_s_h(_s0, 0); + const v8i16 _sign1 = __msa_clti_s_h(_s1, 0); + _sum0x = __msa_addv_w(_sum0x, (v4i32)__msa_ilvr_h(_sign0, _s0)); + _sum1x = __msa_addv_w(_sum1x, (v4i32)__msa_ilvr_h(_sign1, _s1)); + _sum2x = __msa_addv_w(_sum2x, (v4i32)__msa_ilvl_h(_sign0, _s0)); + _sum3x = __msa_addv_w(_sum3x, (v4i32)__msa_ilvl_h(_sign1, _s1)); + pA += 8; + pB += 2; + } + + const v4f32 _descaleA0 = (v4f32)__msa_ld_w(pAD, 0); + const v4f32 _descaleA1 = (v4f32)__msa_ld_w(pAD + 4, 0); + _fsum0 = __msa_fadd_w(_fsum0, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum0x), __msa_fmul_w(_descaleA0, __msa_fill_w_f32(pBD[0])))); + _fsum1 = __msa_fadd_w(_fsum1, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum2x), __msa_fmul_w(_descaleA1, __msa_fill_w_f32(pBD[0])))); + _fsum2 = __msa_fadd_w(_fsum2, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum1x), __msa_fmul_w(_descaleA0, __msa_fill_w_f32(pBD[1])))); + _fsum3 = __msa_fadd_w(_fsum3, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum3x), __msa_fmul_w(_descaleA1, __msa_fill_w_f32(pBD[1])))); + pAD += 8; + pBD += 2; + } + + __msa_st_w((v4i32)_fsum0, outptr, 0); + __msa_st_w((v4i32)_fsum1, outptr + 4, 0); + __msa_st_w((v4i32)_fsum2, outptr + 8, 0); + __msa_st_w((v4i32)_fsum3, outptr + 12, 0); + outptr += 16; + } + for (; jj < max_jj; jj++) + { + const signed char* pA = pAT; + const float* pAD = pAT_descales; + v4f32 _fsum0 = (v4f32)__msa_fill_w(0); + v4f32 _fsum1 = (v4f32)__msa_fill_w(0); + + for (int k = 0; k < K; k += block_size) + { + const int max_kk = std::min(K - k, block_size); + v4i32 _sum0 = __msa_fill_w(0); + v4i32 _sum1 = __msa_fill_w(0); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + const v16i8 _pA0 = __msa_ld_b(pA, 0); + const v16i8 _pA1 = __msa_ld_b(pA + 16, 0); + const v16i8 _pB = (v16i8)__msa_fill_w(*(const int*)pB); + _sum0 = __msa_dpadd_s_w(_sum0, __msa_dotp_s_h(_pA0, _pB), _one); + _sum1 = __msa_dpadd_s_w(_sum1, __msa_dotp_s_h(_pA1, _pB), _one); + pA += 32; + pB += 4; + } + if (kk + 1 < max_kk) + { + const v16i8 _pA = __msa_ld_b(pA, 0); + const v16i8 _pB = (v16i8)__msa_fill_h(*(const short*)pB); + const v8i16 _s = __msa_dotp_s_h(_pA, _pB); + const v8i16 _sign = __msa_clti_s_h(_s, 0); + _sum0 = __msa_addv_w(_sum0, (v4i32)__msa_ilvr_h(_sign, _s)); + _sum1 = __msa_addv_w(_sum1, (v4i32)__msa_ilvl_h(_sign, _s)); + pA += 16; + pB += 2; + kk += 2; + } + if (kk < max_kk) + { + const v16i8 _pA8 = (v16i8)__msa_fill_d_ptr(pA); + const v8i16 _pA = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pA8, 0), _pA8); + const v8i16 _s = __msa_mulv_h(_pA, __msa_fill_h(pB[0])); + const v8i16 _sign = __msa_clti_s_h(_s, 0); + _sum0 = __msa_addv_w(_sum0, (v4i32)__msa_ilvr_h(_sign, _s)); + _sum1 = __msa_addv_w(_sum1, (v4i32)__msa_ilvl_h(_sign, _s)); + pA += 8; + pB++; + } + + const v4f32 _descaleA0 = (v4f32)__msa_ld_w(pAD, 0); + const v4f32 _descaleA1 = (v4f32)__msa_ld_w(pAD + 4, 0); + _fsum0 = __msa_fadd_w(_fsum0, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum0), __msa_fmul_w(_descaleA0, __msa_fill_w_f32(pBD[0])))); + _fsum1 = __msa_fadd_w(_fsum1, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum1), __msa_fmul_w(_descaleA1, __msa_fill_w_f32(pBD[0])))); + pAD += 8; + pBD++; + } + + __msa_st_w((v4i32)_fsum0, outptr, 0); + __msa_st_w((v4i32)_fsum1, outptr + 4, 0); + outptr += 8; + } + + pAT += (size_t)8 * A_hstep; + pAT_descales += (size_t)8 * A_descales_hstep; + } + for (; ii + 3 < max_ii; ii += 4) + { + const signed char* pB = pBT; + const float* pBD = pBT_descales; + const v8i16 _one = __msa_fill_h(1); + + int jj = 0; + for (; jj + 7 < max_jj; jj += 8) + { + const signed char* pA = pAT; + const float* pAD = pAT_descales; + const signed char* pB0 = pB; + const signed char* pB1 = pB + (size_t)4 * K; + const float* pBD0 = pBD; + const float* pBD1 = pBD + (size_t)4 * num_blocks; + v4f32 _fsum0 = (v4f32)__msa_fill_w(0); + v4f32 _fsum1 = (v4f32)__msa_fill_w(0); + v4f32 _fsum2 = (v4f32)__msa_fill_w(0); + v4f32 _fsum3 = (v4f32)__msa_fill_w(0); + v4f32 _fsum4 = (v4f32)__msa_fill_w(0); + v4f32 _fsum5 = (v4f32)__msa_fill_w(0); + v4f32 _fsum6 = (v4f32)__msa_fill_w(0); + v4f32 _fsum7 = (v4f32)__msa_fill_w(0); + + for (int k = 0; k < K; k += block_size) + { + const int max_kk = std::min(K - k, block_size); + v4i32 _sum0 = __msa_fill_w(0); + v4i32 _sum1 = __msa_fill_w(0); + v4i32 _sum2 = __msa_fill_w(0); + v4i32 _sum3 = __msa_fill_w(0); + v4i32 _sum4 = __msa_fill_w(0); + v4i32 _sum5 = __msa_fill_w(0); + v4i32 _sum6 = __msa_fill_w(0); + v4i32 _sum7 = __msa_fill_w(0); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + const v16i8 _pA = __msa_ld_b(pA, 0); + const v16i8 _pAr = (v16i8)__msa_shf_w((v4i32)_pA, _MSA_SHUFFLE(1, 0, 3, 2)); + const v16i8 _pB0 = __msa_ld_b(pB0, 0); + const v16i8 _pB1 = __msa_ld_b(pB1, 0); + const v16i8 _pB0r = (v16i8)__msa_shf_w((v4i32)_pB0, _MSA_SHUFFLE(0, 3, 2, 1)); + const v16i8 _pB1r = (v16i8)__msa_shf_w((v4i32)_pB1, _MSA_SHUFFLE(0, 3, 2, 1)); + _sum0 = __msa_dpadd_s_w(_sum0, __msa_dotp_s_h(_pA, _pB0), _one); + _sum1 = __msa_dpadd_s_w(_sum1, __msa_dotp_s_h(_pA, _pB0r), _one); + _sum2 = __msa_dpadd_s_w(_sum2, __msa_dotp_s_h(_pAr, _pB0), _one); + _sum3 = __msa_dpadd_s_w(_sum3, __msa_dotp_s_h(_pAr, _pB0r), _one); + _sum4 = __msa_dpadd_s_w(_sum4, __msa_dotp_s_h(_pA, _pB1), _one); + _sum5 = __msa_dpadd_s_w(_sum5, __msa_dotp_s_h(_pA, _pB1r), _one); + _sum6 = __msa_dpadd_s_w(_sum6, __msa_dotp_s_h(_pAr, _pB1), _one); + _sum7 = __msa_dpadd_s_w(_sum7, __msa_dotp_s_h(_pAr, _pB1r), _one); + pA += 16; + pB0 += 16; + pB1 += 16; + } + + const signed char* pA2 = pA; + const signed char* pB02 = pB0; + const signed char* pB12 = pB1; + const bool has_k2 = kk + 1 < max_kk; + if (kk + 1 < max_kk) + { + pA += 8; + pB0 += 8; + pB1 += 8; + kk += 2; + } + for (; kk < max_kk; kk++) + { + v8i16 _pA = (v8i16)__msa_fill_w(*(const int*)pA); + _pA = (v8i16)__msa_ilvr_b(__msa_clti_s_b((v16i8)_pA, 0), (v16i8)_pA); + const v8i16 _pAr = __msa_shf_h(_pA, _MSA_SHUFFLE(1, 0, 3, 2)); + v8i16 _pB0 = (v8i16)__msa_fill_w(*(const int*)pB0); + v8i16 _pB1 = (v8i16)__msa_fill_w(*(const int*)pB1); + _pB0 = (v8i16)__msa_ilvr_b(__msa_clti_s_b((v16i8)_pB0, 0), (v16i8)_pB0); + _pB1 = (v8i16)__msa_ilvr_b(__msa_clti_s_b((v16i8)_pB1, 0), (v16i8)_pB1); + const v8i16 _pB0r = __msa_shf_h(_pB0, _MSA_SHUFFLE(0, 3, 2, 1)); + const v8i16 _pB1r = __msa_shf_h(_pB1, _MSA_SHUFFLE(0, 3, 2, 1)); + const v8i16 _s0 = __msa_mulv_h(_pA, _pB0); + const v8i16 _s1 = __msa_mulv_h(_pA, _pB0r); + const v8i16 _s2 = __msa_mulv_h(_pAr, _pB0); + const v8i16 _s3 = __msa_mulv_h(_pAr, _pB0r); + const v8i16 _s4 = __msa_mulv_h(_pA, _pB1); + const v8i16 _s5 = __msa_mulv_h(_pA, _pB1r); + const v8i16 _s6 = __msa_mulv_h(_pAr, _pB1); + const v8i16 _s7 = __msa_mulv_h(_pAr, _pB1r); + _sum0 = __msa_addv_w(_sum0, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s0, 0), _s0)); + _sum1 = __msa_addv_w(_sum1, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s1, 0), _s1)); + _sum2 = __msa_addv_w(_sum2, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s2, 0), _s2)); + _sum3 = __msa_addv_w(_sum3, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s3, 0), _s3)); + _sum4 = __msa_addv_w(_sum4, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s4, 0), _s4)); + _sum5 = __msa_addv_w(_sum5, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s5, 0), _s5)); + _sum6 = __msa_addv_w(_sum6, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s6, 0), _s6)); + _sum7 = __msa_addv_w(_sum7, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s7, 0), _s7)); + pA += 4; + pB0 += 4; + pB1 += 4; + } + + _sum2 = __msa_shf_w(_sum2, _MSA_SHUFFLE(1, 0, 3, 2)); + _sum3 = __msa_shf_w(_sum3, _MSA_SHUFFLE(1, 0, 3, 2)); + transpose4x4_epi32(_sum0, _sum1, _sum2, _sum3); + _sum1 = __msa_shf_w(_sum1, _MSA_SHUFFLE(2, 1, 0, 3)); + _sum2 = __msa_shf_w(_sum2, _MSA_SHUFFLE(1, 0, 3, 2)); + _sum3 = __msa_shf_w(_sum3, _MSA_SHUFFLE(0, 3, 2, 1)); + _sum6 = __msa_shf_w(_sum6, _MSA_SHUFFLE(1, 0, 3, 2)); + _sum7 = __msa_shf_w(_sum7, _MSA_SHUFFLE(1, 0, 3, 2)); + transpose4x4_epi32(_sum4, _sum5, _sum6, _sum7); + _sum5 = __msa_shf_w(_sum5, _MSA_SHUFFLE(2, 1, 0, 3)); + _sum6 = __msa_shf_w(_sum6, _MSA_SHUFFLE(1, 0, 3, 2)); + _sum7 = __msa_shf_w(_sum7, _MSA_SHUFFLE(0, 3, 2, 1)); + if (has_k2) + { + const v8i16 _pA = (v8i16)__msa_fill_d_ptr(pA2); + const v16i8 _pB0 = (v16i8)__msa_fill_d_ptr(pB02); + const v16i8 _pB1 = (v16i8)__msa_fill_d_ptr(pB12); + v8i16 _s = __msa_dotp_s_h((v16i8)__msa_splati_h(_pA, 0), _pB0); + _sum0 = __msa_addv_w(_sum0, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + _s = __msa_dotp_s_h((v16i8)__msa_splati_h(_pA, 1), _pB0); + _sum1 = __msa_addv_w(_sum1, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + _s = __msa_dotp_s_h((v16i8)__msa_splati_h(_pA, 2), _pB0); + _sum2 = __msa_addv_w(_sum2, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + _s = __msa_dotp_s_h((v16i8)__msa_splati_h(_pA, 3), _pB0); + _sum3 = __msa_addv_w(_sum3, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + _s = __msa_dotp_s_h((v16i8)__msa_splati_h(_pA, 0), _pB1); + _sum4 = __msa_addv_w(_sum4, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + _s = __msa_dotp_s_h((v16i8)__msa_splati_h(_pA, 1), _pB1); + _sum5 = __msa_addv_w(_sum5, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + _s = __msa_dotp_s_h((v16i8)__msa_splati_h(_pA, 2), _pB1); + _sum6 = __msa_addv_w(_sum6, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + _s = __msa_dotp_s_h((v16i8)__msa_splati_h(_pA, 3), _pB1); + _sum7 = __msa_addv_w(_sum7, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + } + + const v4f32 _descaleB0 = (v4f32)__msa_ld_w(pBD0, 0); + const v4f32 _descaleB1 = (v4f32)__msa_ld_w(pBD1, 0); + _fsum0 = __msa_fadd_w(_fsum0, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum0), __msa_fmul_w(_descaleB0, __msa_fill_w_f32(pAD[0])))); + _fsum1 = __msa_fadd_w(_fsum1, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum1), __msa_fmul_w(_descaleB0, __msa_fill_w_f32(pAD[1])))); + _fsum2 = __msa_fadd_w(_fsum2, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum2), __msa_fmul_w(_descaleB0, __msa_fill_w_f32(pAD[2])))); + _fsum3 = __msa_fadd_w(_fsum3, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum3), __msa_fmul_w(_descaleB0, __msa_fill_w_f32(pAD[3])))); + _fsum4 = __msa_fadd_w(_fsum4, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum4), __msa_fmul_w(_descaleB1, __msa_fill_w_f32(pAD[0])))); + _fsum5 = __msa_fadd_w(_fsum5, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum5), __msa_fmul_w(_descaleB1, __msa_fill_w_f32(pAD[1])))); + _fsum6 = __msa_fadd_w(_fsum6, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum6), __msa_fmul_w(_descaleB1, __msa_fill_w_f32(pAD[2])))); + _fsum7 = __msa_fadd_w(_fsum7, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum7), __msa_fmul_w(_descaleB1, __msa_fill_w_f32(pAD[3])))); + pAD += 4; + pBD0 += 4; + pBD1 += 4; + } + + transpose4x4_ps(_fsum0, _fsum1, _fsum2, _fsum3); + transpose4x4_ps(_fsum4, _fsum5, _fsum6, _fsum7); + __msa_st_w((v4i32)_fsum0, outptr, 0); + __msa_st_w((v4i32)_fsum1, outptr + 4, 0); + __msa_st_w((v4i32)_fsum2, outptr + 8, 0); + __msa_st_w((v4i32)_fsum3, outptr + 12, 0); + __msa_st_w((v4i32)_fsum4, outptr + 16, 0); + __msa_st_w((v4i32)_fsum5, outptr + 20, 0); + __msa_st_w((v4i32)_fsum6, outptr + 24, 0); + __msa_st_w((v4i32)_fsum7, outptr + 28, 0); + outptr += 32; + pB = pB1; + pBD = pBD1; + } + for (; jj + 3 < max_jj; jj += 4) + { + const signed char* pA = pAT; + const float* pAD = pAT_descales; + v4f32 _fsum0 = (v4f32)__msa_fill_w(0); + v4f32 _fsum1 = (v4f32)__msa_fill_w(0); + v4f32 _fsum2 = (v4f32)__msa_fill_w(0); + v4f32 _fsum3 = (v4f32)__msa_fill_w(0); + + for (int k = 0; k < K; k += block_size) + { + const int max_kk = std::min(K - k, block_size); + v4i32 _sum0 = __msa_fill_w(0); + v4i32 _sum1 = __msa_fill_w(0); + v4i32 _sum2 = __msa_fill_w(0); + v4i32 _sum3 = __msa_fill_w(0); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + const v16i8 _pA = __msa_ld_b(pA, 0); + const v16i8 _pAr = (v16i8)__msa_shf_w((v4i32)_pA, _MSA_SHUFFLE(1, 0, 3, 2)); + const v16i8 _pB0 = __msa_ld_b(pB, 0); + const v16i8 _pB0r = (v16i8)__msa_shf_w((v4i32)_pB0, _MSA_SHUFFLE(0, 3, 2, 1)); + _sum0 = __msa_dpadd_s_w(_sum0, __msa_dotp_s_h(_pA, _pB0), _one); + _sum1 = __msa_dpadd_s_w(_sum1, __msa_dotp_s_h(_pA, _pB0r), _one); + _sum2 = __msa_dpadd_s_w(_sum2, __msa_dotp_s_h(_pAr, _pB0), _one); + _sum3 = __msa_dpadd_s_w(_sum3, __msa_dotp_s_h(_pAr, _pB0r), _one); + pA += 16; + pB += 16; + } + v8i16 _sum2_0 = __msa_fill_h(0); + v8i16 _sum2_1 = __msa_fill_h(0); + v8i16 _sum2_2 = __msa_fill_h(0); + v8i16 _sum2_3 = __msa_fill_h(0); + if (kk + 1 < max_kk) + { + const v16i8 _pB = (v16i8)__msa_fill_d_ptr(pB); + const v8i16 _pA = (v8i16)__msa_fill_d_ptr(pA); + const v16i8 _pA0 = (v16i8)__msa_splati_h(_pA, 0); + const v16i8 _pA1 = (v16i8)__msa_splati_h(_pA, 1); + const v16i8 _pA2 = (v16i8)__msa_splati_h(_pA, 2); + const v16i8 _pA3 = (v16i8)__msa_splati_h(_pA, 3); + _sum2_0 = __msa_dotp_s_h(_pA0, _pB); + _sum2_1 = __msa_dotp_s_h(_pA1, _pB); + _sum2_2 = __msa_dotp_s_h(_pA2, _pB); + _sum2_3 = __msa_dotp_s_h(_pA3, _pB); + pA += 8; + pB += 8; + kk += 2; + } + for (; kk < max_kk; kk++) + { + v8i16 _pA = (v8i16)__msa_fill_w(*(const int*)pA); + _pA = (v8i16)__msa_ilvr_b(__msa_clti_s_b((v16i8)_pA, 0), (v16i8)_pA); + const v8i16 _pAr = __msa_shf_h(_pA, _MSA_SHUFFLE(1, 0, 3, 2)); + v8i16 _pB0 = (v8i16)__msa_fill_w(*(const int*)pB); + _pB0 = (v8i16)__msa_ilvr_b(__msa_clti_s_b((v16i8)_pB0, 0), (v16i8)_pB0); + const v8i16 _pB0r = __msa_shf_h(_pB0, _MSA_SHUFFLE(0, 3, 2, 1)); + const v8i16 _s0 = __msa_mulv_h(_pA, _pB0); + const v8i16 _s1 = __msa_mulv_h(_pA, _pB0r); + const v8i16 _s2 = __msa_mulv_h(_pAr, _pB0); + const v8i16 _s3 = __msa_mulv_h(_pAr, _pB0r); + _sum0 = __msa_addv_w(_sum0, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s0, 0), _s0)); + _sum1 = __msa_addv_w(_sum1, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s1, 0), _s1)); + _sum2 = __msa_addv_w(_sum2, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s2, 0), _s2)); + _sum3 = __msa_addv_w(_sum3, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s3, 0), _s3)); + pA += 4; + pB += 4; + } + _sum2 = __msa_shf_w(_sum2, _MSA_SHUFFLE(1, 0, 3, 2)); + _sum3 = __msa_shf_w(_sum3, _MSA_SHUFFLE(1, 0, 3, 2)); + transpose4x4_epi32(_sum0, _sum1, _sum2, _sum3); + _sum1 = __msa_shf_w(_sum1, _MSA_SHUFFLE(2, 1, 0, 3)); + _sum2 = __msa_shf_w(_sum2, _MSA_SHUFFLE(1, 0, 3, 2)); + _sum3 = __msa_shf_w(_sum3, _MSA_SHUFFLE(0, 3, 2, 1)); + _sum0 = __msa_addv_w(_sum0, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_sum2_0, 0), _sum2_0)); + _sum1 = __msa_addv_w(_sum1, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_sum2_1, 0), _sum2_1)); + _sum2 = __msa_addv_w(_sum2, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_sum2_2, 0), _sum2_2)); + _sum3 = __msa_addv_w(_sum3, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_sum2_3, 0), _sum2_3)); + const v4f32 _descaleB = (v4f32)__msa_ld_w(pBD, 0); + _fsum0 = __msa_fadd_w(_fsum0, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum0), __msa_fmul_w(_descaleB, __msa_fill_w_f32(pAD[0])))); + _fsum1 = __msa_fadd_w(_fsum1, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum1), __msa_fmul_w(_descaleB, __msa_fill_w_f32(pAD[1])))); + _fsum2 = __msa_fadd_w(_fsum2, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum2), __msa_fmul_w(_descaleB, __msa_fill_w_f32(pAD[2])))); + _fsum3 = __msa_fadd_w(_fsum3, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum3), __msa_fmul_w(_descaleB, __msa_fill_w_f32(pAD[3])))); + pAD += 4; + pBD += 4; + } + transpose4x4_ps(_fsum0, _fsum1, _fsum2, _fsum3); + __msa_st_w((v4i32)_fsum0, outptr, 0); + __msa_st_w((v4i32)_fsum1, outptr + 4, 0); + __msa_st_w((v4i32)_fsum2, outptr + 8, 0); + __msa_st_w((v4i32)_fsum3, outptr + 12, 0); + outptr += 16; + } + for (; jj + 1 < max_jj; jj += 2) + { + const signed char* pA = pAT; + const float* pAD = pAT_descales; + v4f32 _fsum0 = (v4f32)__msa_fill_w(0); + v4f32 _fsum1 = (v4f32)__msa_fill_w(0); + for (int k = 0; k < K; k += block_size) + { + const int max_kk = std::min(K - k, block_size); + v4i32 _sum0 = __msa_fill_w(0); + v4i32 _sum1 = __msa_fill_w(0); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + const v16i8 _pA = __msa_ld_b(pA, 0); + const v16i8 _pB0 = (v16i8)__msa_fill_d_ptr(pB); + const v16i8 _pB0r = (v16i8)__msa_shf_w((v4i32)_pB0, _MSA_SHUFFLE(0, 3, 2, 1)); + _sum0 = __msa_dpadd_s_w(_sum0, __msa_dotp_s_h(_pA, _pB0), _one); + _sum1 = __msa_dpadd_s_w(_sum1, __msa_dotp_s_h(_pA, _pB0r), _one); + pA += 16; + pB += 8; + } + v8i16 _sum2_0 = __msa_fill_h(0); + v8i16 _sum2_1 = __msa_fill_h(0); + if (kk + 1 < max_kk) + { + const v16i8 _pA = (v16i8)__msa_fill_d_ptr(pA); + const v8i16 _pB = (v8i16)__msa_fill_w(*(const int*)pB); + const v16i8 _pB0 = (v16i8)__msa_splati_h(_pB, 0); + const v16i8 _pB1 = (v16i8)__msa_splati_h(_pB, 1); + _sum2_0 = __msa_dotp_s_h(_pA, _pB0); + _sum2_1 = __msa_dotp_s_h(_pA, _pB1); + pA += 8; + pB += 4; + kk += 2; + } + for (; kk < max_kk; kk++) + { + v8i16 _pA = (v8i16)__msa_fill_w(*(const int*)pA); + _pA = (v8i16)__msa_ilvr_b(__msa_clti_s_b((v16i8)_pA, 0), (v16i8)_pA); + const v16i8 _pB8 = (v16i8)__msa_fill_h(*(const short*)pB); + const v8i16 _pB0 = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pB8, 0), _pB8); + const v8i16 _pB0r = __msa_shf_h(_pB0, _MSA_SHUFFLE(0, 3, 2, 1)); + const v8i16 _s0 = __msa_mulv_h(_pA, _pB0); + const v8i16 _s1 = __msa_mulv_h(_pA, _pB0r); + _sum0 = __msa_addv_w(_sum0, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s0, 0), _s0)); + _sum1 = __msa_addv_w(_sum1, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s1, 0), _s1)); + pA += 4; + pB += 2; + } + const v4i32 _sum0e = __msa_shf_w(_sum0, _MSA_SHUFFLE(3, 1, 2, 0)); + const v4i32 _sum0o = __msa_shf_w(_sum0, _MSA_SHUFFLE(2, 0, 3, 1)); + const v4i32 _sum1e = __msa_shf_w(_sum1, _MSA_SHUFFLE(3, 1, 2, 0)); + const v4i32 _sum1o = __msa_shf_w(_sum1, _MSA_SHUFFLE(2, 0, 3, 1)); + v4i32 _sum0x = (v4i32)__msa_ilvr_w(_sum1o, _sum0e); + v4i32 _sum1x = (v4i32)__msa_ilvr_w(_sum0o, _sum1e); + _sum0x = __msa_addv_w(_sum0x, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_sum2_0, 0), _sum2_0)); + _sum1x = __msa_addv_w(_sum1x, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_sum2_1, 0), _sum2_1)); + const v4f32 _descaleA = (v4f32)__msa_ld_w(pAD, 0); + _fsum0 = __msa_fadd_w(_fsum0, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum0x), __msa_fmul_w(_descaleA, __msa_fill_w_f32(pBD[0])))); + _fsum1 = __msa_fadd_w(_fsum1, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum1x), __msa_fmul_w(_descaleA, __msa_fill_w_f32(pBD[1])))); + pAD += 4; + pBD += 2; + } + __msa_st_w((v4i32)_fsum0, outptr, 0); + __msa_st_w((v4i32)_fsum1, outptr + 4, 0); + outptr += 8; + } + for (; jj < max_jj; jj++) + { + const signed char* pA = pAT; + const float* pAD = pAT_descales; + v4f32 _fsum0 = (v4f32)__msa_fill_w(0); + for (int k = 0; k < K; k += block_size) + { + const int max_kk = std::min(K - k, block_size); + v4i32 _sum0 = __msa_fill_w(0); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + const v16i8 _pA = __msa_ld_b(pA, 0); + const v16i8 _pB0 = (v16i8)__msa_fill_w(*(const int*)pB); + _sum0 = __msa_dpadd_s_w(_sum0, __msa_dotp_s_h(_pA, _pB0), _one); + pA += 16; + pB += 4; + } + v8i16 _sum2_0 = __msa_fill_h(0); + if (kk + 1 < max_kk) + { + const v16i8 _pA = (v16i8)__msa_fill_d_ptr(pA); + const v16i8 _pB = (v16i8)__msa_fill_h(*(const short*)pB); + _sum2_0 = __msa_dotp_s_h(_pA, _pB); + pA += 8; + pB += 2; + kk += 2; + } + for (; kk < max_kk; kk++) + { + v8i16 _pA = (v8i16)__msa_fill_w(*(const int*)pA); + _pA = (v8i16)__msa_ilvr_b(__msa_clti_s_b((v16i8)_pA, 0), (v16i8)_pA); + const v8i16 _pB0 = __msa_fill_h(pB[0]); + const v8i16 _s0 = __msa_mulv_h(_pA, _pB0); + _sum0 = __msa_addv_w(_sum0, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s0, 0), _s0)); + pA += 4; + pB++; + } + _sum0 = __msa_addv_w(_sum0, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_sum2_0, 0), _sum2_0)); + const v4f32 _descaleA = (v4f32)__msa_ld_w(pAD, 0); + _fsum0 = __msa_fadd_w(_fsum0, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum0), __msa_fmul_w(_descaleA, __msa_fill_w_f32(pBD[0])))); + pAD += 4; + pBD++; + } + __msa_st_w((v4i32)_fsum0, outptr, 0); + outptr += 4; + } + + pAT += (size_t)4 * A_hstep; + pAT_descales += (size_t)4 * A_descales_hstep; + } +#endif // __mips_msa + for (; ii + 1 < max_ii; ii += 2) + { + const signed char* pB = pBT; + const float* pBD = pBT_descales; + + int jj = 0; +#if __mips_msa + const v8i16 _one = __msa_fill_h(1); + for (; jj + 7 < max_jj; jj += 8) + { + const signed char* pA = pAT; + const float* pAD = pAT_descales; + const signed char* pB0 = pB; + const signed char* pB1 = pB + (size_t)4 * K; + const float* pBD0 = pBD; + const float* pBD1 = pBD + (size_t)4 * num_blocks; + v4f32 _fsum0 = (v4f32)__msa_fill_w(0); + v4f32 _fsum1 = (v4f32)__msa_fill_w(0); + v4f32 _fsum2 = (v4f32)__msa_fill_w(0); + v4f32 _fsum3 = (v4f32)__msa_fill_w(0); + for (int k = 0; k < K; k += block_size) + { + const int max_kk = std::min(K - k, block_size); + v4i32 _sum0 = __msa_fill_w(0); + v4i32 _sum1 = __msa_fill_w(0); + v4i32 _sum2 = __msa_fill_w(0); + v4i32 _sum3 = __msa_fill_w(0); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + const v16i8 _pA = (v16i8)__msa_fill_d_ptr(pA); + const v16i8 _pB0 = __msa_ld_b(pB0, 0); + const v16i8 _pB00 = (v16i8)__msa_ilvr_w((v4i32)_pB0, (v4i32)_pB0); + const v16i8 _pB01 = (v16i8)__msa_ilvl_w((v4i32)_pB0, (v4i32)_pB0); + const v16i8 _pB1 = __msa_ld_b(pB1, 0); + const v16i8 _pB10 = (v16i8)__msa_ilvr_w((v4i32)_pB1, (v4i32)_pB1); + const v16i8 _pB11 = (v16i8)__msa_ilvl_w((v4i32)_pB1, (v4i32)_pB1); + _sum0 = __msa_dpadd_s_w(_sum0, __msa_dotp_s_h(_pA, _pB00), _one); + _sum1 = __msa_dpadd_s_w(_sum1, __msa_dotp_s_h(_pA, _pB01), _one); + _sum2 = __msa_dpadd_s_w(_sum2, __msa_dotp_s_h(_pA, _pB10), _one); + _sum3 = __msa_dpadd_s_w(_sum3, __msa_dotp_s_h(_pA, _pB11), _one); + pA += 8; + pB0 += 16; + pB1 += 16; + } + if (kk + 1 < max_kk) + { + const v8i16 _pA = (v8i16)__msa_fill_w(*(const int*)pA); + const v16i8 _pA0 = (v16i8)__msa_splati_h(_pA, 0); + const v16i8 _pA1 = (v16i8)__msa_splati_h(_pA, 1); + const v16i8 _pB0 = (v16i8)__msa_fill_d_ptr(pB0); + const v16i8 _pB1 = (v16i8)__msa_fill_d_ptr(pB1); + const v8i16 _s00 = __msa_dotp_s_h(_pA0, _pB0); + const v8i16 _s01 = __msa_dotp_s_h(_pA1, _pB0); + const v8i16 _s10 = __msa_dotp_s_h(_pA0, _pB1); + const v8i16 _s11 = __msa_dotp_s_h(_pA1, _pB1); + const v8i16 _s0 = (v8i16)__msa_ilvr_h(_s01, _s00); + const v8i16 _s2 = (v8i16)__msa_ilvr_h(_s11, _s10); + const v8i16 _sign0 = __msa_clti_s_h(_s0, 0); + const v8i16 _sign2 = __msa_clti_s_h(_s2, 0); + _sum0 = __msa_addv_w(_sum0, (v4i32)__msa_ilvr_h(_sign0, _s0)); + _sum1 = __msa_addv_w(_sum1, (v4i32)__msa_ilvl_h(_sign0, _s0)); + _sum2 = __msa_addv_w(_sum2, (v4i32)__msa_ilvr_h(_sign2, _s2)); + _sum3 = __msa_addv_w(_sum3, (v4i32)__msa_ilvl_h(_sign2, _s2)); + pA += 4; + pB0 += 8; + pB1 += 8; + kk += 2; + } + if (kk < max_kk) + { + const v16i8 _pA8 = (v16i8)__msa_fill_h(*(const short*)pA); + const v16i8 _pA0b = __msa_splati_b(_pA8, 0); + const v16i8 _pA1b = __msa_splati_b(_pA8, 1); + const v8i16 _pA0 = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pA0b, 0), _pA0b); + const v8i16 _pA1 = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pA1b, 0), _pA1b); + const v16i8 _pB08 = (v16i8)__msa_fill_w(*(const int*)pB0); + const v16i8 _pB18 = (v16i8)__msa_fill_w(*(const int*)pB1); + const v8i16 _pB0 = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pB08, 0), _pB08); + const v8i16 _pB1 = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pB18, 0), _pB18); + const v8i16 _s00 = __msa_mulv_h(_pA0, _pB0); + const v8i16 _s01 = __msa_mulv_h(_pA1, _pB0); + const v8i16 _s10 = __msa_mulv_h(_pA0, _pB1); + const v8i16 _s11 = __msa_mulv_h(_pA1, _pB1); + const v8i16 _s0 = (v8i16)__msa_ilvr_h(_s01, _s00); + const v8i16 _s2 = (v8i16)__msa_ilvr_h(_s11, _s10); + const v8i16 _sign0 = __msa_clti_s_h(_s0, 0); + const v8i16 _sign2 = __msa_clti_s_h(_s2, 0); + _sum0 = __msa_addv_w(_sum0, (v4i32)__msa_ilvr_h(_sign0, _s0)); + _sum1 = __msa_addv_w(_sum1, (v4i32)__msa_ilvl_h(_sign0, _s0)); + _sum2 = __msa_addv_w(_sum2, (v4i32)__msa_ilvr_h(_sign2, _s2)); + _sum3 = __msa_addv_w(_sum3, (v4i32)__msa_ilvl_h(_sign2, _s2)); + pA += 2; + pB0 += 4; + pB1 += 4; + } + const v4f32 _descaleA = (v4f32){pAD[0], pAD[1], pAD[0], pAD[1]}; + const v4f32 _descaleB0 = (v4f32){pBD0[0], pBD0[0], pBD0[1], pBD0[1]}; + const v4f32 _descaleB1 = (v4f32){pBD0[2], pBD0[2], pBD0[3], pBD0[3]}; + const v4f32 _descaleB2 = (v4f32){pBD1[0], pBD1[0], pBD1[1], pBD1[1]}; + const v4f32 _descaleB3 = (v4f32){pBD1[2], pBD1[2], pBD1[3], pBD1[3]}; + _fsum0 = __msa_fadd_w(_fsum0, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum0), __msa_fmul_w(_descaleA, _descaleB0))); + _fsum1 = __msa_fadd_w(_fsum1, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum1), __msa_fmul_w(_descaleA, _descaleB1))); + _fsum2 = __msa_fadd_w(_fsum2, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum2), __msa_fmul_w(_descaleA, _descaleB2))); + _fsum3 = __msa_fadd_w(_fsum3, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum3), __msa_fmul_w(_descaleA, _descaleB3))); + pAD += 2; + pBD0 += 4; + pBD1 += 4; + } + __msa_st_w((v4i32)_fsum0, outptr, 0); + __msa_st_w((v4i32)_fsum1, outptr + 4, 0); + __msa_st_w((v4i32)_fsum2, outptr + 8, 0); + __msa_st_w((v4i32)_fsum3, outptr + 12, 0); + outptr += 16; + pB = pB1; + pBD = pBD1; + } + for (; jj + 3 < max_jj; jj += 4) + { + const signed char* pA = pAT; + const float* pAD = pAT_descales; + v4f32 _fsum0 = (v4f32)__msa_fill_w(0); + v4f32 _fsum1 = (v4f32)__msa_fill_w(0); + for (int k = 0; k < K; k += block_size) + { + const int max_kk = std::min(K - k, block_size); + v4i32 _sum0 = __msa_fill_w(0); + v4i32 _sum1 = __msa_fill_w(0); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + const v16i8 _pA = (v16i8)__msa_fill_d_ptr(pA); + const v16i8 _pB0 = __msa_ld_b(pB, 0); + const v16i8 _pB01 = (v16i8)__msa_ilvr_w((v4i32)_pB0, (v4i32)_pB0); + const v16i8 _pB23 = (v16i8)__msa_ilvl_w((v4i32)_pB0, (v4i32)_pB0); + _sum0 = __msa_dpadd_s_w(_sum0, __msa_dotp_s_h(_pA, _pB01), _one); + _sum1 = __msa_dpadd_s_w(_sum1, __msa_dotp_s_h(_pA, _pB23), _one); + pA += 8; + pB += 16; + } + if (kk + 1 < max_kk) + { + const v8i16 _pA = (v8i16)__msa_fill_w(*(const int*)pA); + const v16i8 _pA0 = (v16i8)__msa_splati_h(_pA, 0); + const v16i8 _pA1 = (v16i8)__msa_splati_h(_pA, 1); + const v16i8 _pB = (v16i8)__msa_fill_d_ptr(pB); + const v8i16 _s00 = __msa_dotp_s_h(_pA0, _pB); + const v8i16 _s01 = __msa_dotp_s_h(_pA1, _pB); + const v8i16 _s0 = (v8i16)__msa_ilvr_h(_s01, _s00); + const v8i16 _sign = __msa_clti_s_h(_s0, 0); + _sum0 = __msa_addv_w(_sum0, (v4i32)__msa_ilvr_h(_sign, _s0)); + _sum1 = __msa_addv_w(_sum1, (v4i32)__msa_ilvl_h(_sign, _s0)); + pA += 4; + pB += 8; + kk += 2; + } + if (kk < max_kk) + { + const v16i8 _pA8 = (v16i8)__msa_fill_h(*(const short*)pA); + const v16i8 _pA0b = __msa_splati_b(_pA8, 0); + const v16i8 _pA1b = __msa_splati_b(_pA8, 1); + const v8i16 _pA0 = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pA0b, 0), _pA0b); + const v8i16 _pA1 = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pA1b, 0), _pA1b); + const v16i8 _pB8 = (v16i8)__msa_fill_w(*(const int*)pB); + const v8i16 _pB = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pB8, 0), _pB8); + const v8i16 _s00 = __msa_mulv_h(_pA0, _pB); + const v8i16 _s01 = __msa_mulv_h(_pA1, _pB); + const v8i16 _s0 = (v8i16)__msa_ilvr_h(_s01, _s00); + const v8i16 _sign = __msa_clti_s_h(_s0, 0); + _sum0 = __msa_addv_w(_sum0, (v4i32)__msa_ilvr_h(_sign, _s0)); + _sum1 = __msa_addv_w(_sum1, (v4i32)__msa_ilvl_h(_sign, _s0)); + pA += 2; + pB += 4; + } + const v4f32 _descaleA = (v4f32){pAD[0], pAD[1], pAD[0], pAD[1]}; + const v4f32 _descaleB0 = (v4f32){pBD[0], pBD[0], pBD[1], pBD[1]}; + const v4f32 _descaleB1 = (v4f32){pBD[2], pBD[2], pBD[3], pBD[3]}; + _fsum0 = __msa_fadd_w(_fsum0, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum0), __msa_fmul_w(_descaleA, _descaleB0))); + _fsum1 = __msa_fadd_w(_fsum1, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum1), __msa_fmul_w(_descaleA, _descaleB1))); + pAD += 2; + pBD += 4; + } + __msa_st_w((v4i32)_fsum0, outptr, 0); + __msa_st_w((v4i32)_fsum1, outptr + 4, 0); + outptr += 8; + } +#endif // __mips_msa + for (; jj + 1 < max_jj; jj += 2) + { + const signed char* pA = pAT; + const float* pAD = pAT_descales; + float sum00 = 0.f; + float sum01 = 0.f; + float sum10 = 0.f; + float sum11 = 0.f; + + for (int k = 0; k < K; k += block_size) + { + const int max_kk = std::min(K - k, block_size); + int sum00_i = 0; + int sum01_i = 0; + int sum10_i = 0; + int sum11_i = 0; + int kk = 0; +#if __mips_loongson_mmi + int32x2_t _sum00 = __mmi_pzerow_s(); + int32x2_t _sum01 = __mmi_pzerow_s(); + int32x2_t _sum10 = __mmi_pzerow_s(); + int32x2_t _sum11 = __mmi_pzerow_s(); +#if NCNN_GNU_INLINE_ASM + double _tmp0; + double _tmp1; + double _tmp2; + double _tmp3; + double _tmp4; + double _tmp5; + double _tmp6; + double _tmp7; + double _shift; + const int shift_8 = 8; + for (; kk + 3 < max_kk; kk += 4) + { + asm volatile( + "ld $0, 32(%1) \n" + "ldc1 %6, 0(%0) \n" + "ldc1 %8, 0(%1) \n" +#if __mips64 + "dmtc1 $0, %7 \n" +#else + "mtc1 $0, %7 \n" +#endif + "mtc1 %21, %14 \n" +#if __mips64 + "daddiu %0, %0, 8 \n" + "daddiu %1, %1, 8 \n" +#else + "addiu %0, %0, 8 \n" + "addiu %1, %1, 8 \n" +#endif + "punpcklbh %10, %6, %7 \n" + "punpckhbh %11, %6, %7 \n" + "punpcklbh %12, %8, %7 \n" + "punpckhbh %13, %8, %7 \n" + "psllh %10, %10, %14 \n" + "psllh %11, %11, %14 \n" + "psllh %12, %12, %14 \n" + "psllh %13, %13, %14 \n" + "psrah %10, %10, %14 \n" + "psrah %11, %11, %14 \n" + "psrah %12, %12, %14 \n" + "psrah %13, %13, %14 \n" + "pmaddhw %6, %10, %12 \n" + "pmaddhw %7, %11, %12 \n" + "pmaddhw %8, %10, %13 \n" + "pmaddhw %9, %11, %13 \n" + "paddw %2, %2, %6 \n" + "paddw %3, %3, %7 \n" + "paddw %4, %4, %8 \n" + "paddw %5, %5, %9 \n" + : "=r"(pA), + "=r"(pB), + "=f"(_sum00), + "=f"(_sum01), + "=f"(_sum10), + "=f"(_sum11), + "=&f"(_tmp0), + "=&f"(_tmp1), + "=&f"(_tmp2), + "=&f"(_tmp3), + "=&f"(_tmp4), + "=&f"(_tmp5), + "=&f"(_tmp6), + "=&f"(_tmp7), + "=&f"(_shift) + : "0"(pA), + "1"(pB), + "2"(_sum00), + "3"(_sum01), + "4"(_sum10), + "5"(_sum11), + "r"(shift_8) + : "memory"); + } +#else // NCNN_GNU_INLINE_ASM + const int8x8_t _zero = __mmi_pzerob_s(); + for (; kk + 3 < max_kk; kk += 4) + { + __builtin_prefetch(pB + 32); + const int8x8_t _pA = __mmi_pldb_s(pA); + const int8x8_t _pB = __mmi_pldb_s(pB); + int16x4_t _pA0 = (int16x4_t)__mmi_punpcklbh_s(_pA, _zero); + int16x4_t _pA1 = (int16x4_t)__mmi_punpckhbh_s(_pA, _zero); + int16x4_t _pB0 = (int16x4_t)__mmi_punpcklbh_s(_pB, _zero); + int16x4_t _pB1 = (int16x4_t)__mmi_punpckhbh_s(_pB, _zero); + _pA0 = __mmi_psrah_s(__mmi_psllh_s(_pA0, 8), 8); + _pA1 = __mmi_psrah_s(__mmi_psllh_s(_pA1, 8), 8); + _pB0 = __mmi_psrah_s(__mmi_psllh_s(_pB0, 8), 8); + _pB1 = __mmi_psrah_s(__mmi_psllh_s(_pB1, 8), 8); + _sum00 = __mmi_paddw_s(_sum00, __mmi_pmaddhw(_pA0, _pB0)); + _sum01 = __mmi_paddw_s(_sum01, __mmi_pmaddhw(_pA1, _pB0)); + _sum10 = __mmi_paddw_s(_sum10, __mmi_pmaddhw(_pA0, _pB1)); + _sum11 = __mmi_paddw_s(_sum11, __mmi_pmaddhw(_pA1, _pB1)); + pA += 8; + pB += 8; + } +#endif // NCNN_GNU_INLINE_ASM + _sum00 = __mmi_paddw_s(_sum00, __mmi_punpckhwd_s(_sum00, _sum00)); + _sum01 = __mmi_paddw_s(_sum01, __mmi_punpckhwd_s(_sum01, _sum01)); + _sum10 = __mmi_paddw_s(_sum10, __mmi_punpckhwd_s(_sum10, _sum10)); + _sum11 = __mmi_paddw_s(_sum11, __mmi_punpckhwd_s(_sum11, _sum11)); + sum00_i += _sum00[0]; + sum01_i += _sum01[0]; + sum10_i += _sum10[0]; + sum11_i += _sum11[0]; +#endif // __mips_loongson_mmi + for (; kk + 3 < max_kk; kk += 4) + { + sum00_i += pA[0] * pB[0] + pA[1] * pB[1] + pA[2] * pB[2] + pA[3] * pB[3]; + sum01_i += pA[4] * pB[0] + pA[5] * pB[1] + pA[6] * pB[2] + pA[7] * pB[3]; + sum10_i += pA[0] * pB[4] + pA[1] * pB[5] + pA[2] * pB[6] + pA[3] * pB[7]; + sum11_i += pA[4] * pB[4] + pA[5] * pB[5] + pA[6] * pB[6] + pA[7] * pB[7]; + pA += 8; + pB += 8; + } + if (kk + 1 < max_kk) + { + sum00_i += pA[0] * pB[0] + pA[1] * pB[1]; + sum01_i += pA[2] * pB[0] + pA[3] * pB[1]; + sum10_i += pA[0] * pB[2] + pA[1] * pB[3]; + sum11_i += pA[2] * pB[2] + pA[3] * pB[3]; + pA += 4; + pB += 4; + kk += 2; + } + for (; kk < max_kk; kk++) + { + sum00_i += pA[0] * pB[0]; + sum01_i += pA[1] * pB[0]; + sum10_i += pA[0] * pB[1]; + sum11_i += pA[1] * pB[1]; + pA += 2; + pB += 2; + } + sum00 += sum00_i * pAD[0] * pBD[0]; + sum01 += sum01_i * pAD[1] * pBD[0]; + sum10 += sum10_i * pAD[0] * pBD[1]; + sum11 += sum11_i * pAD[1] * pBD[1]; + pAD += 2; + pBD += 2; + } + + outptr[0] = sum00; + outptr[1] = sum01; + outptr[2] = sum10; + outptr[3] = sum11; + outptr += 4; + } + for (; jj < max_jj; jj++) + { + const signed char* pA = pAT; + const float* pAD = pAT_descales; + float sum0 = 0.f; + float sum1 = 0.f; + + for (int k = 0; k < K; k += block_size) + { + const int max_kk = std::min(K - k, block_size); + int sum0_i = 0; + int sum1_i = 0; + int kk = 0; +#if __mips_loongson_mmi + int32x2_t _sum0 = __mmi_pzerow_s(); + int32x2_t _sum1 = __mmi_pzerow_s(); +#if NCNN_GNU_INLINE_ASM + double _tmp0; + double _tmp1; + double _tmp2; + double _tmp3; + double _tmp4; + double _tmp5; + double _tmp6; + double _shift; + const int shift_8 = 8; + for (; kk + 3 < max_kk; kk += 4) + { + asm volatile( + "ld $0, 16(%1) \n" + "ldc1 %4, 0(%0) \n" + "lwc1 %6, 0(%1) \n" +#if __mips64 + "dmtc1 $0, %5 \n" +#else + "mtc1 $0, %5 \n" +#endif + "mtc1 %16, %11 \n" + "punpcklwd %6, %6, %6 \n" +#if __mips64 + "daddiu %0, %0, 8 \n" + "daddiu %1, %1, 4 \n" +#else + "addiu %0, %0, 8 \n" + "addiu %1, %1, 4 \n" +#endif + "punpcklbh %8, %4, %5 \n" + "punpckhbh %9, %4, %5 \n" + "punpcklbh %10, %6, %5 \n" + "psllh %8, %8, %11 \n" + "psllh %9, %9, %11 \n" + "psllh %10, %10, %11 \n" + "psrah %8, %8, %11 \n" + "psrah %9, %9, %11 \n" + "psrah %10, %10, %11 \n" + "pmaddhw %4, %8, %10 \n" + "pmaddhw %5, %9, %10 \n" + "paddw %2, %2, %4 \n" + "paddw %3, %3, %5 \n" + : "=r"(pA), + "=r"(pB), + "=f"(_sum0), + "=f"(_sum1), + "=&f"(_tmp0), + "=&f"(_tmp1), + "=&f"(_tmp2), + "=&f"(_tmp3), + "=&f"(_tmp4), + "=&f"(_tmp5), + "=&f"(_tmp6), + "=&f"(_shift) + : "0"(pA), + "1"(pB), + "2"(_sum0), + "3"(_sum1), + "r"(shift_8) + : "memory"); + } +#else // NCNN_GNU_INLINE_ASM + const int8x8_t _zero = __mmi_pzerob_s(); + for (; kk + 3 < max_kk; kk += 4) + { + __builtin_prefetch(pB + 16); + const int8x8_t _pA = __mmi_pldb_s(pA); + const int8x8_t _pB = (int8x8_t)__mmi_pfillw_s(*(const int*)pB); + int16x4_t _pA0 = (int16x4_t)__mmi_punpcklbh_s(_pA, _zero); + int16x4_t _pA1 = (int16x4_t)__mmi_punpckhbh_s(_pA, _zero); + int16x4_t _pB0 = (int16x4_t)__mmi_punpcklbh_s(_pB, _zero); + _pA0 = __mmi_psrah_s(__mmi_psllh_s(_pA0, 8), 8); + _pA1 = __mmi_psrah_s(__mmi_psllh_s(_pA1, 8), 8); + _pB0 = __mmi_psrah_s(__mmi_psllh_s(_pB0, 8), 8); + _sum0 = __mmi_paddw_s(_sum0, __mmi_pmaddhw(_pA0, _pB0)); + _sum1 = __mmi_paddw_s(_sum1, __mmi_pmaddhw(_pA1, _pB0)); + pA += 8; + pB += 4; + } +#endif // NCNN_GNU_INLINE_ASM + _sum0 = __mmi_paddw_s(_sum0, __mmi_punpckhwd_s(_sum0, _sum0)); + _sum1 = __mmi_paddw_s(_sum1, __mmi_punpckhwd_s(_sum1, _sum1)); + sum0_i += _sum0[0]; + sum1_i += _sum1[0]; +#endif // __mips_loongson_mmi + for (; kk + 3 < max_kk; kk += 4) + { + sum0_i += pA[0] * pB[0] + pA[1] * pB[1] + pA[2] * pB[2] + pA[3] * pB[3]; + sum1_i += pA[4] * pB[0] + pA[5] * pB[1] + pA[6] * pB[2] + pA[7] * pB[3]; + pA += 8; + pB += 4; + } + if (kk + 1 < max_kk) + { + sum0_i += pA[0] * pB[0] + pA[1] * pB[1]; + sum1_i += pA[2] * pB[0] + pA[3] * pB[1]; + pA += 4; + pB += 2; + kk += 2; + } + for (; kk < max_kk; kk++) + { + sum0_i += pA[0] * pB[0]; + sum1_i += pA[1] * pB[0]; + pA += 2; + pB++; + } + sum0 += sum0_i * pAD[0] * pBD[0]; + sum1 += sum1_i * pAD[1] * pBD[0]; + pAD += 2; + pBD++; + } + + outptr[0] = sum0; + outptr[1] = sum1; + outptr += 2; + } + + pAT += (size_t)2 * A_hstep; + pAT_descales += (size_t)2 * A_descales_hstep; + } + for (; ii < max_ii; ii++) + { + const signed char* pB = pBT; + const float* pBD = pBT_descales; + + int jj = 0; +#if __mips_msa + const v8i16 _one = __msa_fill_h(1); + for (; jj + 7 < max_jj; jj += 8) + { + const signed char* pA = pAT; + const float* pAD = pAT_descales; + const signed char* pB0 = pB; + const signed char* pB1 = pB + (size_t)4 * K; + const float* pBD0 = pBD; + const float* pBD1 = pBD + (size_t)4 * num_blocks; + v4f32 _fsum0 = (v4f32)__msa_fill_w(0); + v4f32 _fsum1 = (v4f32)__msa_fill_w(0); + for (int k = 0; k < K; k += block_size) + { + const int max_kk = std::min(K - k, block_size); + v4i32 _sum0 = __msa_fill_w(0); + v4i32 _sum1 = __msa_fill_w(0); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + const v16i8 _pA = (v16i8)__msa_fill_w(*(const int*)pA); + const v16i8 _pB0 = __msa_ld_b(pB0, 0); + const v16i8 _pB1 = __msa_ld_b(pB1, 0); + _sum0 = __msa_dpadd_s_w(_sum0, __msa_dotp_s_h(_pA, _pB0), _one); + _sum1 = __msa_dpadd_s_w(_sum1, __msa_dotp_s_h(_pA, _pB1), _one); + pA += 4; + pB0 += 16; + pB1 += 16; + } + if (kk + 1 < max_kk) + { + const v16i8 _pA = (v16i8)__msa_fill_h(*(const short*)pA); + const v8i16 _s0 = __msa_dotp_s_h(_pA, (v16i8)__msa_fill_d_ptr(pB0)); + const v8i16 _s1 = __msa_dotp_s_h(_pA, (v16i8)__msa_fill_d_ptr(pB1)); + _sum0 = __msa_addv_w(_sum0, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s0, 0), _s0)); + _sum1 = __msa_addv_w(_sum1, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s1, 0), _s1)); + pA += 2; + pB0 += 8; + pB1 += 8; + kk += 2; + } + if (kk < max_kk) + { + const v8i16 _pA = __msa_fill_h(pA[0]); + const v16i8 _pB08 = (v16i8)__msa_fill_w(*(const int*)pB0); + const v16i8 _pB18 = (v16i8)__msa_fill_w(*(const int*)pB1); + const v8i16 _pB0 = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pB08, 0), _pB08); + const v8i16 _pB1 = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pB18, 0), _pB18); + const v8i16 _s0 = __msa_mulv_h(_pA, _pB0); + const v8i16 _s1 = __msa_mulv_h(_pA, _pB1); + _sum0 = __msa_addv_w(_sum0, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s0, 0), _s0)); + _sum1 = __msa_addv_w(_sum1, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s1, 0), _s1)); + pA++; + pB0 += 4; + pB1 += 4; + } + const v4f32 _descaleB0 = (v4f32)__msa_ld_w(pBD0, 0); + const v4f32 _descaleB1 = (v4f32)__msa_ld_w(pBD1, 0); + const v4f32 _descaleA = __msa_fill_w_f32(pAD[0]); + _fsum0 = __msa_fadd_w(_fsum0, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum0), __msa_fmul_w(_descaleA, _descaleB0))); + _fsum1 = __msa_fadd_w(_fsum1, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum1), __msa_fmul_w(_descaleA, _descaleB1))); + pAD++; + pBD0 += 4; + pBD1 += 4; + } + __msa_st_w((v4i32)_fsum0, outptr, 0); + __msa_st_w((v4i32)_fsum1, outptr + 4, 0); + outptr += 8; + pB = pB1; + pBD = pBD1; + } + for (; jj + 3 < max_jj; jj += 4) + { + const signed char* pA = pAT; + const float* pAD = pAT_descales; + v4f32 _fsum0 = (v4f32)__msa_fill_w(0); + for (int k = 0; k < K; k += block_size) + { + const int max_kk = std::min(K - k, block_size); + v4i32 _sum0 = __msa_fill_w(0); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + const v16i8 _pA = (v16i8)__msa_fill_w(*(const int*)pA); + const v16i8 _pB0 = __msa_ld_b(pB, 0); + _sum0 = __msa_dpadd_s_w(_sum0, __msa_dotp_s_h(_pA, _pB0), _one); + pA += 4; + pB += 16; + } + if (kk + 1 < max_kk) + { + const v16i8 _pA = (v16i8)__msa_fill_h(*(const short*)pA); + const v8i16 _s = __msa_dotp_s_h(_pA, (v16i8)__msa_fill_d_ptr(pB)); + _sum0 = __msa_addv_w(_sum0, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + pA += 2; + pB += 8; + kk += 2; + } + if (kk < max_kk) + { + const v16i8 _pB8 = (v16i8)__msa_fill_w(*(const int*)pB); + const v8i16 _pB = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pB8, 0), _pB8); + const v8i16 _s = __msa_mulv_h(__msa_fill_h(pA[0]), _pB); + _sum0 = __msa_addv_w(_sum0, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + pA++; + pB += 4; + } + const v4f32 _descaleB = (v4f32)__msa_ld_w(pBD, 0); + _fsum0 = __msa_fadd_w(_fsum0, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum0), __msa_fmul_w(_descaleB, __msa_fill_w_f32(pAD[0])))); + pAD++; + pBD += 4; + } + __msa_st_w((v4i32)_fsum0, outptr, 0); + outptr += 4; + } +#endif // __mips_msa + for (; jj + 1 < max_jj; jj += 2) + { + const signed char* pA = pAT; + const float* pAD = pAT_descales; + float sum0 = 0.f; + float sum1 = 0.f; + + for (int k = 0; k < K; k += block_size) + { + const int max_kk = std::min(K - k, block_size); + int sum0_i = 0; + int sum1_i = 0; + int kk = 0; +#if __mips_loongson_mmi + int32x2_t _sum0 = __mmi_pzerow_s(); + int32x2_t _sum1 = __mmi_pzerow_s(); +#if NCNN_GNU_INLINE_ASM + double _tmp0; + double _tmp1; + double _tmp2; + double _tmp3; + double _tmp4; + double _tmp5; + double _tmp6; + double _shift; + const int shift_8 = 8; + for (; kk + 3 < max_kk; kk += 4) + { + asm volatile( + "ld $0, 32(%1) \n" + "lwc1 %4, 0(%0) \n" + "ldc1 %6, 0(%1) \n" +#if __mips64 + "dmtc1 $0, %5 \n" +#else + "mtc1 $0, %5 \n" +#endif + "mtc1 %16, %11 \n" + "punpcklwd %4, %4, %4 \n" +#if __mips64 + "daddiu %0, %0, 4 \n" + "daddiu %1, %1, 8 \n" +#else + "addiu %0, %0, 4 \n" + "addiu %1, %1, 8 \n" +#endif + "punpcklbh %8, %4, %5 \n" + "punpcklbh %9, %6, %5 \n" + "punpckhbh %10, %6, %5 \n" + "psllh %8, %8, %11 \n" + "psllh %9, %9, %11 \n" + "psllh %10, %10, %11 \n" + "psrah %8, %8, %11 \n" + "psrah %9, %9, %11 \n" + "psrah %10, %10, %11 \n" + "pmaddhw %4, %8, %9 \n" + "pmaddhw %5, %8, %10 \n" + "paddw %2, %2, %4 \n" + "paddw %3, %3, %5 \n" + : "=r"(pA), + "=r"(pB), + "=f"(_sum0), + "=f"(_sum1), + "=&f"(_tmp0), + "=&f"(_tmp1), + "=&f"(_tmp2), + "=&f"(_tmp3), + "=&f"(_tmp4), + "=&f"(_tmp5), + "=&f"(_tmp6), + "=&f"(_shift) + : "0"(pA), + "1"(pB), + "2"(_sum0), + "3"(_sum1), + "r"(shift_8) + : "memory"); + } +#else // NCNN_GNU_INLINE_ASM + const int8x8_t _zero = __mmi_pzerob_s(); + for (; kk + 3 < max_kk; kk += 4) + { + __builtin_prefetch(pB + 32); + const int8x8_t _pA = (int8x8_t)__mmi_pfillw_s(*(const int*)pA); + const int8x8_t _pB = __mmi_pldb_s(pB); + int16x4_t _pA0 = (int16x4_t)__mmi_punpcklbh_s(_pA, _zero); + int16x4_t _pB0 = (int16x4_t)__mmi_punpcklbh_s(_pB, _zero); + int16x4_t _pB1 = (int16x4_t)__mmi_punpckhbh_s(_pB, _zero); + _pA0 = __mmi_psrah_s(__mmi_psllh_s(_pA0, 8), 8); + _pB0 = __mmi_psrah_s(__mmi_psllh_s(_pB0, 8), 8); + _pB1 = __mmi_psrah_s(__mmi_psllh_s(_pB1, 8), 8); + _sum0 = __mmi_paddw_s(_sum0, __mmi_pmaddhw(_pA0, _pB0)); + _sum1 = __mmi_paddw_s(_sum1, __mmi_pmaddhw(_pA0, _pB1)); + pA += 4; + pB += 8; + } +#endif // NCNN_GNU_INLINE_ASM + _sum0 = __mmi_paddw_s(_sum0, __mmi_punpckhwd_s(_sum0, _sum0)); + _sum1 = __mmi_paddw_s(_sum1, __mmi_punpckhwd_s(_sum1, _sum1)); + sum0_i += _sum0[0]; + sum1_i += _sum1[0]; +#endif // __mips_loongson_mmi + for (; kk + 3 < max_kk; kk += 4) + { + sum0_i += pA[0] * pB[0] + pA[1] * pB[1] + pA[2] * pB[2] + pA[3] * pB[3]; + sum1_i += pA[0] * pB[4] + pA[1] * pB[5] + pA[2] * pB[6] + pA[3] * pB[7]; + pA += 4; + pB += 8; + } + if (kk + 1 < max_kk) + { + sum0_i += pA[0] * pB[0] + pA[1] * pB[1]; + sum1_i += pA[0] * pB[2] + pA[1] * pB[3]; + pA += 2; + pB += 4; + kk += 2; + } + for (; kk < max_kk; kk++) + { + sum0_i += pA[0] * pB[0]; + sum1_i += pA[0] * pB[1]; + pA++; + pB += 2; + } + sum0 += sum0_i * pAD[0] * pBD[0]; + sum1 += sum1_i * pAD[0] * pBD[1]; + pAD++; + pBD += 2; + } + + outptr[0] = sum0; + outptr[1] = sum1; + outptr += 2; + } + for (; jj < max_jj; jj++) + { + const signed char* pA = pAT; + const float* pAD = pAT_descales; + float sum0 = 0.f; + + for (int k = 0; k < K; k += block_size) + { + const int max_kk = std::min(K - k, block_size); + int sum0_i = 0; + int kk = 0; +#if __mips_loongson_mmi + int32x2_t _sum0 = __mmi_pzerow_s(); +#if NCNN_GNU_INLINE_ASM + double _tmp0; + double _tmp1; + double _tmp2; + double _tmp3; + double _tmp4; + double _tmp5; + double _shift; + const int shift_8 = 8; + for (; kk + 3 < max_kk; kk += 4) + { + asm volatile( + "ld $0, 16(%1) \n" + "lwc1 %3, 0(%0) \n" + "lwc1 %5, 0(%1) \n" +#if __mips64 + "dmtc1 $0, %4 \n" +#else + "mtc1 $0, %4 \n" +#endif + "mtc1 %13, %9 \n" + "punpcklwd %3, %3, %3 \n" + "punpcklwd %5, %5, %5 \n" +#if __mips64 + "daddiu %0, %0, 4 \n" + "daddiu %1, %1, 4 \n" +#else + "addiu %0, %0, 4 \n" + "addiu %1, %1, 4 \n" +#endif + "punpcklbh %7, %3, %4 \n" + "punpcklbh %8, %5, %4 \n" + "psllh %7, %7, %9 \n" + "psllh %8, %8, %9 \n" + "psrah %7, %7, %9 \n" + "psrah %8, %8, %9 \n" + "pmaddhw %3, %7, %8 \n" + "paddw %2, %2, %3 \n" + : "=r"(pA), + "=r"(pB), + "=f"(_sum0), + "=&f"(_tmp0), + "=&f"(_tmp1), + "=&f"(_tmp2), + "=&f"(_tmp3), + "=&f"(_tmp4), + "=&f"(_tmp5), + "=&f"(_shift) + : "0"(pA), + "1"(pB), + "2"(_sum0), + "r"(shift_8) + : "memory"); + } +#else // NCNN_GNU_INLINE_ASM + const int8x8_t _zero = __mmi_pzerob_s(); + for (; kk + 3 < max_kk; kk += 4) + { + const int8x8_t _pA = (int8x8_t)__mmi_pfillw_s(*(const int*)pA); + const int8x8_t _pB = (int8x8_t)__mmi_pfillw_s(*(const int*)pB); + int16x4_t _pA0 = (int16x4_t)__mmi_punpcklbh_s(_pA, _zero); + int16x4_t _pB0 = (int16x4_t)__mmi_punpcklbh_s(_pB, _zero); + _pA0 = __mmi_psrah_s(__mmi_psllh_s(_pA0, 8), 8); + _pB0 = __mmi_psrah_s(__mmi_psllh_s(_pB0, 8), 8); + _sum0 = __mmi_paddw_s(_sum0, __mmi_pmaddhw(_pA0, _pB0)); + pA += 4; + pB += 4; + } +#endif // NCNN_GNU_INLINE_ASM + _sum0 = __mmi_paddw_s(_sum0, __mmi_punpckhwd_s(_sum0, _sum0)); + sum0_i += _sum0[0]; +#endif // __mips_loongson_mmi + for (; kk + 3 < max_kk; kk += 4) + { + sum0_i += pA[0] * pB[0] + pA[1] * pB[1] + pA[2] * pB[2] + pA[3] * pB[3]; + pA += 4; + pB += 4; + } + if (kk + 1 < max_kk) + { + sum0_i += pA[0] * pB[0] + pA[1] * pB[1]; + pA += 2; + pB += 2; + kk += 2; + } + for (; kk < max_kk; kk++) + sum0_i += *pA++ * *pB++; + sum0 += sum0_i * pAD[0] * pBD[0]; + pAD++; + pBD++; + } + + *outptr++ = sum0; + } + + pAT += A_hstep; + pAT_descales += A_descales_hstep; + } +} + +static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, float alpha, float beta) +{ + const size_t out_hstep = top_blob.dims == 3 ? top_blob.cstep : (size_t)top_blob.w; + const size_t c_hstep = C.dims == 3 ? C.cstep : (size_t)C.w; + float* outptr = (float*)top_blob + (size_t)i * out_hstep + j; + + int ii = 0; +#if __mips_msa + for (; ii + 7 < max_ii; ii += 8) + { + float* outptr0 = outptr; + float* outptr1 = outptr0 + out_hstep; + float* outptr2 = outptr0 + out_hstep * 2; + float* outptr3 = outptr0 + out_hstep * 3; + float* outptr4 = outptr0 + out_hstep * 4; + float* outptr5 = outptr0 + out_hstep * 5; + float* outptr6 = outptr0 + out_hstep * 6; + float* outptr7 = outptr0 + out_hstep * 7; + + const float* pC = C; + if (pC) + { + if (broadcast_type_C == 1 || broadcast_type_C == 2) + pC += i + ii; + if (broadcast_type_C == 3) + pC += (size_t)(i + ii) * c_hstep + j; + if (broadcast_type_C == 4) + pC += j; + } + + v4f32 _c0123 = (v4f32)__msa_fill_w(0); + v4f32 _c4567 = (v4f32)__msa_fill_w(0); + if (pC) + { + if (broadcast_type_C == 0) + { + float c = pC[0]; + if (beta != 1.f) + c *= beta; + _c0123 = __msa_fill_w_f32(c); + _c4567 = _c0123; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _c0123 = (v4f32)__msa_ld_w(pC, 0); + _c4567 = (v4f32)__msa_ld_w(pC + 4, 0); + if (beta != 1.f) + { + const v4f32 _beta = __msa_fill_w_f32(beta); + _c0123 = __msa_fmul_w(_c0123, _beta); + _c4567 = __msa_fmul_w(_c4567, _beta); + } + } + } + + const float* pC0 = pC && broadcast_type_C == 3 ? pC : 0; + const float* pC1 = pC0 ? pC0 + c_hstep : 0; + const float* pC2 = pC0 ? pC0 + c_hstep * 2 : 0; + const float* pC3 = pC0 ? pC0 + c_hstep * 3 : 0; + const float* pC4 = pC0 ? pC0 + c_hstep * 4 : 0; + const float* pC5 = pC0 ? pC0 + c_hstep * 5 : 0; + const float* pC6 = pC0 ? pC0 + c_hstep * 6 : 0; + const float* pC7 = pC0 ? pC0 + c_hstep * 7 : 0; + + int jj = 0; + for (; jj + 3 < max_jj; jj += 4) + { + v4f32 _f0 = (v4f32)__msa_ld_w(pp, 0); + v4f32 _f4 = (v4f32)__msa_ld_w(pp + 4, 0); + v4f32 _f1 = (v4f32)__msa_ld_w(pp + 8, 0); + v4f32 _f5 = (v4f32)__msa_ld_w(pp + 12, 0); + v4f32 _f2 = (v4f32)__msa_ld_w(pp + 16, 0); + v4f32 _f6 = (v4f32)__msa_ld_w(pp + 20, 0); + v4f32 _f3 = (v4f32)__msa_ld_w(pp + 24, 0); + v4f32 _f7 = (v4f32)__msa_ld_w(pp + 28, 0); + pp += 32; + transpose4x4_ps(_f0, _f1, _f2, _f3); + transpose4x4_ps(_f4, _f5, _f6, _f7); + + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = __msa_fadd_w(_f0, (v4f32)__msa_splati_w((v4i32)_c0123, 0)); + _f1 = __msa_fadd_w(_f1, (v4f32)__msa_splati_w((v4i32)_c0123, 1)); + _f2 = __msa_fadd_w(_f2, (v4f32)__msa_splati_w((v4i32)_c0123, 2)); + _f3 = __msa_fadd_w(_f3, (v4f32)__msa_splati_w((v4i32)_c0123, 3)); + _f4 = __msa_fadd_w(_f4, (v4f32)__msa_splati_w((v4i32)_c4567, 0)); + _f5 = __msa_fadd_w(_f5, (v4f32)__msa_splati_w((v4i32)_c4567, 1)); + _f6 = __msa_fadd_w(_f6, (v4f32)__msa_splati_w((v4i32)_c4567, 2)); + _f7 = __msa_fadd_w(_f7, (v4f32)__msa_splati_w((v4i32)_c4567, 3)); + } + if (broadcast_type_C == 3) + { + const v4f32 _beta = __msa_fill_w_f32(beta); + v4f32 _c = (v4f32)__msa_ld_w(pC0, 0); + _f0 = __msa_fadd_w(_f0, beta == 1.f ? _c : __msa_fmul_w(_c, _beta)); + _c = (v4f32)__msa_ld_w(pC1, 0); + _f1 = __msa_fadd_w(_f1, beta == 1.f ? _c : __msa_fmul_w(_c, _beta)); + _c = (v4f32)__msa_ld_w(pC2, 0); + _f2 = __msa_fadd_w(_f2, beta == 1.f ? _c : __msa_fmul_w(_c, _beta)); + _c = (v4f32)__msa_ld_w(pC3, 0); + _f3 = __msa_fadd_w(_f3, beta == 1.f ? _c : __msa_fmul_w(_c, _beta)); + _c = (v4f32)__msa_ld_w(pC4, 0); + _f4 = __msa_fadd_w(_f4, beta == 1.f ? _c : __msa_fmul_w(_c, _beta)); + _c = (v4f32)__msa_ld_w(pC5, 0); + _f5 = __msa_fadd_w(_f5, beta == 1.f ? _c : __msa_fmul_w(_c, _beta)); + _c = (v4f32)__msa_ld_w(pC6, 0); + _f6 = __msa_fadd_w(_f6, beta == 1.f ? _c : __msa_fmul_w(_c, _beta)); + _c = (v4f32)__msa_ld_w(pC7, 0); + _f7 = __msa_fadd_w(_f7, beta == 1.f ? _c : __msa_fmul_w(_c, _beta)); + } + if (broadcast_type_C == 4) + { + v4f32 _c = (v4f32)__msa_ld_w(pC, 0); + if (beta != 1.f) + _c = __msa_fmul_w(_c, __msa_fill_w_f32(beta)); + _f0 = __msa_fadd_w(_f0, _c); + _f1 = __msa_fadd_w(_f1, _c); + _f2 = __msa_fadd_w(_f2, _c); + _f3 = __msa_fadd_w(_f3, _c); + _f4 = __msa_fadd_w(_f4, _c); + _f5 = __msa_fadd_w(_f5, _c); + _f6 = __msa_fadd_w(_f6, _c); + _f7 = __msa_fadd_w(_f7, _c); + } + } + + if (alpha != 1.f) + { + const v4f32 _alpha = __msa_fill_w_f32(alpha); + _f0 = __msa_fmul_w(_f0, _alpha); + _f1 = __msa_fmul_w(_f1, _alpha); + _f2 = __msa_fmul_w(_f2, _alpha); + _f3 = __msa_fmul_w(_f3, _alpha); + _f4 = __msa_fmul_w(_f4, _alpha); + _f5 = __msa_fmul_w(_f5, _alpha); + _f6 = __msa_fmul_w(_f6, _alpha); + _f7 = __msa_fmul_w(_f7, _alpha); + } + __msa_st_w((v4i32)_f0, outptr0, 0); + __msa_st_w((v4i32)_f1, outptr1, 0); + __msa_st_w((v4i32)_f2, outptr2, 0); + __msa_st_w((v4i32)_f3, outptr3, 0); + __msa_st_w((v4i32)_f4, outptr4, 0); + __msa_st_w((v4i32)_f5, outptr5, 0); + __msa_st_w((v4i32)_f6, outptr6, 0); + __msa_st_w((v4i32)_f7, outptr7, 0); + outptr0 += 4; + outptr1 += 4; + outptr2 += 4; + outptr3 += 4; + outptr4 += 4; + outptr5 += 4; + outptr6 += 4; + outptr7 += 4; + if (pC0) + { + pC0 += 4; + pC1 += 4; + pC2 += 4; + pC3 += 4; + pC4 += 4; + pC5 += 4; + pC6 += 4; + pC7 += 4; + } + if (pC && broadcast_type_C == 4) + pC += 4; + } + for (; jj + 1 < max_jj; jj += 2) + { + v4f32 _f0 = (v4f32)__msa_ld_w(pp, 0); + v4f32 _f4 = (v4f32)__msa_ld_w(pp + 4, 0); + v4f32 _f1 = (v4f32)__msa_ld_w(pp + 8, 0); + v4f32 _f5 = (v4f32)__msa_ld_w(pp + 12, 0); + pp += 16; + + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = __msa_fadd_w(_f0, _c0123); + _f4 = __msa_fadd_w(_f4, _c4567); + _f1 = __msa_fadd_w(_f1, _c0123); + _f5 = __msa_fadd_w(_f5, _c4567); + } + if (broadcast_type_C == 3) + { + v4f32 _c0 = (v4f32){pC0[0], pC1[0], pC2[0], pC3[0]}; + v4f32 _c4 = (v4f32){pC4[0], pC5[0], pC6[0], pC7[0]}; + v4f32 _c1 = (v4f32){pC0[1], pC1[1], pC2[1], pC3[1]}; + v4f32 _c5 = (v4f32){pC4[1], pC5[1], pC6[1], pC7[1]}; + if (beta != 1.f) + { + const v4f32 _beta = __msa_fill_w_f32(beta); + _c0 = __msa_fmul_w(_c0, _beta); + _c4 = __msa_fmul_w(_c4, _beta); + _c1 = __msa_fmul_w(_c1, _beta); + _c5 = __msa_fmul_w(_c5, _beta); + } + _f0 = __msa_fadd_w(_f0, _c0); + _f4 = __msa_fadd_w(_f4, _c4); + _f1 = __msa_fadd_w(_f1, _c1); + _f5 = __msa_fadd_w(_f5, _c5); + } + if (broadcast_type_C == 4) + { + float c0 = pC[0]; + float c1 = pC[1]; + if (beta != 1.f) + { + c0 *= beta; + c1 *= beta; + } + _f0 = __msa_fadd_w(_f0, __msa_fill_w_f32(c0)); + _f4 = __msa_fadd_w(_f4, __msa_fill_w_f32(c0)); + _f1 = __msa_fadd_w(_f1, __msa_fill_w_f32(c1)); + _f5 = __msa_fadd_w(_f5, __msa_fill_w_f32(c1)); + } + } + + if (alpha != 1.f) + { + const v4f32 _alpha = __msa_fill_w_f32(alpha); + _f0 = __msa_fmul_w(_f0, _alpha); + _f4 = __msa_fmul_w(_f4, _alpha); + _f1 = __msa_fmul_w(_f1, _alpha); + _f5 = __msa_fmul_w(_f5, _alpha); + } + ((int*)outptr0)[0] = __msa_copy_s_w((v4i32)_f0, 0); + ((int*)outptr0)[1] = __msa_copy_s_w((v4i32)_f1, 0); + ((int*)outptr1)[0] = __msa_copy_s_w((v4i32)_f0, 1); + ((int*)outptr1)[1] = __msa_copy_s_w((v4i32)_f1, 1); + ((int*)outptr2)[0] = __msa_copy_s_w((v4i32)_f0, 2); + ((int*)outptr2)[1] = __msa_copy_s_w((v4i32)_f1, 2); + ((int*)outptr3)[0] = __msa_copy_s_w((v4i32)_f0, 3); + ((int*)outptr3)[1] = __msa_copy_s_w((v4i32)_f1, 3); + ((int*)outptr4)[0] = __msa_copy_s_w((v4i32)_f4, 0); + ((int*)outptr4)[1] = __msa_copy_s_w((v4i32)_f5, 0); + ((int*)outptr5)[0] = __msa_copy_s_w((v4i32)_f4, 1); + ((int*)outptr5)[1] = __msa_copy_s_w((v4i32)_f5, 1); + ((int*)outptr6)[0] = __msa_copy_s_w((v4i32)_f4, 2); + ((int*)outptr6)[1] = __msa_copy_s_w((v4i32)_f5, 2); + ((int*)outptr7)[0] = __msa_copy_s_w((v4i32)_f4, 3); + ((int*)outptr7)[1] = __msa_copy_s_w((v4i32)_f5, 3); + outptr0 += 2; + outptr1 += 2; + outptr2 += 2; + outptr3 += 2; + outptr4 += 2; + outptr5 += 2; + outptr6 += 2; + outptr7 += 2; + if (pC0) + { + pC0 += 2; + pC1 += 2; + pC2 += 2; + pC3 += 2; + pC4 += 2; + pC5 += 2; + pC6 += 2; + pC7 += 2; + } + if (pC && broadcast_type_C == 4) + pC += 2; + } + for (; jj < max_jj; jj++) + { + v4f32 _f0 = (v4f32)__msa_ld_w(pp, 0); + v4f32 _f4 = (v4f32)__msa_ld_w(pp + 4, 0); + pp += 8; + + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = __msa_fadd_w(_f0, _c0123); + _f4 = __msa_fadd_w(_f4, _c4567); + } + if (broadcast_type_C == 3) + { + v4f32 _c0 = (v4f32){pC0[0], pC1[0], pC2[0], pC3[0]}; + v4f32 _c4 = (v4f32){pC4[0], pC5[0], pC6[0], pC7[0]}; + if (beta != 1.f) + { + const v4f32 _beta = __msa_fill_w_f32(beta); + _c0 = __msa_fmul_w(_c0, _beta); + _c4 = __msa_fmul_w(_c4, _beta); + } + _f0 = __msa_fadd_w(_f0, _c0); + _f4 = __msa_fadd_w(_f4, _c4); + } + if (broadcast_type_C == 4) + { + float c = pC[0]; + if (beta != 1.f) + c *= beta; + _f0 = __msa_fadd_w(_f0, __msa_fill_w_f32(c)); + _f4 = __msa_fadd_w(_f4, __msa_fill_w_f32(c)); + } + } + + if (alpha != 1.f) + { + const v4f32 _alpha = __msa_fill_w_f32(alpha); + _f0 = __msa_fmul_w(_f0, _alpha); + _f4 = __msa_fmul_w(_f4, _alpha); + } + ((int*)outptr0)[0] = __msa_copy_s_w((v4i32)_f0, 0); + ((int*)outptr1)[0] = __msa_copy_s_w((v4i32)_f0, 1); + ((int*)outptr2)[0] = __msa_copy_s_w((v4i32)_f0, 2); + ((int*)outptr3)[0] = __msa_copy_s_w((v4i32)_f0, 3); + ((int*)outptr4)[0] = __msa_copy_s_w((v4i32)_f4, 0); + ((int*)outptr5)[0] = __msa_copy_s_w((v4i32)_f4, 1); + ((int*)outptr6)[0] = __msa_copy_s_w((v4i32)_f4, 2); + ((int*)outptr7)[0] = __msa_copy_s_w((v4i32)_f4, 3); + outptr0++; + outptr1++; + outptr2++; + outptr3++; + outptr4++; + outptr5++; + outptr6++; + outptr7++; + } + outptr += out_hstep * 8; + } + for (; ii + 3 < max_ii; ii += 4) + { + float* outptr0 = outptr; + float* outptr1 = outptr0 + out_hstep; + float* outptr2 = outptr0 + out_hstep * 2; + float* outptr3 = outptr0 + out_hstep * 3; + + const float* pC = C; + if (pC) + { + if (broadcast_type_C == 1 || broadcast_type_C == 2) + pC += i + ii; + if (broadcast_type_C == 3) + pC += (size_t)(i + ii) * c_hstep + j; + if (broadcast_type_C == 4) + pC += j; + } + + v4f32 _c0123 = (v4f32)__msa_fill_w(0); + if (pC) + { + if (broadcast_type_C == 0) + { + float c = pC[0]; + if (beta != 1.f) + c *= beta; + _c0123 = __msa_fill_w_f32(c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _c0123 = (v4f32)__msa_ld_w(pC, 0); + if (beta != 1.f) + _c0123 = __msa_fmul_w(_c0123, __msa_fill_w_f32(beta)); + } + } + + const float* pC0 = pC && broadcast_type_C == 3 ? pC : 0; + const float* pC1 = pC0 ? pC0 + c_hstep : 0; + const float* pC2 = pC0 ? pC0 + c_hstep * 2 : 0; + const float* pC3 = pC0 ? pC0 + c_hstep * 3 : 0; + + int jj = 0; + for (; jj + 7 < max_jj; jj += 8) + { + v4f32 _f0 = (v4f32)__msa_ld_w(pp, 0); + v4f32 _f1 = (v4f32)__msa_ld_w(pp + 4, 0); + v4f32 _f2 = (v4f32)__msa_ld_w(pp + 8, 0); + v4f32 _f3 = (v4f32)__msa_ld_w(pp + 12, 0); + v4f32 _f4 = (v4f32)__msa_ld_w(pp + 16, 0); + v4f32 _f5 = (v4f32)__msa_ld_w(pp + 20, 0); + v4f32 _f6 = (v4f32)__msa_ld_w(pp + 24, 0); + v4f32 _f7 = (v4f32)__msa_ld_w(pp + 28, 0); + pp += 32; + + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = __msa_fadd_w(_f0, _c0123); + _f1 = __msa_fadd_w(_f1, _c0123); + _f2 = __msa_fadd_w(_f2, _c0123); + _f3 = __msa_fadd_w(_f3, _c0123); + _f4 = __msa_fadd_w(_f4, _c0123); + _f5 = __msa_fadd_w(_f5, _c0123); + _f6 = __msa_fadd_w(_f6, _c0123); + _f7 = __msa_fadd_w(_f7, _c0123); + } + if (broadcast_type_C == 3) + { + v4f32 _c0 = (v4f32)__msa_ld_w(pC0, 0); + v4f32 _c1 = (v4f32)__msa_ld_w(pC1, 0); + v4f32 _c2 = (v4f32)__msa_ld_w(pC2, 0); + v4f32 _c3 = (v4f32)__msa_ld_w(pC3, 0); + transpose4x4_ps(_c0, _c1, _c2, _c3); + v4f32 _c4 = (v4f32)__msa_ld_w(pC0 + 4, 0); + v4f32 _c5 = (v4f32)__msa_ld_w(pC1 + 4, 0); + v4f32 _c6 = (v4f32)__msa_ld_w(pC2 + 4, 0); + v4f32 _c7 = (v4f32)__msa_ld_w(pC3 + 4, 0); + transpose4x4_ps(_c4, _c5, _c6, _c7); + if (beta != 1.f) + { + v4f32 _beta = __msa_fill_w_f32(beta); + _c0 = __msa_fmul_w(_c0, _beta); + _c1 = __msa_fmul_w(_c1, _beta); + _c2 = __msa_fmul_w(_c2, _beta); + _c3 = __msa_fmul_w(_c3, _beta); + _c4 = __msa_fmul_w(_c4, _beta); + _c5 = __msa_fmul_w(_c5, _beta); + _c6 = __msa_fmul_w(_c6, _beta); + _c7 = __msa_fmul_w(_c7, _beta); + } + _f0 = __msa_fadd_w(_f0, _c0); + _f1 = __msa_fadd_w(_f1, _c1); + _f2 = __msa_fadd_w(_f2, _c2); + _f3 = __msa_fadd_w(_f3, _c3); + _f4 = __msa_fadd_w(_f4, _c4); + _f5 = __msa_fadd_w(_f5, _c5); + _f6 = __msa_fadd_w(_f6, _c6); + _f7 = __msa_fadd_w(_f7, _c7); + } + if (broadcast_type_C == 4) + { + v4f32 _c0 = (v4f32)__msa_ld_w(pC, 0); + v4f32 _c4 = (v4f32)__msa_ld_w(pC + 4, 0); + if (beta != 1.f) + { + v4f32 _beta = __msa_fill_w_f32(beta); + _c0 = __msa_fmul_w(_c0, _beta); + _c4 = __msa_fmul_w(_c4, _beta); + } + _f0 = __msa_fadd_w(_f0, (v4f32)__msa_splati_w((v4i32)_c0, 0)); + _f1 = __msa_fadd_w(_f1, (v4f32)__msa_splati_w((v4i32)_c0, 1)); + _f2 = __msa_fadd_w(_f2, (v4f32)__msa_splati_w((v4i32)_c0, 2)); + _f3 = __msa_fadd_w(_f3, (v4f32)__msa_splati_w((v4i32)_c0, 3)); + _f4 = __msa_fadd_w(_f4, (v4f32)__msa_splati_w((v4i32)_c4, 0)); + _f5 = __msa_fadd_w(_f5, (v4f32)__msa_splati_w((v4i32)_c4, 1)); + _f6 = __msa_fadd_w(_f6, (v4f32)__msa_splati_w((v4i32)_c4, 2)); + _f7 = __msa_fadd_w(_f7, (v4f32)__msa_splati_w((v4i32)_c4, 3)); + } + } + + if (alpha != 1.f) + { + v4f32 _alpha = __msa_fill_w_f32(alpha); + _f0 = __msa_fmul_w(_f0, _alpha); + _f1 = __msa_fmul_w(_f1, _alpha); + _f2 = __msa_fmul_w(_f2, _alpha); + _f3 = __msa_fmul_w(_f3, _alpha); + _f4 = __msa_fmul_w(_f4, _alpha); + _f5 = __msa_fmul_w(_f5, _alpha); + _f6 = __msa_fmul_w(_f6, _alpha); + _f7 = __msa_fmul_w(_f7, _alpha); + } + + transpose4x4_ps(_f0, _f1, _f2, _f3); + transpose4x4_ps(_f4, _f5, _f6, _f7); + __msa_st_w((v4i32)_f0, outptr0, 0); + __msa_st_w((v4i32)_f4, outptr0 + 4, 0); + __msa_st_w((v4i32)_f1, outptr1, 0); + __msa_st_w((v4i32)_f5, outptr1 + 4, 0); + __msa_st_w((v4i32)_f2, outptr2, 0); + __msa_st_w((v4i32)_f6, outptr2 + 4, 0); + __msa_st_w((v4i32)_f3, outptr3, 0); + __msa_st_w((v4i32)_f7, outptr3 + 4, 0); + outptr0 += 8; + outptr1 += 8; + outptr2 += 8; + outptr3 += 8; + if (pC0) + { + pC0 += 8; + pC1 += 8; + pC2 += 8; + pC3 += 8; + } + if (pC && broadcast_type_C == 4) + pC += 8; + } + for (; jj + 3 < max_jj; jj += 4) + { + v4f32 _f0 = (v4f32)__msa_ld_w(pp, 0); + v4f32 _f1 = (v4f32)__msa_ld_w(pp + 4, 0); + v4f32 _f2 = (v4f32)__msa_ld_w(pp + 8, 0); + v4f32 _f3 = (v4f32)__msa_ld_w(pp + 12, 0); + pp += 16; + + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = __msa_fadd_w(_f0, _c0123); + _f1 = __msa_fadd_w(_f1, _c0123); + _f2 = __msa_fadd_w(_f2, _c0123); + _f3 = __msa_fadd_w(_f3, _c0123); + } + if (broadcast_type_C == 3) + { + v4f32 _c0 = (v4f32)__msa_ld_w(pC0, 0); + v4f32 _c1 = (v4f32)__msa_ld_w(pC1, 0); + v4f32 _c2 = (v4f32)__msa_ld_w(pC2, 0); + v4f32 _c3 = (v4f32)__msa_ld_w(pC3, 0); + transpose4x4_ps(_c0, _c1, _c2, _c3); + if (beta != 1.f) + { + v4f32 _beta = __msa_fill_w_f32(beta); + _c0 = __msa_fmul_w(_c0, _beta); + _c1 = __msa_fmul_w(_c1, _beta); + _c2 = __msa_fmul_w(_c2, _beta); + _c3 = __msa_fmul_w(_c3, _beta); + } + _f0 = __msa_fadd_w(_f0, _c0); + _f1 = __msa_fadd_w(_f1, _c1); + _f2 = __msa_fadd_w(_f2, _c2); + _f3 = __msa_fadd_w(_f3, _c3); + } + if (broadcast_type_C == 4) + { + v4f32 _c0 = (v4f32)__msa_ld_w(pC, 0); + if (beta != 1.f) + _c0 = __msa_fmul_w(_c0, __msa_fill_w_f32(beta)); + _f0 = __msa_fadd_w(_f0, (v4f32)__msa_splati_w((v4i32)_c0, 0)); + _f1 = __msa_fadd_w(_f1, (v4f32)__msa_splati_w((v4i32)_c0, 1)); + _f2 = __msa_fadd_w(_f2, (v4f32)__msa_splati_w((v4i32)_c0, 2)); + _f3 = __msa_fadd_w(_f3, (v4f32)__msa_splati_w((v4i32)_c0, 3)); + } + } + + if (alpha != 1.f) + { + v4f32 _alpha = __msa_fill_w_f32(alpha); + _f0 = __msa_fmul_w(_f0, _alpha); + _f1 = __msa_fmul_w(_f1, _alpha); + _f2 = __msa_fmul_w(_f2, _alpha); + _f3 = __msa_fmul_w(_f3, _alpha); + } + + transpose4x4_ps(_f0, _f1, _f2, _f3); + __msa_st_w((v4i32)_f0, outptr0, 0); + __msa_st_w((v4i32)_f1, outptr1, 0); + __msa_st_w((v4i32)_f2, outptr2, 0); + __msa_st_w((v4i32)_f3, outptr3, 0); + outptr0 += 4; + outptr1 += 4; + outptr2 += 4; + outptr3 += 4; + if (pC0) + { + pC0 += 4; + pC1 += 4; + pC2 += 4; + pC3 += 4; + } + if (pC && broadcast_type_C == 4) + pC += 4; + } + for (; jj + 1 < max_jj; jj += 2) + { + v4f32 _f0 = (v4f32)__msa_ld_w(pp, 0); + v4f32 _f1 = (v4f32)__msa_ld_w(pp + 4, 0); + pp += 8; + + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = __msa_fadd_w(_f0, _c0123); + _f1 = __msa_fadd_w(_f1, _c0123); + } + if (broadcast_type_C == 3) + { + v4f32 _c0 = (v4f32){pC0[0], pC1[0], pC2[0], pC3[0]}; + v4f32 _c1 = (v4f32){pC0[1], pC1[1], pC2[1], pC3[1]}; + if (beta != 1.f) + { + v4f32 _beta = __msa_fill_w_f32(beta); + _c0 = __msa_fmul_w(_c0, _beta); + _c1 = __msa_fmul_w(_c1, _beta); + } + _f0 = __msa_fadd_w(_f0, _c0); + _f1 = __msa_fadd_w(_f1, _c1); + } + if (broadcast_type_C == 4) + { + float c0 = pC[0]; + float c1 = pC[1]; + if (beta != 1.f) + { + c0 *= beta; + c1 *= beta; + } + _f0 = __msa_fadd_w(_f0, __msa_fill_w_f32(c0)); + _f1 = __msa_fadd_w(_f1, __msa_fill_w_f32(c1)); + } + } + + if (alpha != 1.f) + { + v4f32 _alpha = __msa_fill_w_f32(alpha); + _f0 = __msa_fmul_w(_f0, _alpha); + _f1 = __msa_fmul_w(_f1, _alpha); + } + + v4f32 _f2 = (v4f32)__msa_fill_w(0); + v4f32 _f3 = (v4f32)__msa_fill_w(0); + transpose4x4_ps(_f0, _f1, _f2, _f3); + *(int64_t*)outptr0 = __msa_copy_s_d((v2i64)_f0, 0); + *(int64_t*)outptr1 = __msa_copy_s_d((v2i64)_f1, 0); + *(int64_t*)outptr2 = __msa_copy_s_d((v2i64)_f2, 0); + *(int64_t*)outptr3 = __msa_copy_s_d((v2i64)_f3, 0); + outptr0 += 2; + outptr1 += 2; + outptr2 += 2; + outptr3 += 2; + if (pC0) + { + pC0 += 2; + pC1 += 2; + pC2 += 2; + pC3 += 2; + } + if (pC && broadcast_type_C == 4) + pC += 2; + } + for (; jj < max_jj; jj++) + { + v4f32 _f0 = (v4f32)__msa_ld_w(pp, 0); + pp += 4; + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = __msa_fadd_w(_f0, _c0123); + } + if (broadcast_type_C == 3) + { + v4f32 _c0 = (v4f32){pC0[0], pC1[0], pC2[0], pC3[0]}; + if (beta != 1.f) + _c0 = __msa_fmul_w(_c0, __msa_fill_w_f32(beta)); + _f0 = __msa_fadd_w(_f0, _c0); + } + if (broadcast_type_C == 4) + { + float c = pC[0]; + if (beta != 1.f) + c *= beta; + _f0 = __msa_fadd_w(_f0, __msa_fill_w_f32(c)); + } + } + if (alpha != 1.f) + _f0 = __msa_fmul_w(_f0, __msa_fill_w_f32(alpha)); + outptr0[0] = _f0[0]; + outptr1[0] = _f0[1]; + outptr2[0] = _f0[2]; + outptr3[0] = _f0[3]; + outptr0++; + outptr1++; + outptr2++; + outptr3++; + if (pC0) + { + pC0++; + pC1++; + pC2++; + pC3++; + } + if (pC && broadcast_type_C == 4) + pC++; + } + outptr += out_hstep * 4; + } +#endif // __mips_msa + for (; ii + 1 < max_ii; ii += 2) + { + float* outptr0 = outptr; + float* outptr1 = outptr0 + out_hstep; + const float* pC = C; + if (pC) + { + if (broadcast_type_C == 1 || broadcast_type_C == 2) + pC += i + ii; + if (broadcast_type_C == 3) + pC += (size_t)(i + ii) * c_hstep + j; + if (broadcast_type_C == 4) + pC += j; + } + const float* pC0 = pC && broadcast_type_C == 3 ? pC : 0; + const float* pC1 = pC0 ? pC0 + c_hstep : 0; + + float c0 = 0.f; + float c1 = 0.f; + if (pC && (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2)) + { + c0 = pC[0]; + c1 = pC[broadcast_type_C == 0 ? 0 : 1]; + if (beta != 1.f) + { + c0 *= beta; + c1 *= beta; + } + } + + int jj = 0; +#if __mips_msa + for (; jj + 7 < max_jj; jj += 8) + { + v4i32 _s0 = __msa_ld_w(pp, 0); + v4i32 _s1 = __msa_ld_w(pp + 4, 0); + v4i32 _s2 = __msa_ld_w(pp + 8, 0); + v4i32 _s3 = __msa_ld_w(pp + 12, 0); + pp += 16; + + v4f32 _f0 = (v4f32)__msa_pckev_w(_s1, _s0); + v4f32 _f1 = (v4f32)__msa_pckev_w(_s3, _s2); + v4f32 _f2 = (v4f32)__msa_pckod_w(_s1, _s0); + v4f32 _f3 = (v4f32)__msa_pckod_w(_s3, _s2); + + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = __msa_fadd_w(_f0, __msa_fill_w_f32(c0)); + _f1 = __msa_fadd_w(_f1, __msa_fill_w_f32(c0)); + _f2 = __msa_fadd_w(_f2, __msa_fill_w_f32(c1)); + _f3 = __msa_fadd_w(_f3, __msa_fill_w_f32(c1)); + } + if (broadcast_type_C == 3) + { + v4f32 _c0 = (v4f32)__msa_ld_w(pC0, 0); + v4f32 _c1 = (v4f32)__msa_ld_w(pC0 + 4, 0); + v4f32 _c2 = (v4f32)__msa_ld_w(pC1, 0); + v4f32 _c3 = (v4f32)__msa_ld_w(pC1 + 4, 0); + if (beta != 1.f) + { + v4f32 _beta = __msa_fill_w_f32(beta); + _c0 = __msa_fmul_w(_c0, _beta); + _c1 = __msa_fmul_w(_c1, _beta); + _c2 = __msa_fmul_w(_c2, _beta); + _c3 = __msa_fmul_w(_c3, _beta); + } + _f0 = __msa_fadd_w(_f0, _c0); + _f1 = __msa_fadd_w(_f1, _c1); + _f2 = __msa_fadd_w(_f2, _c2); + _f3 = __msa_fadd_w(_f3, _c3); + } + if (broadcast_type_C == 4) + { + v4f32 _c0 = (v4f32)__msa_ld_w(pC, 0); + v4f32 _c1 = (v4f32)__msa_ld_w(pC + 4, 0); + if (beta != 1.f) + { + v4f32 _beta = __msa_fill_w_f32(beta); + _c0 = __msa_fmul_w(_c0, _beta); + _c1 = __msa_fmul_w(_c1, _beta); + } + _f0 = __msa_fadd_w(_f0, _c0); + _f1 = __msa_fadd_w(_f1, _c1); + _f2 = __msa_fadd_w(_f2, _c0); + _f3 = __msa_fadd_w(_f3, _c1); + } + } + + if (alpha != 1.f) + { + v4f32 _alpha = __msa_fill_w_f32(alpha); + _f0 = __msa_fmul_w(_f0, _alpha); + _f1 = __msa_fmul_w(_f1, _alpha); + _f2 = __msa_fmul_w(_f2, _alpha); + _f3 = __msa_fmul_w(_f3, _alpha); + } + + __msa_st_w((v4i32)_f0, outptr0, 0); + __msa_st_w((v4i32)_f1, outptr0 + 4, 0); + __msa_st_w((v4i32)_f2, outptr1, 0); + __msa_st_w((v4i32)_f3, outptr1 + 4, 0); + outptr0 += 8; + outptr1 += 8; + if (pC0) + { + pC0 += 8; + pC1 += 8; + } + if (pC && broadcast_type_C == 4) + pC += 8; + } + for (; jj + 3 < max_jj; jj += 4) + { + v4i32 _s0 = __msa_ld_w(pp, 0); + v4i32 _s1 = __msa_ld_w(pp + 4, 0); + pp += 8; + + v4f32 _f0 = (v4f32)__msa_pckev_w(_s1, _s0); + v4f32 _f1 = (v4f32)__msa_pckod_w(_s1, _s0); + + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = __msa_fadd_w(_f0, __msa_fill_w_f32(c0)); + _f1 = __msa_fadd_w(_f1, __msa_fill_w_f32(c1)); + } + if (broadcast_type_C == 3) + { + v4f32 _c0 = (v4f32)__msa_ld_w(pC0, 0); + v4f32 _c1 = (v4f32)__msa_ld_w(pC1, 0); + if (beta != 1.f) + { + v4f32 _beta = __msa_fill_w_f32(beta); + _c0 = __msa_fmul_w(_c0, _beta); + _c1 = __msa_fmul_w(_c1, _beta); + } + _f0 = __msa_fadd_w(_f0, _c0); + _f1 = __msa_fadd_w(_f1, _c1); + } + if (broadcast_type_C == 4) + { + v4f32 _c0 = (v4f32)__msa_ld_w(pC, 0); + if (beta != 1.f) + _c0 = __msa_fmul_w(_c0, __msa_fill_w_f32(beta)); + _f0 = __msa_fadd_w(_f0, _c0); + _f1 = __msa_fadd_w(_f1, _c0); + } + } + + if (alpha != 1.f) + { + v4f32 _alpha = __msa_fill_w_f32(alpha); + _f0 = __msa_fmul_w(_f0, _alpha); + _f1 = __msa_fmul_w(_f1, _alpha); + } + + __msa_st_w((v4i32)_f0, outptr0, 0); + __msa_st_w((v4i32)_f1, outptr1, 0); + outptr0 += 4; + outptr1 += 4; + if (pC0) + { + pC0 += 4; + pC1 += 4; + } + if (pC && broadcast_type_C == 4) + pC += 4; + } + for (; jj + 1 < max_jj; jj += 2) + { + v4f32 _f = (v4f32)__msa_ld_w(pp, 0); + pp += 4; + + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + _f = __msa_fadd_w(_f, (v4f32){c0, c1, c0, c1}); + if (broadcast_type_C == 3) + { + v4f32 _c = (v4f32){pC0[0], pC1[0], pC0[1], pC1[1]}; + if (beta != 1.f) + _c = __msa_fmul_w(_c, __msa_fill_w_f32(beta)); + _f = __msa_fadd_w(_f, _c); + } + if (broadcast_type_C == 4) + { + float cc0 = pC[0]; + float cc1 = pC[1]; + if (beta != 1.f) + { + cc0 *= beta; + cc1 *= beta; + } + _f = __msa_fadd_w(_f, (v4f32){cc0, cc0, cc1, cc1}); + } + } + + if (alpha != 1.f) + _f = __msa_fmul_w(_f, __msa_fill_w_f32(alpha)); + + v4i32 _f0 = __msa_pckev_w((v4i32)_f, (v4i32)_f); + v4i32 _f1 = __msa_pckod_w((v4i32)_f, (v4i32)_f); + __msa_storel_d(_f0, outptr0); + __msa_storel_d(_f1, outptr1); + outptr0 += 2; + outptr1 += 2; + if (pC0) + { + pC0 += 2; + pC1 += 2; + } + if (pC && broadcast_type_C == 4) + pC += 2; + } +#endif // __mips_msa + for (; jj + 1 < max_jj; jj += 2) + { + float sum00 = pp[0]; + float sum01 = pp[1]; + float sum10 = pp[2]; + float sum11 = pp[3]; + pp += 4; + + if (pC) + { + if (broadcast_type_C == 0) + { + sum00 += c0; + sum01 += c1; + sum10 += c0; + sum11 += c1; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + sum00 += c0; + sum10 += c0; + sum01 += c1; + sum11 += c1; + } + if (broadcast_type_C == 3) + { + float c00 = pC0[0]; + float c01 = pC1[0]; + float c10 = pC0[1]; + float c11 = pC1[1]; + if (beta != 1.f) + { + c00 *= beta; + c01 *= beta; + c10 *= beta; + c11 *= beta; + } + sum00 += c00; + sum01 += c01; + sum10 += c10; + sum11 += c11; + } + if (broadcast_type_C == 4) + { + float c0 = pC[0]; + float c1 = pC[1]; + if (beta != 1.f) + { + c0 *= beta; + c1 *= beta; + } + sum00 += c0; + sum01 += c0; + sum10 += c1; + sum11 += c1; + } + } + + if (alpha != 1.f) + { + sum00 *= alpha; + sum01 *= alpha; + sum10 *= alpha; + sum11 *= alpha; + } + + outptr0[0] = sum00; + outptr1[0] = sum01; + outptr0[1] = sum10; + outptr1[1] = sum11; + outptr0 += 2; + outptr1 += 2; + if (pC0) + { + pC0 += 2; + pC1 += 2; + } + if (pC && broadcast_type_C == 4) + pC += 2; + } + for (; jj < max_jj; jj++) + { + float sum0 = pp[0]; + float sum1 = pp[1]; + pp += 2; + if (pC) + { + if (broadcast_type_C == 0) + { + sum0 += c0; + sum1 += c1; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + sum0 += c0; + sum1 += c1; + } + if (broadcast_type_C == 3) + { + float c0 = pC0[0]; + float c1 = pC1[0]; + if (beta != 1.f) + { + c0 *= beta; + c1 *= beta; + } + sum0 += c0; + sum1 += c1; + } + if (broadcast_type_C == 4) + { + float c = pC[0]; + if (beta != 1.f) c *= beta; + sum0 += c; + sum1 += c; + } + } + if (alpha != 1.f) + { + sum0 *= alpha; + sum1 *= alpha; + } + outptr0[0] = sum0; + outptr1[0] = sum1; + outptr0++; + outptr1++; + if (pC0) + { + pC0++; + pC1++; + } + if (pC && broadcast_type_C == 4) + pC++; + } + outptr += out_hstep * 2; + } + for (; ii < max_ii; ii++) + { + float* outptr0 = outptr; + const float* pC = C; + if (pC) + { + if (broadcast_type_C == 1 || broadcast_type_C == 2) + pC += i + ii; + if (broadcast_type_C == 3) + pC += (size_t)(i + ii) * c_hstep + j; + if (broadcast_type_C == 4) + pC += j; + } + const float* pC0 = pC && (broadcast_type_C == 3 || broadcast_type_C == 4) ? pC : 0; + + float c0 = 0.f; + if (pC && (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2)) + { + c0 = pC[0]; + if (beta != 1.f) + c0 *= beta; + } + + int jj = 0; +#if __mips_msa + for (; jj + 7 < max_jj; jj += 8) + { + v4f32 _f0 = (v4f32)__msa_ld_w(pp, 0); + v4f32 _f1 = (v4f32)__msa_ld_w(pp + 4, 0); + pp += 8; + + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + v4f32 _c0 = __msa_fill_w_f32(c0); + _f0 = __msa_fadd_w(_f0, _c0); + _f1 = __msa_fadd_w(_f1, _c0); + } + if (broadcast_type_C == 3 || broadcast_type_C == 4) + { + v4f32 _c0 = (v4f32)__msa_ld_w(pC0, 0); + v4f32 _c1 = (v4f32)__msa_ld_w(pC0 + 4, 0); + if (beta != 1.f) + { + v4f32 _beta = __msa_fill_w_f32(beta); + _c0 = __msa_fmul_w(_c0, _beta); + _c1 = __msa_fmul_w(_c1, _beta); + } + _f0 = __msa_fadd_w(_f0, _c0); + _f1 = __msa_fadd_w(_f1, _c1); + } + } + + if (alpha != 1.f) + { + v4f32 _alpha = __msa_fill_w_f32(alpha); + _f0 = __msa_fmul_w(_f0, _alpha); + _f1 = __msa_fmul_w(_f1, _alpha); + } + + __msa_st_w((v4i32)_f0, outptr0, 0); + __msa_st_w((v4i32)_f1, outptr0 + 4, 0); + outptr0 += 8; + if (pC0) + pC0 += 8; + } + for (; jj + 3 < max_jj; jj += 4) + { + v4f32 _f0 = (v4f32)__msa_ld_w(pp, 0); + pp += 4; + + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + _f0 = __msa_fadd_w(_f0, __msa_fill_w_f32(c0)); + if (broadcast_type_C == 3 || broadcast_type_C == 4) + { + v4f32 _c0 = (v4f32)__msa_ld_w(pC0, 0); + if (beta != 1.f) + _c0 = __msa_fmul_w(_c0, __msa_fill_w_f32(beta)); + _f0 = __msa_fadd_w(_f0, _c0); + } + } + + if (alpha != 1.f) + _f0 = __msa_fmul_w(_f0, __msa_fill_w_f32(alpha)); + + __msa_st_w((v4i32)_f0, outptr0, 0); + outptr0 += 4; + if (pC0) + pC0 += 4; + } + for (; jj + 1 < max_jj; jj += 2) + { + v4i32 _fi = __msa_fill_w(0); + _fi = __msa_insert_w(_fi, 0, ((const int*)pp)[0]); + _fi = __msa_insert_w(_fi, 1, ((const int*)pp)[1]); + v4f32 _f = (v4f32)_fi; + pp += 2; + + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + _f = __msa_fadd_w(_f, __msa_fill_w_f32(c0)); + if (broadcast_type_C == 3 || broadcast_type_C == 4) + { + v4i32 _ci = __msa_fill_w(0); + _ci = __msa_insert_w(_ci, 0, ((const int*)pC0)[0]); + _ci = __msa_insert_w(_ci, 1, ((const int*)pC0)[1]); + v4f32 _c = (v4f32)_ci; + if (beta != 1.f) + _c = __msa_fmul_w(_c, __msa_fill_w_f32(beta)); + _f = __msa_fadd_w(_f, _c); + } + } + + if (alpha != 1.f) + _f = __msa_fmul_w(_f, __msa_fill_w_f32(alpha)); + + __msa_storel_d((v4i32)_f, outptr0); + outptr0 += 2; + if (pC0) + pC0 += 2; + } +#endif // __mips_msa + for (; jj + 1 < max_jj; jj += 2) + { + float sum0 = pp[0]; + float sum1 = pp[1]; + pp += 2; + if (pC) + { + if (broadcast_type_C == 0) + { + sum0 += c0; + sum1 += c0; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + sum0 += c0; + sum1 += c0; + } + if (broadcast_type_C == 3) + { + float c0 = pC0[0]; + float c1 = pC0[1]; + if (beta != 1.f) + { + c0 *= beta; + c1 *= beta; + } + sum0 += c0; + sum1 += c1; + } + if (broadcast_type_C == 4) + { + float c0 = pC0[0]; + float c1 = pC0[1]; + if (beta != 1.f) + { + c0 *= beta; + c1 *= beta; + } + sum0 += c0; + sum1 += c1; + } + } + if (alpha != 1.f) + { + sum0 *= alpha; + sum1 *= alpha; + } + outptr0[0] = sum0; + outptr0[1] = sum1; + outptr0 += 2; + if (pC0) + pC0 += 2; + } + for (; jj < max_jj; jj++) + { + float sum0 = *pp++; + if (pC) + { + float c = 0.f; + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) c = c0; + if (broadcast_type_C == 3 || broadcast_type_C == 4) c = pC0[0]; + if ((broadcast_type_C == 3 || broadcast_type_C == 4) && beta != 1.f) c *= beta; + sum0 += c; + } + if (alpha != 1.f) sum0 *= alpha; + outptr0[0] = sum0; + outptr0++; + if (pC0) + pC0++; + } + outptr += out_hstep; + } +} + +static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, float alpha, float beta) +{ + const size_t out_hstep = top_blob.dims == 3 ? top_blob.cstep : (size_t)top_blob.w; + const size_t c_hstep = C.dims == 3 ? C.cstep : (size_t)C.w; + float* outptr0 = (float*)top_blob + (size_t)j * out_hstep + i; + + int ii = 0; +#if __mips_msa + for (; ii + 7 < max_ii; ii += 8) + { + float* outptr = outptr0; + + const float* pC = C; + if (pC) + { + if (broadcast_type_C == 1 || broadcast_type_C == 2) + pC += i + ii; + if (broadcast_type_C == 3) + pC += (size_t)(i + ii) * c_hstep + j; + if (broadcast_type_C == 4) + pC += j; + } + + v4f32 _c0123 = (v4f32)__msa_fill_w(0); + v4f32 _c4567 = (v4f32)__msa_fill_w(0); + if (pC) + { + if (broadcast_type_C == 0) + { + float c = pC[0]; + if (beta != 1.f) + c *= beta; + _c0123 = __msa_fill_w_f32(c); + _c4567 = _c0123; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _c0123 = (v4f32)__msa_ld_w(pC, 0); + _c4567 = (v4f32)__msa_ld_w(pC + 4, 0); + if (beta != 1.f) + { + const v4f32 _beta = __msa_fill_w_f32(beta); + _c0123 = __msa_fmul_w(_c0123, _beta); + _c4567 = __msa_fmul_w(_c4567, _beta); + } + } + } + + const float* pC0 = pC && broadcast_type_C == 3 ? pC : 0; + const float* pC1 = pC0 ? pC0 + c_hstep : 0; + const float* pC2 = pC0 ? pC0 + c_hstep * 2 : 0; + const float* pC3 = pC0 ? pC0 + c_hstep * 3 : 0; + const float* pC4 = pC0 ? pC0 + c_hstep * 4 : 0; + const float* pC5 = pC0 ? pC0 + c_hstep * 5 : 0; + const float* pC6 = pC0 ? pC0 + c_hstep * 6 : 0; + const float* pC7 = pC0 ? pC0 + c_hstep * 7 : 0; + + int jj = 0; + for (; jj + 3 < max_jj; jj += 4) + { + v4f32 _f0 = (v4f32)__msa_ld_w(pp, 0); + v4f32 _f4 = (v4f32)__msa_ld_w(pp + 4, 0); + v4f32 _f1 = (v4f32)__msa_ld_w(pp + 8, 0); + v4f32 _f5 = (v4f32)__msa_ld_w(pp + 12, 0); + v4f32 _f2 = (v4f32)__msa_ld_w(pp + 16, 0); + v4f32 _f6 = (v4f32)__msa_ld_w(pp + 20, 0); + v4f32 _f3 = (v4f32)__msa_ld_w(pp + 24, 0); + v4f32 _f7 = (v4f32)__msa_ld_w(pp + 28, 0); + pp += 32; + + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = __msa_fadd_w(_f0, _c0123); + _f1 = __msa_fadd_w(_f1, _c0123); + _f2 = __msa_fadd_w(_f2, _c0123); + _f3 = __msa_fadd_w(_f3, _c0123); + _f4 = __msa_fadd_w(_f4, _c4567); + _f5 = __msa_fadd_w(_f5, _c4567); + _f6 = __msa_fadd_w(_f6, _c4567); + _f7 = __msa_fadd_w(_f7, _c4567); + } + if (broadcast_type_C == 3) + { + v4f32 _cl0 = (v4f32){pC0[0], pC1[0], pC2[0], pC3[0]}; + v4f32 _ch0 = (v4f32){pC4[0], pC5[0], pC6[0], pC7[0]}; + v4f32 _cl1 = (v4f32){pC0[1], pC1[1], pC2[1], pC3[1]}; + v4f32 _ch1 = (v4f32){pC4[1], pC5[1], pC6[1], pC7[1]}; + v4f32 _cl2 = (v4f32){pC0[2], pC1[2], pC2[2], pC3[2]}; + v4f32 _ch2 = (v4f32){pC4[2], pC5[2], pC6[2], pC7[2]}; + v4f32 _cl3 = (v4f32){pC0[3], pC1[3], pC2[3], pC3[3]}; + v4f32 _ch3 = (v4f32){pC4[3], pC5[3], pC6[3], pC7[3]}; + if (beta != 1.f) + { + const v4f32 _beta = __msa_fill_w_f32(beta); + _cl0 = __msa_fmul_w(_cl0, _beta); + _ch0 = __msa_fmul_w(_ch0, _beta); + _cl1 = __msa_fmul_w(_cl1, _beta); + _ch1 = __msa_fmul_w(_ch1, _beta); + _cl2 = __msa_fmul_w(_cl2, _beta); + _ch2 = __msa_fmul_w(_ch2, _beta); + _cl3 = __msa_fmul_w(_cl3, _beta); + _ch3 = __msa_fmul_w(_ch3, _beta); + } + _f0 = __msa_fadd_w(_f0, _cl0); + _f4 = __msa_fadd_w(_f4, _ch0); + _f1 = __msa_fadd_w(_f1, _cl1); + _f5 = __msa_fadd_w(_f5, _ch1); + _f2 = __msa_fadd_w(_f2, _cl2); + _f6 = __msa_fadd_w(_f6, _ch2); + _f3 = __msa_fadd_w(_f3, _cl3); + _f7 = __msa_fadd_w(_f7, _ch3); + } + if (broadcast_type_C == 4) + { + v4f32 _c = (v4f32)__msa_ld_w(pC, 0); + if (beta != 1.f) + _c = __msa_fmul_w(_c, __msa_fill_w_f32(beta)); + _f0 = __msa_fadd_w(_f0, (v4f32)__msa_splati_w((v4i32)_c, 0)); + _f4 = __msa_fadd_w(_f4, (v4f32)__msa_splati_w((v4i32)_c, 0)); + _f1 = __msa_fadd_w(_f1, (v4f32)__msa_splati_w((v4i32)_c, 1)); + _f5 = __msa_fadd_w(_f5, (v4f32)__msa_splati_w((v4i32)_c, 1)); + _f2 = __msa_fadd_w(_f2, (v4f32)__msa_splati_w((v4i32)_c, 2)); + _f6 = __msa_fadd_w(_f6, (v4f32)__msa_splati_w((v4i32)_c, 2)); + _f3 = __msa_fadd_w(_f3, (v4f32)__msa_splati_w((v4i32)_c, 3)); + _f7 = __msa_fadd_w(_f7, (v4f32)__msa_splati_w((v4i32)_c, 3)); + } + } + + if (alpha != 1.f) + { + const v4f32 _alpha = __msa_fill_w_f32(alpha); + _f0 = __msa_fmul_w(_f0, _alpha); + _f1 = __msa_fmul_w(_f1, _alpha); + _f2 = __msa_fmul_w(_f2, _alpha); + _f3 = __msa_fmul_w(_f3, _alpha); + _f4 = __msa_fmul_w(_f4, _alpha); + _f5 = __msa_fmul_w(_f5, _alpha); + _f6 = __msa_fmul_w(_f6, _alpha); + _f7 = __msa_fmul_w(_f7, _alpha); + } + __msa_st_w((v4i32)_f0, outptr, 0); + __msa_st_w((v4i32)_f4, outptr + 4, 0); + __msa_st_w((v4i32)_f1, outptr + out_hstep, 0); + __msa_st_w((v4i32)_f5, outptr + out_hstep + 4, 0); + __msa_st_w((v4i32)_f2, outptr + out_hstep * 2, 0); + __msa_st_w((v4i32)_f6, outptr + out_hstep * 2 + 4, 0); + __msa_st_w((v4i32)_f3, outptr + out_hstep * 3, 0); + __msa_st_w((v4i32)_f7, outptr + out_hstep * 3 + 4, 0); + outptr += out_hstep * 4; + if (pC0) + { + pC0 += 4; + pC1 += 4; + pC2 += 4; + pC3 += 4; + pC4 += 4; + pC5 += 4; + pC6 += 4; + pC7 += 4; + } + if (pC && broadcast_type_C == 4) + pC += 4; + } + for (; jj + 1 < max_jj; jj += 2) + { + v4f32 _f0 = (v4f32)__msa_ld_w(pp, 0); + v4f32 _f4 = (v4f32)__msa_ld_w(pp + 4, 0); + v4f32 _f1 = (v4f32)__msa_ld_w(pp + 8, 0); + v4f32 _f5 = (v4f32)__msa_ld_w(pp + 12, 0); + pp += 16; + + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = __msa_fadd_w(_f0, _c0123); + _f4 = __msa_fadd_w(_f4, _c4567); + _f1 = __msa_fadd_w(_f1, _c0123); + _f5 = __msa_fadd_w(_f5, _c4567); + } + if (broadcast_type_C == 3) + { + v4f32 _c0 = (v4f32){pC0[0], pC1[0], pC2[0], pC3[0]}; + v4f32 _c4 = (v4f32){pC4[0], pC5[0], pC6[0], pC7[0]}; + v4f32 _c1 = (v4f32){pC0[1], pC1[1], pC2[1], pC3[1]}; + v4f32 _c5 = (v4f32){pC4[1], pC5[1], pC6[1], pC7[1]}; + if (beta != 1.f) + { + const v4f32 _beta = __msa_fill_w_f32(beta); + _c0 = __msa_fmul_w(_c0, _beta); + _c4 = __msa_fmul_w(_c4, _beta); + _c1 = __msa_fmul_w(_c1, _beta); + _c5 = __msa_fmul_w(_c5, _beta); + } + _f0 = __msa_fadd_w(_f0, _c0); + _f4 = __msa_fadd_w(_f4, _c4); + _f1 = __msa_fadd_w(_f1, _c1); + _f5 = __msa_fadd_w(_f5, _c5); + } + if (broadcast_type_C == 4) + { + float c0 = pC[0]; + float c1 = pC[1]; + if (beta != 1.f) + { + c0 *= beta; + c1 *= beta; + } + _f0 = __msa_fadd_w(_f0, __msa_fill_w_f32(c0)); + _f4 = __msa_fadd_w(_f4, __msa_fill_w_f32(c0)); + _f1 = __msa_fadd_w(_f1, __msa_fill_w_f32(c1)); + _f5 = __msa_fadd_w(_f5, __msa_fill_w_f32(c1)); + } + } + + if (alpha != 1.f) + { + const v4f32 _alpha = __msa_fill_w_f32(alpha); + _f0 = __msa_fmul_w(_f0, _alpha); + _f4 = __msa_fmul_w(_f4, _alpha); + _f1 = __msa_fmul_w(_f1, _alpha); + _f5 = __msa_fmul_w(_f5, _alpha); + } + __msa_st_w((v4i32)_f0, outptr, 0); + __msa_st_w((v4i32)_f4, outptr + 4, 0); + __msa_st_w((v4i32)_f1, outptr + out_hstep, 0); + __msa_st_w((v4i32)_f5, outptr + out_hstep + 4, 0); + outptr += out_hstep * 2; + if (pC0) + { + pC0 += 2; + pC1 += 2; + pC2 += 2; + pC3 += 2; + pC4 += 2; + pC5 += 2; + pC6 += 2; + pC7 += 2; + } + if (pC && broadcast_type_C == 4) + pC += 2; + } + for (; jj < max_jj; jj++) + { + v4f32 _f0 = (v4f32)__msa_ld_w(pp, 0); + v4f32 _f4 = (v4f32)__msa_ld_w(pp + 4, 0); + pp += 8; + + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = __msa_fadd_w(_f0, _c0123); + _f4 = __msa_fadd_w(_f4, _c4567); + } + if (broadcast_type_C == 3) + { + v4f32 _c0 = (v4f32){pC0[0], pC1[0], pC2[0], pC3[0]}; + v4f32 _c4 = (v4f32){pC4[0], pC5[0], pC6[0], pC7[0]}; + if (beta != 1.f) + { + const v4f32 _beta = __msa_fill_w_f32(beta); + _c0 = __msa_fmul_w(_c0, _beta); + _c4 = __msa_fmul_w(_c4, _beta); + } + _f0 = __msa_fadd_w(_f0, _c0); + _f4 = __msa_fadd_w(_f4, _c4); + } + if (broadcast_type_C == 4) + { + float c = pC[0]; + if (beta != 1.f) + c *= beta; + _f0 = __msa_fadd_w(_f0, __msa_fill_w_f32(c)); + _f4 = __msa_fadd_w(_f4, __msa_fill_w_f32(c)); + } + } + + if (alpha != 1.f) + { + const v4f32 _alpha = __msa_fill_w_f32(alpha); + _f0 = __msa_fmul_w(_f0, _alpha); + _f4 = __msa_fmul_w(_f4, _alpha); + } + __msa_st_w((v4i32)_f0, outptr, 0); + __msa_st_w((v4i32)_f4, outptr + 4, 0); + outptr += out_hstep; + } + outptr0 += 8; + } + for (; ii + 3 < max_ii; ii += 4) + { + float* outptr = outptr0; + + const float* pC = C; + if (pC) + { + if (broadcast_type_C == 1 || broadcast_type_C == 2) + pC += i + ii; + if (broadcast_type_C == 3) + pC += (size_t)(i + ii) * c_hstep + j; + if (broadcast_type_C == 4) + pC += j; + } + + v4f32 _c0123 = (v4f32)__msa_fill_w(0); + if (pC) + { + if (broadcast_type_C == 0) + { + float c = pC[0]; + if (beta != 1.f) + c *= beta; + _c0123 = __msa_fill_w_f32(c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _c0123 = (v4f32)__msa_ld_w(pC, 0); + if (beta != 1.f) + _c0123 = __msa_fmul_w(_c0123, __msa_fill_w_f32(beta)); + } + } + + const float* pC0 = pC && broadcast_type_C == 3 ? pC : 0; + const float* pC1 = pC0 ? pC0 + c_hstep : 0; + const float* pC2 = pC0 ? pC0 + c_hstep * 2 : 0; + const float* pC3 = pC0 ? pC0 + c_hstep * 3 : 0; + + int jj = 0; + for (; jj + 7 < max_jj; jj += 8) + { + v4f32 _f0 = (v4f32)__msa_ld_w(pp, 0); + v4f32 _f1 = (v4f32)__msa_ld_w(pp + 4, 0); + v4f32 _f2 = (v4f32)__msa_ld_w(pp + 8, 0); + v4f32 _f3 = (v4f32)__msa_ld_w(pp + 12, 0); + v4f32 _f4 = (v4f32)__msa_ld_w(pp + 16, 0); + v4f32 _f5 = (v4f32)__msa_ld_w(pp + 20, 0); + v4f32 _f6 = (v4f32)__msa_ld_w(pp + 24, 0); + v4f32 _f7 = (v4f32)__msa_ld_w(pp + 28, 0); + pp += 32; + + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = __msa_fadd_w(_f0, _c0123); + _f1 = __msa_fadd_w(_f1, _c0123); + _f2 = __msa_fadd_w(_f2, _c0123); + _f3 = __msa_fadd_w(_f3, _c0123); + _f4 = __msa_fadd_w(_f4, _c0123); + _f5 = __msa_fadd_w(_f5, _c0123); + _f6 = __msa_fadd_w(_f6, _c0123); + _f7 = __msa_fadd_w(_f7, _c0123); + } + if (broadcast_type_C == 3) + { + v4f32 _c0 = (v4f32)__msa_ld_w(pC0, 0); + v4f32 _c1 = (v4f32)__msa_ld_w(pC1, 0); + v4f32 _c2 = (v4f32)__msa_ld_w(pC2, 0); + v4f32 _c3 = (v4f32)__msa_ld_w(pC3, 0); + transpose4x4_ps(_c0, _c1, _c2, _c3); + v4f32 _c4 = (v4f32)__msa_ld_w(pC0 + 4, 0); + v4f32 _c5 = (v4f32)__msa_ld_w(pC1 + 4, 0); + v4f32 _c6 = (v4f32)__msa_ld_w(pC2 + 4, 0); + v4f32 _c7 = (v4f32)__msa_ld_w(pC3 + 4, 0); + transpose4x4_ps(_c4, _c5, _c6, _c7); + if (beta != 1.f) + { + v4f32 _beta = __msa_fill_w_f32(beta); + _c0 = __msa_fmul_w(_c0, _beta); + _c1 = __msa_fmul_w(_c1, _beta); + _c2 = __msa_fmul_w(_c2, _beta); + _c3 = __msa_fmul_w(_c3, _beta); + _c4 = __msa_fmul_w(_c4, _beta); + _c5 = __msa_fmul_w(_c5, _beta); + _c6 = __msa_fmul_w(_c6, _beta); + _c7 = __msa_fmul_w(_c7, _beta); + } + _f0 = __msa_fadd_w(_f0, _c0); + _f1 = __msa_fadd_w(_f1, _c1); + _f2 = __msa_fadd_w(_f2, _c2); + _f3 = __msa_fadd_w(_f3, _c3); + _f4 = __msa_fadd_w(_f4, _c4); + _f5 = __msa_fadd_w(_f5, _c5); + _f6 = __msa_fadd_w(_f6, _c6); + _f7 = __msa_fadd_w(_f7, _c7); + } + if (broadcast_type_C == 4) + { + v4f32 _c0 = (v4f32)__msa_ld_w(pC, 0); + v4f32 _c4 = (v4f32)__msa_ld_w(pC + 4, 0); + if (beta != 1.f) + { + v4f32 _beta = __msa_fill_w_f32(beta); + _c0 = __msa_fmul_w(_c0, _beta); + _c4 = __msa_fmul_w(_c4, _beta); + } + _f0 = __msa_fadd_w(_f0, (v4f32)__msa_splati_w((v4i32)_c0, 0)); + _f1 = __msa_fadd_w(_f1, (v4f32)__msa_splati_w((v4i32)_c0, 1)); + _f2 = __msa_fadd_w(_f2, (v4f32)__msa_splati_w((v4i32)_c0, 2)); + _f3 = __msa_fadd_w(_f3, (v4f32)__msa_splati_w((v4i32)_c0, 3)); + _f4 = __msa_fadd_w(_f4, (v4f32)__msa_splati_w((v4i32)_c4, 0)); + _f5 = __msa_fadd_w(_f5, (v4f32)__msa_splati_w((v4i32)_c4, 1)); + _f6 = __msa_fadd_w(_f6, (v4f32)__msa_splati_w((v4i32)_c4, 2)); + _f7 = __msa_fadd_w(_f7, (v4f32)__msa_splati_w((v4i32)_c4, 3)); + } + } + + if (alpha != 1.f) + { + v4f32 _alpha = __msa_fill_w_f32(alpha); + _f0 = __msa_fmul_w(_f0, _alpha); + _f1 = __msa_fmul_w(_f1, _alpha); + _f2 = __msa_fmul_w(_f2, _alpha); + _f3 = __msa_fmul_w(_f3, _alpha); + _f4 = __msa_fmul_w(_f4, _alpha); + _f5 = __msa_fmul_w(_f5, _alpha); + _f6 = __msa_fmul_w(_f6, _alpha); + _f7 = __msa_fmul_w(_f7, _alpha); + } + + __msa_st_w((v4i32)_f0, outptr, 0); + __msa_st_w((v4i32)_f1, outptr + out_hstep, 0); + __msa_st_w((v4i32)_f2, outptr + out_hstep * 2, 0); + __msa_st_w((v4i32)_f3, outptr + out_hstep * 3, 0); + __msa_st_w((v4i32)_f4, outptr + out_hstep * 4, 0); + __msa_st_w((v4i32)_f5, outptr + out_hstep * 5, 0); + __msa_st_w((v4i32)_f6, outptr + out_hstep * 6, 0); + __msa_st_w((v4i32)_f7, outptr + out_hstep * 7, 0); + outptr += out_hstep * 8; + if (pC0) + { + pC0 += 8; + pC1 += 8; + pC2 += 8; + pC3 += 8; + } + if (pC && broadcast_type_C == 4) + pC += 8; + } + for (; jj + 3 < max_jj; jj += 4) + { + v4f32 _f0 = (v4f32)__msa_ld_w(pp, 0); + v4f32 _f1 = (v4f32)__msa_ld_w(pp + 4, 0); + v4f32 _f2 = (v4f32)__msa_ld_w(pp + 8, 0); + v4f32 _f3 = (v4f32)__msa_ld_w(pp + 12, 0); + pp += 16; + + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = __msa_fadd_w(_f0, _c0123); + _f1 = __msa_fadd_w(_f1, _c0123); + _f2 = __msa_fadd_w(_f2, _c0123); + _f3 = __msa_fadd_w(_f3, _c0123); + } + if (broadcast_type_C == 3) + { + v4f32 _c0 = (v4f32)__msa_ld_w(pC0, 0); + v4f32 _c1 = (v4f32)__msa_ld_w(pC1, 0); + v4f32 _c2 = (v4f32)__msa_ld_w(pC2, 0); + v4f32 _c3 = (v4f32)__msa_ld_w(pC3, 0); + transpose4x4_ps(_c0, _c1, _c2, _c3); + if (beta != 1.f) + { + v4f32 _beta = __msa_fill_w_f32(beta); + _c0 = __msa_fmul_w(_c0, _beta); + _c1 = __msa_fmul_w(_c1, _beta); + _c2 = __msa_fmul_w(_c2, _beta); + _c3 = __msa_fmul_w(_c3, _beta); + } + _f0 = __msa_fadd_w(_f0, _c0); + _f1 = __msa_fadd_w(_f1, _c1); + _f2 = __msa_fadd_w(_f2, _c2); + _f3 = __msa_fadd_w(_f3, _c3); + } + if (broadcast_type_C == 4) + { + v4f32 _c0 = (v4f32)__msa_ld_w(pC, 0); + if (beta != 1.f) + _c0 = __msa_fmul_w(_c0, __msa_fill_w_f32(beta)); + _f0 = __msa_fadd_w(_f0, (v4f32)__msa_splati_w((v4i32)_c0, 0)); + _f1 = __msa_fadd_w(_f1, (v4f32)__msa_splati_w((v4i32)_c0, 1)); + _f2 = __msa_fadd_w(_f2, (v4f32)__msa_splati_w((v4i32)_c0, 2)); + _f3 = __msa_fadd_w(_f3, (v4f32)__msa_splati_w((v4i32)_c0, 3)); + } + } + + if (alpha != 1.f) + { + v4f32 _alpha = __msa_fill_w_f32(alpha); + _f0 = __msa_fmul_w(_f0, _alpha); + _f1 = __msa_fmul_w(_f1, _alpha); + _f2 = __msa_fmul_w(_f2, _alpha); + _f3 = __msa_fmul_w(_f3, _alpha); + } + + __msa_st_w((v4i32)_f0, outptr, 0); + __msa_st_w((v4i32)_f1, outptr + out_hstep, 0); + __msa_st_w((v4i32)_f2, outptr + out_hstep * 2, 0); + __msa_st_w((v4i32)_f3, outptr + out_hstep * 3, 0); + outptr += out_hstep * 4; + if (pC0) + { + pC0 += 4; + pC1 += 4; + pC2 += 4; + pC3 += 4; + } + if (pC && broadcast_type_C == 4) + pC += 4; + } + for (; jj + 1 < max_jj; jj += 2) + { + v4f32 _f0 = (v4f32)__msa_ld_w(pp, 0); + v4f32 _f1 = (v4f32)__msa_ld_w(pp + 4, 0); + pp += 8; + + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = __msa_fadd_w(_f0, _c0123); + _f1 = __msa_fadd_w(_f1, _c0123); + } + if (broadcast_type_C == 3) + { + v4f32 _c0 = (v4f32){pC0[0], pC1[0], pC2[0], pC3[0]}; + v4f32 _c1 = (v4f32){pC0[1], pC1[1], pC2[1], pC3[1]}; + if (beta != 1.f) + { + v4f32 _beta = __msa_fill_w_f32(beta); + _c0 = __msa_fmul_w(_c0, _beta); + _c1 = __msa_fmul_w(_c1, _beta); + } + _f0 = __msa_fadd_w(_f0, _c0); + _f1 = __msa_fadd_w(_f1, _c1); + } + if (broadcast_type_C == 4) + { + float c0 = pC[0]; + float c1 = pC[1]; + if (beta != 1.f) + { + c0 *= beta; + c1 *= beta; + } + _f0 = __msa_fadd_w(_f0, __msa_fill_w_f32(c0)); + _f1 = __msa_fadd_w(_f1, __msa_fill_w_f32(c1)); + } + } + + if (alpha != 1.f) + { + v4f32 _alpha = __msa_fill_w_f32(alpha); + _f0 = __msa_fmul_w(_f0, _alpha); + _f1 = __msa_fmul_w(_f1, _alpha); + } + + __msa_st_w((v4i32)_f0, outptr, 0); + __msa_st_w((v4i32)_f1, outptr + out_hstep, 0); + outptr += out_hstep * 2; + if (pC0) + { + pC0 += 2; + pC1 += 2; + pC2 += 2; + pC3 += 2; + } + if (pC && broadcast_type_C == 4) + pC += 2; + } + for (; jj < max_jj; jj++) + { + v4f32 _f0 = (v4f32)__msa_ld_w(pp, 0); + pp += 4; + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = __msa_fadd_w(_f0, _c0123); + } + if (broadcast_type_C == 3) + { + v4f32 _c0 = (v4f32){pC0[0], pC1[0], pC2[0], pC3[0]}; + if (beta != 1.f) + _c0 = __msa_fmul_w(_c0, __msa_fill_w_f32(beta)); + _f0 = __msa_fadd_w(_f0, _c0); + } + if (broadcast_type_C == 4) + { + float c = pC[0]; + if (beta != 1.f) + c *= beta; + _f0 = __msa_fadd_w(_f0, __msa_fill_w_f32(c)); + } + } + if (alpha != 1.f) + _f0 = __msa_fmul_w(_f0, __msa_fill_w_f32(alpha)); + __msa_st_w((v4i32)_f0, outptr, 0); + outptr += out_hstep; + if (pC0) + { + pC0++; + pC1++; + pC2++; + pC3++; + } + if (pC && broadcast_type_C == 4) + pC++; + } + outptr0 += 4; + } +#endif // __mips_msa + for (; ii + 1 < max_ii; ii += 2) + { + float* outptr = outptr0; + const float* pC = C; + if (pC) + { + if (broadcast_type_C == 1 || broadcast_type_C == 2) + pC += i + ii; + if (broadcast_type_C == 3) + pC += (size_t)(i + ii) * c_hstep + j; + if (broadcast_type_C == 4) + pC += j; + } + const float* pC0 = pC && broadcast_type_C == 3 ? pC : 0; + const float* pC1 = pC0 ? pC0 + c_hstep : 0; + + float c0 = 0.f; + float c1 = 0.f; + if (pC && (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2)) + { + c0 = pC[0]; + c1 = pC[broadcast_type_C == 0 ? 0 : 1]; + if (beta != 1.f) + { + c0 *= beta; + c1 *= beta; + } + } + + int jj = 0; +#if __mips_msa + for (; jj + 7 < max_jj; jj += 8) + { + v4f32 _f0 = (v4f32)__msa_ld_w(pp, 0); + v4f32 _f1 = (v4f32)__msa_ld_w(pp + 4, 0); + v4f32 _f2 = (v4f32)__msa_ld_w(pp + 8, 0); + v4f32 _f3 = (v4f32)__msa_ld_w(pp + 12, 0); + pp += 16; + + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + v4f32 _c = (v4f32){c0, c1, c0, c1}; + _f0 = __msa_fadd_w(_f0, _c); + _f1 = __msa_fadd_w(_f1, _c); + _f2 = __msa_fadd_w(_f2, _c); + _f3 = __msa_fadd_w(_f3, _c); + } + if (broadcast_type_C == 3) + { + v4f32 _c0 = (v4f32){pC0[0], pC1[0], pC0[1], pC1[1]}; + v4f32 _c1 = (v4f32){pC0[2], pC1[2], pC0[3], pC1[3]}; + v4f32 _c2 = (v4f32){pC0[4], pC1[4], pC0[5], pC1[5]}; + v4f32 _c3 = (v4f32){pC0[6], pC1[6], pC0[7], pC1[7]}; + if (beta != 1.f) + { + v4f32 _beta = __msa_fill_w_f32(beta); + _c0 = __msa_fmul_w(_c0, _beta); + _c1 = __msa_fmul_w(_c1, _beta); + _c2 = __msa_fmul_w(_c2, _beta); + _c3 = __msa_fmul_w(_c3, _beta); + } + _f0 = __msa_fadd_w(_f0, _c0); + _f1 = __msa_fadd_w(_f1, _c1); + _f2 = __msa_fadd_w(_f2, _c2); + _f3 = __msa_fadd_w(_f3, _c3); + } + if (broadcast_type_C == 4) + { + float c00 = pC[0]; + float c01 = pC[1]; + float c02 = pC[2]; + float c03 = pC[3]; + float c04 = pC[4]; + float c05 = pC[5]; + float c06 = pC[6]; + float c07 = pC[7]; + if (beta != 1.f) + { + c00 *= beta; + c01 *= beta; + c02 *= beta; + c03 *= beta; + c04 *= beta; + c05 *= beta; + c06 *= beta; + c07 *= beta; + } + _f0 = __msa_fadd_w(_f0, (v4f32){c00, c00, c01, c01}); + _f1 = __msa_fadd_w(_f1, (v4f32){c02, c02, c03, c03}); + _f2 = __msa_fadd_w(_f2, (v4f32){c04, c04, c05, c05}); + _f3 = __msa_fadd_w(_f3, (v4f32){c06, c06, c07, c07}); + } + } + + if (alpha != 1.f) + { + v4f32 _alpha = __msa_fill_w_f32(alpha); + _f0 = __msa_fmul_w(_f0, _alpha); + _f1 = __msa_fmul_w(_f1, _alpha); + _f2 = __msa_fmul_w(_f2, _alpha); + _f3 = __msa_fmul_w(_f3, _alpha); + } + + __msa_storel_d((v4i32)_f0, outptr); + __msa_storel_d((v4i32)__msa_sldi_b((v16i8)_f0, (v16i8)_f0, 8), outptr + out_hstep); + __msa_storel_d((v4i32)_f1, outptr + out_hstep * 2); + __msa_storel_d((v4i32)__msa_sldi_b((v16i8)_f1, (v16i8)_f1, 8), outptr + out_hstep * 3); + __msa_storel_d((v4i32)_f2, outptr + out_hstep * 4); + __msa_storel_d((v4i32)__msa_sldi_b((v16i8)_f2, (v16i8)_f2, 8), outptr + out_hstep * 5); + __msa_storel_d((v4i32)_f3, outptr + out_hstep * 6); + __msa_storel_d((v4i32)__msa_sldi_b((v16i8)_f3, (v16i8)_f3, 8), outptr + out_hstep * 7); + outptr += out_hstep * 8; + if (pC0) + { + pC0 += 8; + pC1 += 8; + } + if (pC && broadcast_type_C == 4) + pC += 8; + } + for (; jj + 3 < max_jj; jj += 4) + { + v4f32 _f0 = (v4f32)__msa_ld_w(pp, 0); + v4f32 _f1 = (v4f32)__msa_ld_w(pp + 4, 0); + pp += 8; + + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + v4f32 _c = (v4f32){c0, c1, c0, c1}; + _f0 = __msa_fadd_w(_f0, _c); + _f1 = __msa_fadd_w(_f1, _c); + } + if (broadcast_type_C == 3) + { + v4f32 _c0 = (v4f32){pC0[0], pC1[0], pC0[1], pC1[1]}; + v4f32 _c1 = (v4f32){pC0[2], pC1[2], pC0[3], pC1[3]}; + if (beta != 1.f) + { + v4f32 _beta = __msa_fill_w_f32(beta); + _c0 = __msa_fmul_w(_c0, _beta); + _c1 = __msa_fmul_w(_c1, _beta); + } + _f0 = __msa_fadd_w(_f0, _c0); + _f1 = __msa_fadd_w(_f1, _c1); + } + if (broadcast_type_C == 4) + { + float c00 = pC[0]; + float c01 = pC[1]; + float c02 = pC[2]; + float c03 = pC[3]; + if (beta != 1.f) + { + c00 *= beta; + c01 *= beta; + c02 *= beta; + c03 *= beta; + } + _f0 = __msa_fadd_w(_f0, (v4f32){c00, c00, c01, c01}); + _f1 = __msa_fadd_w(_f1, (v4f32){c02, c02, c03, c03}); + } + } + + if (alpha != 1.f) + { + v4f32 _alpha = __msa_fill_w_f32(alpha); + _f0 = __msa_fmul_w(_f0, _alpha); + _f1 = __msa_fmul_w(_f1, _alpha); + } + + __msa_storel_d((v4i32)_f0, outptr); + __msa_storel_d((v4i32)__msa_sldi_b((v16i8)_f0, (v16i8)_f0, 8), outptr + out_hstep); + __msa_storel_d((v4i32)_f1, outptr + out_hstep * 2); + __msa_storel_d((v4i32)__msa_sldi_b((v16i8)_f1, (v16i8)_f1, 8), outptr + out_hstep * 3); + outptr += out_hstep * 4; + if (pC0) + { + pC0 += 4; + pC1 += 4; + } + if (pC && broadcast_type_C == 4) + pC += 4; + } + for (; jj + 1 < max_jj; jj += 2) + { + v4f32 _f = (v4f32)__msa_ld_w(pp, 0); + pp += 4; + + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + _f = __msa_fadd_w(_f, (v4f32){c0, c1, c0, c1}); + if (broadcast_type_C == 3) + { + v4f32 _c = (v4f32){pC0[0], pC1[0], pC0[1], pC1[1]}; + if (beta != 1.f) + _c = __msa_fmul_w(_c, __msa_fill_w_f32(beta)); + _f = __msa_fadd_w(_f, _c); + } + if (broadcast_type_C == 4) + { + float cc0 = pC[0]; + float cc1 = pC[1]; + if (beta != 1.f) + { + cc0 *= beta; + cc1 *= beta; + } + _f = __msa_fadd_w(_f, (v4f32){cc0, cc0, cc1, cc1}); + } + } + + if (alpha != 1.f) + _f = __msa_fmul_w(_f, __msa_fill_w_f32(alpha)); + + __msa_storel_d((v4i32)_f, outptr); + __msa_storel_d((v4i32)__msa_sldi_b((v16i8)_f, (v16i8)_f, 8), outptr + out_hstep); + outptr += out_hstep * 2; + if (pC0) + { + pC0 += 2; + pC1 += 2; + } + if (pC && broadcast_type_C == 4) + pC += 2; + } +#endif // __mips_msa + for (; jj + 1 < max_jj; jj += 2) + { + float sum00 = pp[0]; + float sum01 = pp[1]; + float sum10 = pp[2]; + float sum11 = pp[3]; + pp += 4; + if (pC) + { + if (broadcast_type_C == 0) + { + sum00 += c0; + sum01 += c1; + sum10 += c0; + sum11 += c1; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + sum00 += c0; + sum10 += c0; + sum01 += c1; + sum11 += c1; + } + if (broadcast_type_C == 3) + { + float c00 = pC0[0]; + float c01 = pC1[0]; + float c10 = pC0[1]; + float c11 = pC1[1]; + if (beta != 1.f) + { + c00 *= beta; + c01 *= beta; + c10 *= beta; + c11 *= beta; + } + sum00 += c00; + sum01 += c01; + sum10 += c10; + sum11 += c11; + } + if (broadcast_type_C == 4) + { + float c0 = pC[0]; + float c1 = pC[1]; + if (beta != 1.f) + { + c0 *= beta; + c1 *= beta; + } + sum00 += c0; + sum01 += c0; + sum10 += c1; + sum11 += c1; + } + } + if (alpha != 1.f) + { + sum00 *= alpha; + sum01 *= alpha; + sum10 *= alpha; + sum11 *= alpha; + } + outptr[0] = sum00; + outptr[1] = sum01; + outptr[out_hstep] = sum10; + outptr[out_hstep + 1] = sum11; + outptr += out_hstep * 2; + if (pC0) + { + pC0 += 2; + pC1 += 2; + } + if (pC && broadcast_type_C == 4) + pC += 2; + } + for (; jj < max_jj; jj++) + { + float sum0 = pp[0]; + float sum1 = pp[1]; + pp += 2; + if (pC) + { + if (broadcast_type_C == 0) + { + sum0 += c0; + sum1 += c1; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + sum0 += c0; + sum1 += c1; + } + if (broadcast_type_C == 3) + { + float c0 = pC0[0]; + float c1 = pC1[0]; + if (beta != 1.f) + { + c0 *= beta; + c1 *= beta; + } + sum0 += c0; + sum1 += c1; + } + if (broadcast_type_C == 4) + { + float c = pC[0]; + if (beta != 1.f) c *= beta; + sum0 += c; + sum1 += c; + } + } + if (alpha != 1.f) + { + sum0 *= alpha; + sum1 *= alpha; + } + outptr[0] = sum0; + outptr[1] = sum1; + outptr += out_hstep; + if (pC0) + { + pC0++; + pC1++; + } + if (pC && broadcast_type_C == 4) + pC++; + } + outptr0 += 2; + } + for (; ii < max_ii; ii++) + { + float* outptr = outptr0; + const float* pC = C; + if (pC) + { + if (broadcast_type_C == 1 || broadcast_type_C == 2) + pC += i + ii; + if (broadcast_type_C == 3) + pC += (size_t)(i + ii) * c_hstep + j; + if (broadcast_type_C == 4) + pC += j; + } + const float* pC0 = pC && (broadcast_type_C == 3 || broadcast_type_C == 4) ? pC : 0; + float c0 = 0.f; + if (pC && (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2)) + { + c0 = pC[0]; + if (beta != 1.f) + c0 *= beta; + } + int jj = 0; +#if __mips_msa + v4f32 _c0 = __msa_fill_w_f32(c0); + + for (; jj + 7 < max_jj; jj += 8) + { + v4f32 _f0 = (v4f32)__msa_ld_w(pp, 0); + v4f32 _f1 = (v4f32)__msa_ld_w(pp + 4, 0); + pp += 8; + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = __msa_fadd_w(_f0, _c0); + _f1 = __msa_fadd_w(_f1, _c0); + } + if (broadcast_type_C == 3 || broadcast_type_C == 4) + { + v4f32 _c0 = (v4f32)__msa_ld_w(pC0, 0); + v4f32 _c1 = (v4f32)__msa_ld_w(pC0 + 4, 0); + if (beta != 1.f) + { + v4f32 _beta = __msa_fill_w_f32(beta); + _c0 = __msa_fmul_w(_c0, _beta); + _c1 = __msa_fmul_w(_c1, _beta); + } + _f0 = __msa_fadd_w(_f0, _c0); + _f1 = __msa_fadd_w(_f1, _c1); + } + } + if (alpha != 1.f) + { + v4f32 _alpha = __msa_fill_w_f32(alpha); + _f0 = __msa_fmul_w(_f0, _alpha); + _f1 = __msa_fmul_w(_f1, _alpha); + } + if (out_hstep == 1) + { + __msa_st_w((v4i32)_f0, outptr, 0); + __msa_st_w((v4i32)_f1, outptr + 4, 0); + } + else + { + *(int*)outptr = __msa_copy_s_w((v4i32)_f0, 0); + *(int*)(outptr + out_hstep) = __msa_copy_s_w((v4i32)_f0, 1); + *(int*)(outptr + out_hstep * 2) = __msa_copy_s_w((v4i32)_f0, 2); + *(int*)(outptr + out_hstep * 3) = __msa_copy_s_w((v4i32)_f0, 3); + *(int*)(outptr + out_hstep * 4) = __msa_copy_s_w((v4i32)_f1, 0); + *(int*)(outptr + out_hstep * 5) = __msa_copy_s_w((v4i32)_f1, 1); + *(int*)(outptr + out_hstep * 6) = __msa_copy_s_w((v4i32)_f1, 2); + *(int*)(outptr + out_hstep * 7) = __msa_copy_s_w((v4i32)_f1, 3); + } + outptr += out_hstep * 8; + if (pC0) + pC0 += 8; + } + for (; jj + 3 < max_jj; jj += 4) + { + v4f32 _f0 = (v4f32)__msa_ld_w(pp, 0); + pp += 4; + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + _f0 = __msa_fadd_w(_f0, _c0); + if (broadcast_type_C == 3 || broadcast_type_C == 4) + { + v4f32 _c = (v4f32)__msa_ld_w(pC0, 0); + if (beta != 1.f) + _c = __msa_fmul_w(_c, __msa_fill_w_f32(beta)); + _f0 = __msa_fadd_w(_f0, _c); + } + } + if (alpha != 1.f) + _f0 = __msa_fmul_w(_f0, __msa_fill_w_f32(alpha)); + if (out_hstep == 1) + { + __msa_st_w((v4i32)_f0, outptr, 0); + } + else + { + *(int*)outptr = __msa_copy_s_w((v4i32)_f0, 0); + *(int*)(outptr + out_hstep) = __msa_copy_s_w((v4i32)_f0, 1); + *(int*)(outptr + out_hstep * 2) = __msa_copy_s_w((v4i32)_f0, 2); + *(int*)(outptr + out_hstep * 3) = __msa_copy_s_w((v4i32)_f0, 3); + } + outptr += out_hstep * 4; + if (pC0) + pC0 += 4; + } + for (; jj + 1 < max_jj; jj += 2) + { + v4i32 _fi = __msa_fill_w(0); + _fi = __msa_insert_w(_fi, 0, ((const int*)pp)[0]); + _fi = __msa_insert_w(_fi, 1, ((const int*)pp)[1]); + v4f32 _f0 = (v4f32)_fi; + pp += 2; + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + _f0 = __msa_fadd_w(_f0, _c0); + if (broadcast_type_C == 3 || broadcast_type_C == 4) + { + v4i32 _ci = __msa_fill_w(0); + _ci = __msa_insert_w(_ci, 0, ((const int*)pC0)[0]); + _ci = __msa_insert_w(_ci, 1, ((const int*)pC0)[1]); + v4f32 _c = (v4f32)_ci; + if (beta != 1.f) + _c = __msa_fmul_w(_c, __msa_fill_w_f32(beta)); + _f0 = __msa_fadd_w(_f0, _c); + } + } + if (alpha != 1.f) + _f0 = __msa_fmul_w(_f0, __msa_fill_w_f32(alpha)); + if (out_hstep == 1) + { + *(int64_t*)outptr = __msa_copy_s_d((v2i64)_f0, 0); + } + else + { + *(int*)outptr = __msa_copy_s_w((v4i32)_f0, 0); + *(int*)(outptr + out_hstep) = __msa_copy_s_w((v4i32)_f0, 1); + } + outptr += out_hstep * 2; + if (pC0) + pC0 += 2; + } +#endif // __mips_msa + for (; jj + 1 < max_jj; jj += 2) + { + float sum0 = pp[0]; + float sum1 = pp[1]; + pp += 2; + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + sum0 += c0; + sum1 += c0; + } + if (broadcast_type_C == 3 || broadcast_type_C == 4) + { + float c0 = pC0[0]; + float c1 = pC0[1]; + if (beta != 1.f) + { + c0 *= beta; + c1 *= beta; + } + sum0 += c0; + sum1 += c1; + } + } + if (alpha != 1.f) + { + sum0 *= alpha; + sum1 *= alpha; + } + outptr[0] = sum0; + outptr[out_hstep] = sum1; + outptr += out_hstep * 2; + if (pC0) + pC0 += 2; + } + for (; jj < max_jj; jj++) + { + float sum0 = *pp++; + if (pC) + { + float c = 0.f; + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) c = c0; + if (broadcast_type_C == 3 || broadcast_type_C == 4) + { + c = pC0[0]; + if (beta != 1.f) c *= beta; + } + sum0 += c; + } + if (alpha != 1.f) sum0 *= alpha; + outptr[0] = sum0; + outptr += out_hstep; + if (pC0) + pC0++; + } + outptr0++; + } +} + +static void get_optimal_tile_mnk_wq_int8(int M, int N, int K, int constant_TILE_M, int constant_TILE_N, int constant_TILE_K, int& TILE_M, int& TILE_N, int& TILE_K, int nT) +{ + // resolve optimal tile size from cache size + const size_t l2_cache_size = get_cpu_level2_cache_size(); + + const int tile_size = std::max(1, (int)((float)l2_cache_size / 2 / sizeof(signed char) / std::max(1, K))); + +#if __mips_msa + const int tile_m_align = 8; + const int tile_n_align = 8; +#else + const int tile_m_align = 4; + const int tile_n_align = 2; +#endif + // one driver M tile follows the natural producer slab + TILE_M = tile_m_align; + TILE_N = std::max(tile_n_align, tile_size / tile_n_align * tile_n_align); + TILE_K = K; + + if (N > 0) + { + const int nn_N = (N + TILE_N - 1) / TILE_N; + TILE_N = std::min(TILE_N, ((N + nn_N - 1) / nn_N + tile_n_align - 1) / tile_n_align * tile_n_align); + } + + // always take constant TILE_N value when provided + if (constant_TILE_N > 0) + TILE_N = (constant_TILE_N + tile_n_align - 1) / tile_n_align * tile_n_align; + + (void)M; + (void)constant_TILE_M; + (void)constant_TILE_K; + (void)nT; +} diff --git a/src/layer/mips/multiheadattention_mips.cpp b/src/layer/mips/multiheadattention_mips.cpp index 16453933fd6e..f7c4aa348c33 100644 --- a/src/layer/mips/multiheadattention_mips.cpp +++ b/src/layer/mips/multiheadattention_mips.cpp @@ -28,10 +28,355 @@ MultiHeadAttention_mips::MultiHeadAttention_mips() o_gemm = 0; } +#if NCNN_WEIGHT_QUANT +int MultiHeadAttention_mips::create_pipeline_wq_int8(const Option& _opt) +{ + if (q_gemm) + return 0; + + Option opt = _opt; + Option opt_wq = opt; + opt_wq.use_packing_layout = false; + opt_wq.use_fp16_packed = false; + opt_wq.use_fp16_storage = false; + opt_wq.use_fp16_arithmetic = false; + opt_wq.use_bf16_packed = false; + opt_wq.use_bf16_storage = false; + + { + qk_softmax = ncnn::create_layer_cpu(ncnn::LayerType::Softmax); + if (!qk_softmax) + return -100; + ncnn::ParamDict pd; + pd.set(0, -1); + pd.set(1, 1); + int ret = qk_softmax->load_param(pd); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = qk_softmax->load_model(ModelBinFromMatArray(0)); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = qk_softmax->create_pipeline(opt); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + } + + const int qdim = weight_data_size / embed_dim; + + { + q_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!q_gemm) + { + destroy_pipeline(_opt); + return -100; + } + ncnn::ParamDict pd; + pd.set(0, scale); + pd.set(1, 1.f); + pd.set(2, 0); // transA + pd.set(3, 1); // transB + pd.set(4, 0); // constantA + pd.set(5, 1); // constantB + pd.set(6, 1); // constantC + pd.set(7, 0); // M + pd.set(8, embed_dim); // N + pd.set(9, qdim); // K + pd.set(10, 4); // constant_broadcast_type_C + pd.set(11, 0); // output_N1M + pd.set(12, 0); // output_elempack + pd.set(14, 1); // output_transpose + pd.set(18, quantize_term); + int ret = q_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + Mat weights[4]; + weights[0] = q_weight_data; + weights[1] = q_bias_data; + weights[2] = q_weight_data_quantize_scales; + weights[3] = q_weight_data_input_scales; + ret = q_gemm->load_model(ModelBinFromMatArray(weights)); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = q_gemm->create_pipeline(opt_wq); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + } + + { + k_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!k_gemm) + { + destroy_pipeline(_opt); + return -100; + } + ncnn::ParamDict pd; + pd.set(2, 0); // transA + pd.set(3, 1); // transB + pd.set(4, 0); // constantA + pd.set(5, 1); // constantB + pd.set(6, 1); // constantC + pd.set(7, 0); // M + pd.set(8, embed_dim); // N + pd.set(9, kdim); // K + pd.set(10, 4); // constant_broadcast_type_C + pd.set(11, 0); // output_N1M + pd.set(12, 0); // output_elempack + pd.set(14, 1); // output_transpose + pd.set(18, quantize_term); + int ret = k_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + Mat weights[4]; + weights[0] = k_weight_data; + weights[1] = k_bias_data; + weights[2] = k_weight_data_quantize_scales; + weights[3] = k_weight_data_input_scales; + ret = k_gemm->load_model(ModelBinFromMatArray(weights)); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = k_gemm->create_pipeline(opt_wq); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + } + + { + v_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!v_gemm) + { + destroy_pipeline(_opt); + return -100; + } + ncnn::ParamDict pd; + pd.set(2, 0); // transA + pd.set(3, 1); // transB + pd.set(4, 0); // constantA + pd.set(5, 1); // constantB + pd.set(6, 1); // constantC + pd.set(7, 0); // M + pd.set(8, embed_dim); // N + pd.set(9, vdim); // K + pd.set(10, 4); // constant_broadcast_type_C + pd.set(11, 0); // output_N1M + pd.set(12, 0); // output_elempack + pd.set(14, 1); // output_transpose + pd.set(18, quantize_term); + int ret = v_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + Mat weights[4]; + weights[0] = v_weight_data; + weights[1] = v_bias_data; + weights[2] = v_weight_data_quantize_scales; + weights[3] = v_weight_data_input_scales; + ret = v_gemm->load_model(ModelBinFromMatArray(weights)); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = v_gemm->create_pipeline(opt_wq); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + } + + { + o_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!o_gemm) + { + destroy_pipeline(_opt); + return -100; + } + ncnn::ParamDict pd; + pd.set(2, 1); // transA + pd.set(3, 1); // transB + pd.set(4, 0); // constantA + pd.set(5, 1); // constantB + pd.set(6, 1); // constantC + pd.set(7, 0); // M = outch + pd.set(8, qdim); // N = size + pd.set(9, embed_dim); // K = maxk*inch + pd.set(10, 4); // constant_broadcast_type_C + pd.set(11, 0); // output_N1M + pd.set(18, quantize_term); + int ret = o_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + Mat weights[4]; + weights[0] = out_weight_data; + weights[1] = out_bias_data; + weights[2] = out_weight_data_quantize_scales; + weights[3] = out_weight_data_input_scales; + ret = o_gemm->load_model(ModelBinFromMatArray(weights)); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = o_gemm->create_pipeline(opt_wq); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + } + + { + qk_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!qk_gemm) + { + destroy_pipeline(_opt); + return -100; + } + ncnn::ParamDict pd; + pd.set(2, 1); // transA + pd.set(3, 0); // transB + pd.set(4, 0); // constantA + pd.set(5, 0); // constantB + pd.set(6, attn_mask ? 0 : 1); // constantC + pd.set(7, 0); // M + pd.set(8, 0); // N + pd.set(9, 0); // K + pd.set(10, attn_mask ? 3 : -1); // constant_broadcast_type_C + pd.set(11, 0); // output_N1M + pd.set(12, 1); // output_elempack + pd.set(13, 1); // output_elemtype = fp32 + int ret = qk_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = qk_gemm->load_model(ModelBinFromMatArray(0)); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + Option opt1 = opt; + opt1.use_bf16_packed = false; + opt1.use_bf16_storage = false; + opt1.num_threads = 1; + ret = qk_gemm->create_pipeline(opt1); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + } + + { + qkv_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!qkv_gemm) + { + destroy_pipeline(_opt); + return -100; + } + ncnn::ParamDict pd; + pd.set(2, 0); // transA + pd.set(3, 1); // transB + pd.set(4, 0); // constantA + pd.set(5, 0); // constantB + pd.set(6, 1); // constantC + pd.set(7, 0); // M + pd.set(8, 0); // N + pd.set(9, 0); // K + pd.set(10, -1); // constant_broadcast_type_C + pd.set(11, 0); // output_N1M + pd.set(12, 1); // output_elempack + pd.set(13, 1); // output_elemtype = fp32 + pd.set(14, 1); // output_transpose + int ret = qkv_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = qkv_gemm->load_model(ModelBinFromMatArray(0)); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + Option opt1 = opt; + opt1.use_bf16_packed = false; + opt1.use_bf16_storage = false; + opt1.num_threads = 1; + ret = qkv_gemm->create_pipeline(opt1); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + } + + q_weight_data.release(); + q_bias_data.release(); + k_weight_data.release(); + k_bias_data.release(); + v_weight_data.release(); + v_bias_data.release(); + out_weight_data.release(); + out_bias_data.release(); + q_weight_data_quantize_scales.release(); + k_weight_data_quantize_scales.release(); + v_weight_data_quantize_scales.release(); + out_weight_data_quantize_scales.release(); + q_weight_data_input_scales.release(); + k_weight_data_input_scales.release(); + v_weight_data_input_scales.release(); + out_weight_data_input_scales.release(); + + return 0; +} +#endif // NCNN_WEIGHT_QUANT + int MultiHeadAttention_mips::create_pipeline(const Option& _opt) { +#if NCNN_WEIGHT_QUANT if (weight_block_quantize) - return 0; + { + if (quantize_term / 100 != 8) + return MultiHeadAttention::create_pipeline(_opt); + + return create_pipeline_wq_int8(_opt); + } +#endif Option opt = _opt; if (int8_scale_term) @@ -259,14 +604,24 @@ int MultiHeadAttention_mips::create_pipeline(const Option& _opt) int MultiHeadAttention_mips::destroy_pipeline(const Option& _opt) { - if (weight_block_quantize) - return 0; + if (weight_block_quantize && quantize_term / 100 != 8) + return MultiHeadAttention::destroy_pipeline(_opt); Option opt = _opt; - if (int8_scale_term) + if (int8_scale_term && !weight_block_quantize) { opt.use_packing_layout = false; // TODO enable packing } + Option opt_wq = opt; + if (weight_block_quantize) + { + opt_wq.use_packing_layout = false; + opt_wq.use_fp16_packed = false; + opt_wq.use_fp16_storage = false; + opt_wq.use_fp16_arithmetic = false; + opt_wq.use_bf16_packed = false; + opt_wq.use_bf16_storage = false; + } if (qk_softmax) { @@ -277,28 +632,28 @@ int MultiHeadAttention_mips::destroy_pipeline(const Option& _opt) if (q_gemm) { - q_gemm->destroy_pipeline(opt); + q_gemm->destroy_pipeline(opt_wq); delete q_gemm; q_gemm = 0; } if (k_gemm) { - k_gemm->destroy_pipeline(opt); + k_gemm->destroy_pipeline(opt_wq); delete k_gemm; k_gemm = 0; } if (v_gemm) { - v_gemm->destroy_pipeline(opt); + v_gemm->destroy_pipeline(opt_wq); delete v_gemm; v_gemm = 0; } if (o_gemm) { - o_gemm->destroy_pipeline(opt); + o_gemm->destroy_pipeline(opt_wq); delete o_gemm; o_gemm = 0; } @@ -321,7 +676,7 @@ int MultiHeadAttention_mips::destroy_pipeline(const Option& _opt) int MultiHeadAttention_mips::forward(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& _opt) const { - if (weight_block_quantize) + if (weight_block_quantize && quantize_term / 100 != 8) return MultiHeadAttention::forward(bottom_blobs, top_blobs, _opt); int q_blob_i = 0; @@ -340,10 +695,20 @@ int MultiHeadAttention_mips::forward(const std::vector& bottom_blobs, std:: const Mat& cached_xv_blob = kv_cache ? bottom_blobs[cached_xv_i] : Mat(); Option opt = _opt; - if (int8_scale_term) + if (int8_scale_term && !weight_block_quantize) { opt.use_packing_layout = false; // TODO enable packing } + Option opt_wq = opt; + if (weight_block_quantize) + { + opt_wq.use_packing_layout = false; + opt_wq.use_fp16_packed = false; + opt_wq.use_fp16_storage = false; + opt_wq.use_fp16_arithmetic = false; + opt_wq.use_bf16_packed = false; + opt_wq.use_bf16_storage = false; + } Mat attn_mask_blob_unpacked; if (attn_mask && attn_mask_blob.elempack != 1) @@ -388,7 +753,7 @@ int MultiHeadAttention_mips::forward(const std::vector& bottom_blobs, std:: const int dst_seqlen = past_seqlen > 0 ? (q_blob_i == k_blob_i ? (past_seqlen + cur_seqlen) : past_seqlen) : cur_seqlen; Mat q_affine; - int retq = q_gemm->forward(q_blob, q_affine, opt); + int retq = q_gemm->forward(q_blob, q_affine, opt_wq); if (retq != 0) return retq; @@ -398,7 +763,7 @@ int MultiHeadAttention_mips::forward(const std::vector& bottom_blobs, std:: if (q_blob_i == k_blob_i) { Mat k_affine_q; - int retk = k_gemm->forward(q_blob, k_affine_q, opt); + int retk = k_gemm->forward(q_blob, k_affine_q, opt_wq); if (retk != 0) return retk; @@ -426,7 +791,7 @@ int MultiHeadAttention_mips::forward(const std::vector& bottom_blobs, std:: } else { - int retk = k_gemm->forward(k_blob, k_affine, opt); + int retk = k_gemm->forward(k_blob, k_affine, opt_wq); if (retk != 0) return retk; } @@ -477,7 +842,7 @@ int MultiHeadAttention_mips::forward(const std::vector& bottom_blobs, std:: if (q_blob_i == v_blob_i) { Mat v_affine_q; - int retk = v_gemm->forward(v_blob, v_affine_q, opt); + int retk = v_gemm->forward(v_blob, v_affine_q, opt_wq); if (retk != 0) return retk; @@ -505,7 +870,7 @@ int MultiHeadAttention_mips::forward(const std::vector& bottom_blobs, std:: } else { - int retv = v_gemm->forward(v_blob, v_affine, opt); + int retv = v_gemm->forward(v_blob, v_affine, opt_wq); if (retv != 0) return retv; } @@ -552,7 +917,7 @@ int MultiHeadAttention_mips::forward(const std::vector& bottom_blobs, std:: v_affine.release(); } - int reto = o_gemm->forward(qkv_cross, top_blobs[0], opt); + int reto = o_gemm->forward(qkv_cross, top_blobs[0], opt_wq); if (reto != 0) return reto; diff --git a/src/layer/mips/multiheadattention_mips.h b/src/layer/mips/multiheadattention_mips.h index bdb1bbbeab97..3bdaf3830edf 100644 --- a/src/layer/mips/multiheadattention_mips.h +++ b/src/layer/mips/multiheadattention_mips.h @@ -18,6 +18,11 @@ class MultiHeadAttention_mips : public MultiHeadAttention virtual int forward(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& opt) const; +protected: +#if NCNN_WEIGHT_QUANT + int create_pipeline_wq_int8(const Option& opt); +#endif + public: Layer* q_gemm; Layer* k_gemm; diff --git a/src/layer/multiheadattention.cpp b/src/layer/multiheadattention.cpp index 47663cee2174..ae32564d4717 100644 --- a/src/layer/multiheadattention.cpp +++ b/src/layer/multiheadattention.cpp @@ -558,6 +558,127 @@ static inline int mha_weight_block_quantize_unpack(const unsigned char* ptr, int return mha_weight_block_quantize_sign_extend((v >> bit_shift) & mask, bits); } +static inline signed char mha_weight_block_quantize_float2int8(float v) +{ + int int32 = static_cast(round(v)); + if (int32 > 127) return 127; + if (int32 < -127) return -127; + return (signed char)int32; +} + +static void mha_weight_block_quantize_activation_row_int8(const Mat& A, int transA, int i, signed char* outptr, float* descale_ptr, int K, int block_size, const float* input_scale_ptr) +{ + const int block_count = (K + block_size - 1) / block_size; + const size_t A_hstep = (size_t)A.w; + const float* ptrA = transA ? 0 : A.row(i); + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = block_size < K - k0 ? block_size : K - k0; + + float absmax = 0.f; + for (int kk = 0; kk < max_kk; kk++) + { + const int k = k0 + kk; + float v = transA ? ((const float*)A)[k * A_hstep + i] : ptrA[k]; + if (input_scale_ptr) + v *= input_scale_ptr[k]; + v = fabsf(v); + if (v > absmax) + absmax = v; + } + + if (absmax == 0.f) + { + descale_ptr[g] = 0.f; + for (int kk = 0; kk < max_kk; kk++) + outptr[k0 + kk] = 0; + continue; + } + + volatile double scale_fp64 = 127.0 / (double)absmax; + const float scale = (float)scale_fp64; + descale_ptr[g] = absmax / 127.f; + + for (int kk = 0; kk < max_kk; kk++) + { + const int k = k0 + kk; + float v = transA ? ((const float*)A)[k * A_hstep + i] : ptrA[k]; + if (input_scale_ptr) + { + v *= input_scale_ptr[k]; + volatile float v_ordered = v; + v = v_ordered; + } + outptr[k] = mha_weight_block_quantize_float2int8(v * scale); + } + } +} + +static int mha_weight_block_quantize_gemm_transB_int8(const Mat& A, int transA, const Mat& BT, const Mat& BT_scales, const Mat& input_scales, const Mat& C, Mat& top_blob, int M, int N, int K, int block_size, float alpha, int output_transpose, int output_m_offset, const Option& opt) +{ + const int block_count = (K + block_size - 1) / block_size; + + Mat A_int8; + A_int8.create(K, M, (size_t)1u, opt.workspace_allocator); + if (A_int8.empty()) + return -100; + + Mat A_descales; + A_descales.create(block_count, M, (size_t)4u, opt.workspace_allocator); + if (A_descales.empty()) + return -100; + + const float* input_scale_ptr = input_scales; + + #pragma omp parallel for num_threads(opt.num_threads) + for (int i = 0; i < M; i++) + { + signed char* outptr = A_int8.row(i); + float* descale_ptr = A_descales.row(i); + mha_weight_block_quantize_activation_row_int8(A, transA, i, outptr, descale_ptr, K, block_size, input_scale_ptr); + } + + const float* bias_ptr = C; + + #pragma omp parallel for num_threads(opt.num_threads) + for (int mn = 0; mn < M * N; mn++) + { + const int i = mn / N; + const int j = mn % N; + const signed char* ptrA = A_int8.row(i); + const signed char* ptrB = BT.row(j); + const float* A_descale_ptr = A_descales.row(i); + const float* B_scale_ptr = BT_scales.row(j); + + float sum = bias_ptr[j]; + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = block_size < K - k0 ? block_size : K - k0; + + int sum_int32 = 0; + for (int kk = 0; kk < max_kk; kk++) + { + const int k = k0 + kk; + sum_int32 += ptrA[k] * ptrB[k]; + } + + sum += sum_int32 * A_descale_ptr[g] / B_scale_ptr[g]; + } + + sum *= alpha; + + if (output_transpose) + top_blob.row(j)[output_m_offset + i] = sum; + else + top_blob.row(output_m_offset + i)[j] = sum; + } + + return 0; +} + int MultiHeadAttention::forward_weight_block_quantize(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& opt) const { int q_blob_i = 0; @@ -603,35 +724,44 @@ int MultiHeadAttention::forward_weight_block_quantize(const std::vector& bo if (q_affine.empty()) return -100; - const int packed_k_bytes = mha_weight_quantize_packed_k_bytes(qdim, weight_bits); - - #pragma omp parallel for num_threads(opt.num_threads) - for (int i = 0; i < src_seqlen; i++) + if (weight_bits == 8) { - for (int j = 0; j < embed_dim; j++) - { - const float* ptr = q_blob.row(i); - const unsigned char* kptr = q_weight_data.row(j); - const float* scale_ptr = q_weight_data_quantize_scales.row(j); + int ret = mha_weight_block_quantize_gemm_transB_int8(q_blob, 0, q_weight_data, q_weight_data_quantize_scales, q_weight_data_input_scales, q_bias_data, q_affine, src_seqlen, embed_dim, qdim, block_size, scale, 1, 0, opt); + if (ret != 0) + return ret; + } + else + { + const int packed_k_bytes = mha_weight_quantize_packed_k_bytes(qdim, weight_bits); - float sum = q_bias_data[j]; - for (int k0 = 0; k0 < qdim; k0 += block_size) + #pragma omp parallel for num_threads(opt.num_threads) + for (int i = 0; i < src_seqlen; i++) + { + for (int j = 0; j < embed_dim; j++) { - const int max_kk = block_size < qdim - k0 ? block_size : qdim - k0; - const float descale = 1.f / scale_ptr[k0 / block_size]; + const float* ptr = q_blob.row(i); + const unsigned char* kptr = q_weight_data.row(j); + const float* scale_ptr = q_weight_data_quantize_scales.row(j); - for (int kk = 0; kk < max_kk; kk++) + float sum = q_bias_data[j]; + for (int k0 = 0; k0 < qdim; k0 += block_size) { - const int k = k0 + kk; - const int q = mha_weight_block_quantize_unpack(kptr, k, weight_bits, packed_k_bytes); - float v = ptr[k]; - if (q_input_scale_ptr) - v *= q_input_scale_ptr[k]; - sum += v * (q * descale); + const int max_kk = block_size < qdim - k0 ? block_size : qdim - k0; + const float descale = 1.f / scale_ptr[k0 / block_size]; + + for (int kk = 0; kk < max_kk; kk++) + { + const int k = k0 + kk; + const int q = mha_weight_block_quantize_unpack(kptr, k, weight_bits, packed_k_bytes); + float v = ptr[k]; + if (q_input_scale_ptr) + v *= q_input_scale_ptr[k]; + sum += v * (q * descale); + } } - } - q_affine.row(j)[i] = sum * scale; + q_affine.row(j)[i] = sum * scale; + } } } } @@ -657,35 +787,44 @@ int MultiHeadAttention::forward_weight_block_quantize(const std::vector& bo } } - const int packed_k_bytes = mha_weight_quantize_packed_k_bytes(kdim, weight_bits); - - #pragma omp parallel for num_threads(opt.num_threads) - for (int i = 0; i < cur_seqlen; i++) + if (weight_bits == 8) { - for (int j = 0; j < embed_dim; j++) - { - const float* ptr = k_blob.row(i); - const unsigned char* kptr = k_weight_data.row(j); - const float* scale_ptr = k_weight_data_quantize_scales.row(j); + int ret = mha_weight_block_quantize_gemm_transB_int8(k_blob, 0, k_weight_data, k_weight_data_quantize_scales, k_weight_data_input_scales, k_bias_data, k_affine, cur_seqlen, embed_dim, kdim, block_size, 1.f, 1, past_seqlen, opt); + if (ret != 0) + return ret; + } + else + { + const int packed_k_bytes = mha_weight_quantize_packed_k_bytes(kdim, weight_bits); - float sum = k_bias_data[j]; - for (int k0 = 0; k0 < kdim; k0 += block_size) + #pragma omp parallel for num_threads(opt.num_threads) + for (int i = 0; i < cur_seqlen; i++) + { + for (int j = 0; j < embed_dim; j++) { - const int max_kk = block_size < kdim - k0 ? block_size : kdim - k0; - const float descale = 1.f / scale_ptr[k0 / block_size]; + const float* ptr = k_blob.row(i); + const unsigned char* kptr = k_weight_data.row(j); + const float* scale_ptr = k_weight_data_quantize_scales.row(j); - for (int kk = 0; kk < max_kk; kk++) + float sum = k_bias_data[j]; + for (int k0 = 0; k0 < kdim; k0 += block_size) { - const int k = k0 + kk; - const int q = mha_weight_block_quantize_unpack(kptr, k, weight_bits, packed_k_bytes); - float v = ptr[k]; - if (k_input_scale_ptr) - v *= k_input_scale_ptr[k]; - sum += v * (q * descale); + const int max_kk = block_size < kdim - k0 ? block_size : kdim - k0; + const float descale = 1.f / scale_ptr[k0 / block_size]; + + for (int kk = 0; kk < max_kk; kk++) + { + const int k = k0 + kk; + const int q = mha_weight_block_quantize_unpack(kptr, k, weight_bits, packed_k_bytes); + float v = ptr[k]; + if (k_input_scale_ptr) + v *= k_input_scale_ptr[k]; + sum += v * (q * descale); + } } - } - k_affine.row(j)[past_seqlen + i] = sum; + k_affine.row(j)[past_seqlen + i] = sum; + } } } } @@ -711,35 +850,44 @@ int MultiHeadAttention::forward_weight_block_quantize(const std::vector& bo } } - const int packed_k_bytes = mha_weight_quantize_packed_k_bytes(vdim, weight_bits); - - #pragma omp parallel for num_threads(opt.num_threads) - for (int i = 0; i < cur_seqlen; i++) + if (weight_bits == 8) { - for (int j = 0; j < embed_dim; j++) - { - const float* ptr = v_blob.row(i); - const unsigned char* kptr = v_weight_data.row(j); - const float* scale_ptr = v_weight_data_quantize_scales.row(j); + int ret = mha_weight_block_quantize_gemm_transB_int8(v_blob, 0, v_weight_data, v_weight_data_quantize_scales, v_weight_data_input_scales, v_bias_data, v_affine, cur_seqlen, embed_dim, vdim, block_size, 1.f, 1, past_seqlen, opt); + if (ret != 0) + return ret; + } + else + { + const int packed_k_bytes = mha_weight_quantize_packed_k_bytes(vdim, weight_bits); - float sum = v_bias_data[j]; - for (int k0 = 0; k0 < vdim; k0 += block_size) + #pragma omp parallel for num_threads(opt.num_threads) + for (int i = 0; i < cur_seqlen; i++) + { + for (int j = 0; j < embed_dim; j++) { - const int max_kk = block_size < vdim - k0 ? block_size : vdim - k0; - const float descale = 1.f / scale_ptr[k0 / block_size]; + const float* ptr = v_blob.row(i); + const unsigned char* kptr = v_weight_data.row(j); + const float* scale_ptr = v_weight_data_quantize_scales.row(j); - for (int kk = 0; kk < max_kk; kk++) + float sum = v_bias_data[j]; + for (int k0 = 0; k0 < vdim; k0 += block_size) { - const int k = k0 + kk; - const int q = mha_weight_block_quantize_unpack(kptr, k, weight_bits, packed_k_bytes); - float v = ptr[k]; - if (v_input_scale_ptr) - v *= v_input_scale_ptr[k]; - sum += v * (q * descale); + const int max_kk = block_size < vdim - k0 ? block_size : vdim - k0; + const float descale = 1.f / scale_ptr[k0 / block_size]; + + for (int kk = 0; kk < max_kk; kk++) + { + const int k = k0 + kk; + const int q = mha_weight_block_quantize_unpack(kptr, k, weight_bits, packed_k_bytes); + float v = ptr[k]; + if (v_input_scale_ptr) + v *= v_input_scale_ptr[k]; + sum += v * (q * descale); + } } - } - v_affine.row(j)[past_seqlen + i] = sum; + v_affine.row(j)[past_seqlen + i] = sum; + } } } } @@ -866,36 +1014,45 @@ int MultiHeadAttention::forward_weight_block_quantize(const std::vector& bo if (top_blob.empty()) return -100; - const int packed_k_bytes = mha_weight_quantize_packed_k_bytes(embed_dim, weight_bits); - - #pragma omp parallel for num_threads(opt.num_threads) - for (int i = 0; i < src_seqlen; i++) + if (weight_bits == 8) { - float* outptr = top_blob.row(i); + int ret = mha_weight_block_quantize_gemm_transB_int8(qkv_cross, 1, out_weight_data, out_weight_data_quantize_scales, out_weight_data_input_scales, out_bias_data, top_blob, src_seqlen, qdim, embed_dim, block_size, 1.f, 0, 0, opt); + if (ret != 0) + return ret; + } + else + { + const int packed_k_bytes = mha_weight_quantize_packed_k_bytes(embed_dim, weight_bits); - for (int j = 0; j < qdim; j++) + #pragma omp parallel for num_threads(opt.num_threads) + for (int i = 0; i < src_seqlen; i++) { - const unsigned char* kptr = out_weight_data.row(j); - const float* scale_ptr = out_weight_data_quantize_scales.row(j); + float* outptr = top_blob.row(i); - float sum = out_bias_data[j]; - for (int k0 = 0; k0 < embed_dim; k0 += block_size) + for (int j = 0; j < qdim; j++) { - const int max_kk = block_size < embed_dim - k0 ? block_size : embed_dim - k0; - const float descale = 1.f / scale_ptr[k0 / block_size]; + const unsigned char* kptr = out_weight_data.row(j); + const float* scale_ptr = out_weight_data_quantize_scales.row(j); - for (int kk = 0; kk < max_kk; kk++) + float sum = out_bias_data[j]; + for (int k0 = 0; k0 < embed_dim; k0 += block_size) { - const int k = k0 + kk; - const int q = mha_weight_block_quantize_unpack(kptr, k, weight_bits, packed_k_bytes); - float v = qkv_cross.row(k)[i]; - if (out_input_scale_ptr) - v *= out_input_scale_ptr[k]; - sum += v * (q * descale); + const int max_kk = block_size < embed_dim - k0 ? block_size : embed_dim - k0; + const float descale = 1.f / scale_ptr[k0 / block_size]; + + for (int kk = 0; kk < max_kk; kk++) + { + const int k = k0 + kk; + const int q = mha_weight_block_quantize_unpack(kptr, k, weight_bits, packed_k_bytes); + float v = qkv_cross.row(k)[i]; + if (out_input_scale_ptr) + v *= out_input_scale_ptr[k]; + sum += v * (q * descale); + } } - } - outptr[j] = sum; + outptr[j] = sum; + } } } } diff --git a/src/layer/riscv/gemm_riscv.cpp b/src/layer/riscv/gemm_riscv.cpp index 9e03af872798..ea5bf48477ca 100644 --- a/src/layer/riscv/gemm_riscv.cpp +++ b/src/layer/riscv/gemm_riscv.cpp @@ -13,6 +13,10 @@ namespace ncnn { +#if NCNN_WEIGHT_QUANT +#include "gemm_wq_int8.h" +#endif + Gemm_riscv::Gemm_riscv() { #if __riscv_vector @@ -1869,11 +1873,251 @@ static int gemm_AT_BT_riscv(const Mat& AT, const Mat& BT, const Mat& C, Mat& top return 0; } +#if NCNN_WEIGHT_QUANT +static int gemm_BT_riscv_wq_int8(const Mat& A, const Mat& packed_B, const Mat& packed_B_descales, const Mat& input_scales, const Mat& C, Mat& top_blob, int broadcast_type_C, int N, int K, int block_size, int transA, int output_transpose, float alpha, float beta, int constant_TILE_M, int constant_TILE_N, int constant_TILE_K, int nT, const Option& opt) +{ + const int M = transA ? A.w : (A.dims == 3 ? A.c : A.h) * A.elempack; + const int block_count = (K + block_size - 1) / block_size; + int TILE_M, TILE_N, TILE_K; + get_optimal_tile_mnk_wq_int8(M, N, K, constant_TILE_M, constant_TILE_N, constant_TILE_K, TILE_M, TILE_N, TILE_K, nT); + + (void)TILE_K; + const int mr = std::min(M, TILE_M); + const int nr = std::min(N, TILE_N); + const int nn_M = (M + TILE_M - 1) / TILE_M; + const int nn_N = (N + TILE_N - 1) / TILE_N; + const float* input_scale_ptr = input_scales; + Mat BT = packed_B.reshape(K, N); + Mat BT_descales = packed_B_descales.reshape(block_count, N); + + Mat topT(mr * nr, 1, nT, (size_t)4u, 1, opt.workspace_allocator); + if (topT.empty()) + return -100; + + if (nT > nn_M) + { + Mat AT(K, mr, nn_M, (size_t)1u, 1, opt.workspace_allocator); + Mat AT_descales(block_count, mr, nn_M, (size_t)4u, 1, opt.workspace_allocator); + if (AT.empty() || AT_descales.empty()) + return -100; + + #pragma omp parallel for num_threads(nT) + for (int ppi = 0; ppi < nn_M; ppi++) + { + const int i = ppi * TILE_M; + const int max_ii = std::min(M - i, TILE_M); + + Mat AT_tile = AT.channel(i / TILE_M).row_range(0, max_ii); + Mat AT_descales_tile = AT_descales.channel(i / TILE_M).row_range(0, max_ii); + + if (transA) + transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, block_size, input_scale_ptr); + else + quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, block_size, input_scale_ptr); + } + + const int nn_MN = nn_M * nn_N; + #pragma omp parallel for num_threads(nT) + for (int ppij = 0; ppij < nn_MN; ppij++) + { + const int ppi = ppij / nn_N; + const int ppj = ppij % nn_N; + + const int i = ppi * TILE_M; + const int j = ppj * TILE_N; + const int max_ii = std::min(M - i, TILE_M); + const int max_jj = std::min(N - j, TILE_N); + + Mat AT_tile = AT.channel(i / TILE_M).row_range(0, max_ii); + Mat AT_descales_tile = AT_descales.channel(i / TILE_M).row_range(0, max_ii); + Mat BT_tile = BT.row_range(j, max_jj); + Mat BT_descales_tile = BT_descales.row_range(j, max_jj); + Mat topT_tile = topT.channel(get_omp_thread_num()); + + gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, block_size); + + if (output_transpose) + transpose_unpack_output_tile_wq_int8(topT_tile, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, N, alpha, beta); + else + unpack_output_tile_wq_int8(topT_tile, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, N, alpha, beta); + } + } + else + { + Mat ATX(K, mr, nT, (size_t)1u, 1, opt.workspace_allocator); + Mat ATX_descales(block_count, mr, nT, (size_t)4u, 1, opt.workspace_allocator); + if (ATX.empty() || ATX_descales.empty()) + return -100; + + #pragma omp parallel for num_threads(nT) + for (int ppi = 0; ppi < nn_M; ppi++) + { + const int i = ppi * TILE_M; + const int max_ii = std::min(M - i, TILE_M); + + Mat AT_tile = ATX.channel(get_omp_thread_num()).row_range(0, max_ii); + Mat AT_descales_tile = ATX_descales.channel(get_omp_thread_num()).row_range(0, max_ii); + Mat topT_tile = topT.channel(get_omp_thread_num()); + + if (transA) + transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, block_size, input_scale_ptr); + else + quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, block_size, input_scale_ptr); + + for (int j = 0; j < N; j += TILE_N) + { + const int max_jj = std::min(N - j, TILE_N); + Mat BT_tile = BT.row_range(j, max_jj); + Mat BT_descales_tile = BT_descales.row_range(j, max_jj); + gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, block_size); + + if (output_transpose) + transpose_unpack_output_tile_wq_int8(topT_tile, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, N, alpha, beta); + else + unpack_output_tile_wq_int8(topT_tile, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, N, alpha, beta); + } + } + } + + return 0; +} + +static int gemm_weight_quantize_bits_riscv(int quantize_term) +{ + return quantize_term / 100; +} + +static int gemm_weight_quantize_block_size_riscv(int quantize_term) +{ + const int block_size_code = quantize_term % 10; + return block_size_code == 0 ? 32 : block_size_code == 1 ? 64 : 128; +} + +int Gemm_riscv::forward_weight_block_quantize_int8(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& opt) const +{ + const Mat& A = bottom_blobs[0]; + if (A.elemsize != 4u || A.elempack != 1) + { + NCNN_LOGE("Gemm unsupported input"); + return -1; + } + + if (transA && A.dims != 2) + { + NCNN_LOGE("Gemm unsupported input"); + return -1; + } + + const int K = transA ? A.h : A.w; + if (K != constantK) + { + NCNN_LOGE("Gemm weight block quantize K mismatch"); + return -1; + } + + const int M = transA ? A.w : A.dims == 3 ? A.c : A.h; + const int N = constantN; + const int block_size = gemm_weight_quantize_block_size_riscv(quantize_term); + + Mat C; + int broadcast_type_C = -1; + if (constantC) + { + C = C_data; + broadcast_type_C = constant_broadcast_type_C; + } + else + { + if (bottom_blobs.size() == 2) + C = bottom_blobs[1]; + + if (!C.empty()) + { + bool matched = false; + if (C.dims == 1 && C.w == 1) + { + broadcast_type_C = 0; + matched = true; + } + if (C.dims == 1 && C.w == M) + { + broadcast_type_C = 1; + matched = true; + } + if (C.dims == 1 && C.w == N) + { + broadcast_type_C = 4; + matched = true; + } + if (C.dims == 2 && C.w == 1 && C.h == M) + { + broadcast_type_C = 2; + matched = true; + } + if (C.dims == 2 && C.w == N && C.h == M) + { + broadcast_type_C = 3; + matched = true; + } + if (C.dims == 2 && C.w == N && C.h == 1) + { + broadcast_type_C = 4; + matched = true; + } + + if (!matched || C.elemsize != 4u || C.elempack != 1) + { + NCNN_LOGE("Gemm unsupported C"); + return -1; + } + } + } + + if (!C.empty() && (C.elemsize != 4u || C.elempack != 1)) + { + NCNN_LOGE("Gemm unsupported C"); + return -1; + } + + Mat& top_blob = top_blobs[0]; + if (output_transpose) + top_blob.create(M, N, (size_t)4u, opt.blob_allocator); + else + top_blob.create(N, M, (size_t)4u, opt.blob_allocator); + if (top_blob.empty()) + return -100; + + return gemm_BT_riscv_wq_int8(A, B_data_w8a8_packed, B_data_w8a8_descales, B_data_input_scales, C, top_blob, broadcast_type_C, N, K, block_size, transA, output_transpose, alpha, beta, constant_TILE_M, constant_TILE_N, constant_TILE_K, opt.num_threads, opt); +} +#endif // NCNN_WEIGHT_QUANT + int Gemm_riscv::create_pipeline(const Option& opt) { if (weight_block_quantize) { - return 0; +#if NCNN_WEIGHT_QUANT + if (gemm_weight_quantize_bits_riscv(quantize_term) == 8) + { + if (!B_data_w8a8_packed.empty()) + return 0; + + Mat B_data_packed; + Mat B_data_descales; + int ret = pack_B_wq_int8(B_data, B_data_quantize_scales, B_data_packed, B_data_descales, constantN, constantK, gemm_weight_quantize_block_size_riscv(quantize_term), opt); + if (ret != 0) + return ret; + + B_data_w8a8_packed = B_data_packed; + B_data_w8a8_descales = B_data_descales; + + B_data.release(); + B_data_quantize_scales.release(); + + return 0; + } +#endif // NCNN_WEIGHT_QUANT + + return Gemm::create_pipeline(opt); } #if NCNN_INT8 @@ -2019,10 +2263,25 @@ int Gemm_riscv::create_pipeline(const Option& opt) return 0; } +int Gemm_riscv::destroy_pipeline(const Option& opt) +{ +#if NCNN_WEIGHT_QUANT + B_data_w8a8_packed.release(); + B_data_w8a8_descales.release(); +#endif + + return Gemm::destroy_pipeline(opt); +} + int Gemm_riscv::forward(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& opt) const { if (weight_block_quantize) { +#if NCNN_WEIGHT_QUANT + if (gemm_weight_quantize_bits_riscv(quantize_term) == 8 && !B_data_w8a8_packed.empty()) + return forward_weight_block_quantize_int8(bottom_blobs, top_blobs, opt); +#endif + return Gemm::forward(bottom_blobs, top_blobs, opt); } diff --git a/src/layer/riscv/gemm_riscv.h b/src/layer/riscv/gemm_riscv.h index 2ef61f268927..88c3a27e78d2 100644 --- a/src/layer/riscv/gemm_riscv.h +++ b/src/layer/riscv/gemm_riscv.h @@ -14,10 +14,14 @@ class Gemm_riscv : public Gemm Gemm_riscv(); virtual int create_pipeline(const Option& opt); + virtual int destroy_pipeline(const Option& opt); virtual int forward(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& opt) const; protected: +#if NCNN_WEIGHT_QUANT + int forward_weight_block_quantize_int8(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& opt) const; +#endif #if NCNN_ZFH int create_pipeline_fp16s(const Option& opt); int forward_fp16s(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& opt) const; @@ -28,6 +32,10 @@ class Gemm_riscv : public Gemm Mat AT_data; Mat BT_data; Mat CT_data; +#if NCNN_WEIGHT_QUANT + Mat B_data_w8a8_packed; + Mat B_data_w8a8_descales; +#endif }; } // namespace ncnn diff --git a/src/layer/riscv/gemm_wq_int8.h b/src/layer/riscv/gemm_wq_int8.h new file mode 100644 index 000000000000..2d585acbb5b8 --- /dev/null +++ b/src/layer/riscv/gemm_wq_int8.h @@ -0,0 +1,2217 @@ +// Copyright 2026 Tencent +// SPDX-License-Identifier: BSD-3-Clause + +// output-major tile, block-major within each output tile +static int pack_B_wq_int8(const Mat& B, const Mat& B_scales, Mat& packed_B, Mat& packed_B_descales, int N, int K, int block_size, const Option& opt) +{ +#if __riscv_vector + const int packn = csrr_vlenb(); +#else + const int packn = 4; +#endif + const int block_count = (K + block_size - 1) / block_size; + + packed_B.create(N * K, (size_t)1u, opt.blob_allocator); + if (packed_B.empty()) + return -100; + packed_B.cstep = (size_t)N * K; + + packed_B_descales.create(N * block_count, (size_t)4u, opt.blob_allocator); + if (packed_B_descales.empty()) + return -100; + packed_B_descales.cstep = (size_t)N * block_count; + + const int nn_N = (N + packn - 1) / packn; + + #pragma omp parallel for num_threads(opt.num_threads) + for (int ppj = 0; ppj < nn_N; ppj++) + { + const int j = ppj * packn; + const int max_jj = std::min(N - j, packn); + signed char* pp = (signed char*)packed_B + (size_t)j * K; + float* pd = (float*)packed_B_descales + (size_t)j * block_count; + + int jj = 0; + for (; jj + 3 < max_jj; jj += 4) + { + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + const signed char* p0 = B.row(j + jj) + k0 + kk; +#if __riscv_vector + const size_t vl = __riscv_vsetvl_e8mf4(4); + const ptrdiff_t B_stride = (ptrdiff_t)B.w; + __riscv_vse8_v_i8mf4(pp, __riscv_vlse8_v_i8mf4(p0, B_stride, vl), vl); + __riscv_vse8_v_i8mf4(pp + 4, __riscv_vlse8_v_i8mf4(p0 + 1, B_stride, vl), vl); + __riscv_vse8_v_i8mf4(pp + 8, __riscv_vlse8_v_i8mf4(p0 + 2, B_stride, vl), vl); + __riscv_vse8_v_i8mf4(pp + 12, __riscv_vlse8_v_i8mf4(p0 + 3, B_stride, vl), vl); +#else + const signed char* p1 = B.row(j + jj + 1) + k0 + kk; + const signed char* p2 = B.row(j + jj + 2) + k0 + kk; + const signed char* p3 = B.row(j + jj + 3) + k0 + kk; + pp[0] = p0[0]; + pp[1] = p0[1]; + pp[2] = p0[2]; + pp[3] = p0[3]; + pp[4] = p1[0]; + pp[5] = p1[1]; + pp[6] = p1[2]; + pp[7] = p1[3]; + pp[8] = p2[0]; + pp[9] = p2[1]; + pp[10] = p2[2]; + pp[11] = p2[3]; + pp[12] = p3[0]; + pp[13] = p3[1]; + pp[14] = p3[2]; + pp[15] = p3[3]; +#endif + pp += 16; + } + for (; kk + 1 < max_kk; kk += 2) + { + const signed char* p0 = B.row(j + jj) + k0 + kk; +#if __riscv_vector + const size_t vl = __riscv_vsetvl_e8mf4(4); + const ptrdiff_t B_stride = (ptrdiff_t)B.w; + __riscv_vse8_v_i8mf4(pp, __riscv_vlse8_v_i8mf4(p0, B_stride, vl), vl); + __riscv_vse8_v_i8mf4(pp + 4, __riscv_vlse8_v_i8mf4(p0 + 1, B_stride, vl), vl); +#else + const signed char* p1 = B.row(j + jj + 1) + k0 + kk; + const signed char* p2 = B.row(j + jj + 2) + k0 + kk; + const signed char* p3 = B.row(j + jj + 3) + k0 + kk; + pp[0] = p0[0]; + pp[1] = p0[1]; + pp[2] = p1[0]; + pp[3] = p1[1]; + pp[4] = p2[0]; + pp[5] = p2[1]; + pp[6] = p3[0]; + pp[7] = p3[1]; +#endif + pp += 8; + } + for (; kk < max_kk; kk++) + { + pp[0] = B.row(j + jj)[k0 + kk]; + pp[1] = B.row(j + jj + 1)[k0 + kk]; + pp[2] = B.row(j + jj + 2)[k0 + kk]; + pp[3] = B.row(j + jj + 3)[k0 + kk]; + pp += 4; + } + + pd[0] = 1.f / B_scales.row(j + jj)[g]; + pd[1] = 1.f / B_scales.row(j + jj + 1)[g]; + pd[2] = 1.f / B_scales.row(j + jj + 2)[g]; + pd[3] = 1.f / B_scales.row(j + jj + 3)[g]; + pd += 4; + } + } + for (; jj + 1 < max_jj; jj += 2) + { + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + const signed char* p0 = B.row(j + jj) + k0 + kk; +#if __riscv_vector + const size_t vl = __riscv_vsetvl_e8mf4(2); + const ptrdiff_t B_stride = (ptrdiff_t)B.w; + __riscv_vse8_v_i8mf4(pp, __riscv_vlse8_v_i8mf4(p0, B_stride, vl), vl); + __riscv_vse8_v_i8mf4(pp + 2, __riscv_vlse8_v_i8mf4(p0 + 1, B_stride, vl), vl); + __riscv_vse8_v_i8mf4(pp + 4, __riscv_vlse8_v_i8mf4(p0 + 2, B_stride, vl), vl); + __riscv_vse8_v_i8mf4(pp + 6, __riscv_vlse8_v_i8mf4(p0 + 3, B_stride, vl), vl); +#else + const signed char* p1 = B.row(j + jj + 1) + k0 + kk; + pp[0] = p0[0]; + pp[1] = p0[1]; + pp[2] = p0[2]; + pp[3] = p0[3]; + pp[4] = p1[0]; + pp[5] = p1[1]; + pp[6] = p1[2]; + pp[7] = p1[3]; +#endif + pp += 8; + } + for (; kk + 1 < max_kk; kk += 2) + { + const signed char* p0 = B.row(j + jj) + k0 + kk; +#if __riscv_vector + const size_t vl = __riscv_vsetvl_e8mf4(2); + const ptrdiff_t B_stride = (ptrdiff_t)B.w; + __riscv_vse8_v_i8mf4(pp, __riscv_vlse8_v_i8mf4(p0, B_stride, vl), vl); + __riscv_vse8_v_i8mf4(pp + 2, __riscv_vlse8_v_i8mf4(p0 + 1, B_stride, vl), vl); +#else + const signed char* p1 = B.row(j + jj + 1) + k0 + kk; + pp[0] = p0[0]; + pp[1] = p0[1]; + pp[2] = p1[0]; + pp[3] = p1[1]; +#endif + pp += 4; + } + for (; kk < max_kk; kk++) + { + pp[0] = B.row(j + jj)[k0 + kk]; + pp[1] = B.row(j + jj + 1)[k0 + kk]; + pp += 2; + } + + pd[0] = 1.f / B_scales.row(j + jj)[g]; + pd[1] = 1.f / B_scales.row(j + jj + 1)[g]; + pd += 2; + } + } + for (; jj < max_jj; jj++) + { + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + const signed char* p0 = B.row(j + jj) + k0; + for (int kk = 0; kk < max_kk; kk++) + *pp++ = p0[kk]; + *pd++ = 1.f / B_scales.row(j + jj)[g]; + } + } + } + + return 0; +} + +// K-major, row-interleaved MR-packn/MR2/MR1 +static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int block_size, const float* input_scale_ptr) +{ + signed char* pp = AT_tile; + float* pd = AT_descales_tile; + const int K = AT_tile.w; + const int block_count = AT_descales_tile.w; + const size_t A_hstep = A.dims == 3 ? A.cstep : (size_t)A.w; + + int ii = 0; +#if __riscv_vector + const int packn = csrr_vlenb() / 4; + const size_t vl = __riscv_vsetvl_e32m1(packn); + const ptrdiff_t A_stride = (ptrdiff_t)A_hstep * sizeof(float); + for (; ii + (packn - 1) < max_ii; ii += packn) + { + const float* p0 = (const float*)A + (size_t)(i + ii) * A_hstep; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + vfloat32m1_t _absmax = __riscv_vfmv_v_f_f32m1(0.f, vl); + + for (int kk = 0; kk < max_kk; kk++) + { + vfloat32m1_t _v = __riscv_vlse32_v_f32m1(p0 + k0 + kk, A_stride, vl); + if (input_scale_ptr) + _v = __riscv_vfmul_vf_f32m1(_v, input_scale_ptr[k0 + kk], vl); + _absmax = __riscv_vfmax_vv_f32m1(_absmax, __riscv_vfabs_v_f32m1(_v, vl), vl); + } + + vfloat32m1_t _scale = __riscv_vfrdiv_vf_f32m1(_absmax, 127.f, vl); + _scale = __riscv_vfmerge_vfm_f32m1(_scale, 0.f, __riscv_vmfeq_vf_f32m1_b32(_absmax, 0.f, vl), vl); + __riscv_vse32_v_f32m1(pd, __riscv_vfmul_vf_f32m1(_absmax, 1.f / 127.f, vl), vl); + pd += packn; + + for (int kk = 0; kk < max_kk; kk++) + { + vfloat32m1_t _v = __riscv_vlse32_v_f32m1(p0 + k0 + kk, A_stride, vl); + if (input_scale_ptr) + _v = __riscv_vfmul_vf_f32m1(_v, input_scale_ptr[k0 + kk], vl); + vint32m1_t _v32 = __riscv_vfcvt_x_f_v_i32m1_rm(__riscv_vfmul_vv_f32m1(_v, _scale, vl), __RISCV_FRM_RMM, vl); + _v32 = __riscv_vmax_vx_i32m1(_v32, -127, vl); + _v32 = __riscv_vmin_vx_i32m1(_v32, 127, vl); + const vint16mf2_t _v16 = __riscv_vnclip_wx_i16mf2(_v32, 0, __RISCV_VXRM_RNU, vl); + __riscv_vse8_v_i8mf4(pp, __riscv_vnclip_wx_i8mf4(_v16, 0, __RISCV_VXRM_RNU, vl), vl); + pp += packn; + } + } + } +#endif + for (; ii + 1 < max_ii; ii += 2) + { + const int i0 = i + ii; + const int i1 = i + ii + 1; + const float* p0 = (const float*)A + (size_t)i0 * A_hstep; + const float* p1 = (const float*)A + (size_t)i1 * A_hstep; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + float absmax0 = 0.f; + float absmax1 = 0.f; + +#if __riscv_vector + int kk = 0; + while (kk < max_kk) + { + const size_t vl = __riscv_vsetvl_e32m8(max_kk - kk); + vfloat32m8_t _v0 = __riscv_vle32_v_f32m8(p0 + k0 + kk, vl); + vfloat32m8_t _v1 = __riscv_vle32_v_f32m8(p1 + k0 + kk, vl); + if (input_scale_ptr) + { + const vfloat32m8_t _s = __riscv_vle32_v_f32m8(input_scale_ptr + k0 + kk, vl); + _v0 = __riscv_vfmul_vv_f32m8(_v0, _s, vl); + _v1 = __riscv_vfmul_vv_f32m8(_v1, _s, vl); + } + _v0 = __riscv_vfabs_v_f32m8(_v0, vl); + _v1 = __riscv_vfabs_v_f32m8(_v1, vl); + absmax0 = __riscv_vfmv_f_s_f32m1_f32(__riscv_vfredmax_vs_f32m8_f32m1(_v0, __riscv_vfmv_s_f_f32m1(absmax0, 1), vl)); + absmax1 = __riscv_vfmv_f_s_f32m1_f32(__riscv_vfredmax_vs_f32m8_f32m1(_v1, __riscv_vfmv_s_f_f32m1(absmax1, 1), vl)); + kk += vl; + } +#else + for (int kk = 0; kk < max_kk; kk++) + { + const int k = k0 + kk; + float v0 = p0[k]; + float v1 = p1[k]; + if (input_scale_ptr) + { + v0 *= input_scale_ptr[k]; + v1 *= input_scale_ptr[k]; + } + absmax0 = std::max(absmax0, fabsf(v0)); + absmax1 = std::max(absmax1, fabsf(v1)); + } +#endif + + volatile double scale0_fp64 = absmax0 == 0.f ? 0.0 : 127.0 / (double)absmax0; + volatile double scale1_fp64 = absmax1 == 0.f ? 0.0 : 127.0 / (double)absmax1; + const float scale0 = (float)scale0_fp64; + const float scale1 = (float)scale1_fp64; + pd[0] = absmax0 / 127.f; + pd[1] = absmax1 / 127.f; + pd += 2; + +#if __riscv_vector + kk = 0; + while (kk < max_kk) + { + const size_t vl = __riscv_vsetvl_e32m8(max_kk - kk); + vfloat32m8_t _v0 = __riscv_vle32_v_f32m8(p0 + k0 + kk, vl); + vfloat32m8_t _v1 = __riscv_vle32_v_f32m8(p1 + k0 + kk, vl); + if (input_scale_ptr) + { + const vfloat32m8_t _s = __riscv_vle32_v_f32m8(input_scale_ptr + k0 + kk, vl); + _v0 = __riscv_vfmul_vv_f32m8(_v0, _s, vl); + _v1 = __riscv_vfmul_vv_f32m8(_v1, _s, vl); + } + const vint8m2_t _q0 = float2int8(__riscv_vfmul_vf_f32m8(_v0, scale0, vl), vl); + const vint8m2_t _q1 = float2int8(__riscv_vfmul_vf_f32m8(_v1, scale1, vl), vl); + __riscv_vsse8_v_i8m2(pp, 2, _q0, vl); + __riscv_vsse8_v_i8m2(pp + 1, 2, _q1, vl); + pp += vl * 2; + kk += vl; + } +#else + for (int kk = 0; kk < max_kk; kk++) + { + const int k = k0 + kk; + float v0 = p0[k]; + float v1 = p1[k]; + if (input_scale_ptr) + { + v0 *= input_scale_ptr[k]; + v1 *= input_scale_ptr[k]; + asm volatile("" : "+f"(v0), "+f"(v1)); + } + pp[0] = float2int8(v0 * scale0); + pp[1] = float2int8(v1 * scale1); + pp += 2; + } +#endif + } + } + for (; ii < max_ii; ii++) + { + const int i0 = i + ii; + const float* p0 = (const float*)A + (size_t)i0 * A_hstep; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + float absmax = 0.f; + +#if __riscv_vector + int kk = 0; + while (kk < max_kk) + { + const size_t vl = __riscv_vsetvl_e32m8(max_kk - kk); + vfloat32m8_t _v = __riscv_vle32_v_f32m8(p0 + k0 + kk, vl); + if (input_scale_ptr) + _v = __riscv_vfmul_vv_f32m8(_v, __riscv_vle32_v_f32m8(input_scale_ptr + k0 + kk, vl), vl); + _v = __riscv_vfabs_v_f32m8(_v, vl); + absmax = __riscv_vfmv_f_s_f32m1_f32(__riscv_vfredmax_vs_f32m8_f32m1(_v, __riscv_vfmv_s_f_f32m1(absmax, 1), vl)); + kk += vl; + } +#else + for (int kk = 0; kk < max_kk; kk++) + { + const int k = k0 + kk; + float v = p0[k]; + if (input_scale_ptr) + v *= input_scale_ptr[k]; + absmax = std::max(absmax, fabsf(v)); + } +#endif + + volatile double scale_fp64 = absmax == 0.f ? 0.0 : 127.0 / (double)absmax; + const float scale = (float)scale_fp64; + *pd++ = absmax / 127.f; + +#if __riscv_vector + kk = 0; + while (kk < max_kk) + { + const size_t vl = __riscv_vsetvl_e32m8(max_kk - kk); + vfloat32m8_t _v = __riscv_vle32_v_f32m8(p0 + k0 + kk, vl); + if (input_scale_ptr) + _v = __riscv_vfmul_vv_f32m8(_v, __riscv_vle32_v_f32m8(input_scale_ptr + k0 + kk, vl), vl); + __riscv_vse8_v_i8m2(pp, float2int8(__riscv_vfmul_vf_f32m8(_v, scale, vl), vl), vl); + pp += vl; + kk += vl; + } +#else + for (int kk = 0; kk < max_kk; kk++) + { + const int k = k0 + kk; + float v = p0[k]; + if (input_scale_ptr) + { + v *= input_scale_ptr[k]; + asm volatile("" : "+f"(v)); + } + *pp++ = float2int8(v * scale); + } +#endif + } + } + +} + +// K-major, row-interleaved MR-packn/MR2/MR1 +static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int block_size, const float* input_scale_ptr) +{ + signed char* pp = AT_tile; + float* pd = AT_descales_tile; + const int K = AT_tile.w; + const int block_count = AT_descales_tile.w; + const size_t A_hstep = A.dims == 3 ? A.cstep : (size_t)A.w; + + int ii = 0; +#if __riscv_vector + const int packn = csrr_vlenb() / 4; + const size_t vl = __riscv_vsetvl_e32m1(packn); + for (; ii + (packn - 1) < max_ii; ii += packn) + { + const int i0 = i + ii; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + vfloat32m1_t _absmax = __riscv_vfmv_v_f_f32m1(0.f, vl); + + for (int kk = 0; kk < max_kk; kk++) + { + vfloat32m1_t _v = __riscv_vle32_v_f32m1((const float*)A + (size_t)(k0 + kk) * A_hstep + i0, vl); + if (input_scale_ptr) + _v = __riscv_vfmul_vf_f32m1(_v, input_scale_ptr[k0 + kk], vl); + _absmax = __riscv_vfmax_vv_f32m1(_absmax, __riscv_vfabs_v_f32m1(_v, vl), vl); + } + + vfloat32m1_t _scale = __riscv_vfrdiv_vf_f32m1(_absmax, 127.f, vl); + _scale = __riscv_vfmerge_vfm_f32m1(_scale, 0.f, __riscv_vmfeq_vf_f32m1_b32(_absmax, 0.f, vl), vl); + __riscv_vse32_v_f32m1(pd, __riscv_vfmul_vf_f32m1(_absmax, 1.f / 127.f, vl), vl); + pd += packn; + + for (int kk = 0; kk < max_kk; kk++) + { + vfloat32m1_t _v = __riscv_vle32_v_f32m1((const float*)A + (size_t)(k0 + kk) * A_hstep + i0, vl); + if (input_scale_ptr) + _v = __riscv_vfmul_vf_f32m1(_v, input_scale_ptr[k0 + kk], vl); + vint32m1_t _v32 = __riscv_vfcvt_x_f_v_i32m1_rm(__riscv_vfmul_vv_f32m1(_v, _scale, vl), __RISCV_FRM_RMM, vl); + _v32 = __riscv_vmax_vx_i32m1(_v32, -127, vl); + _v32 = __riscv_vmin_vx_i32m1(_v32, 127, vl); + const vint16mf2_t _v16 = __riscv_vnclip_wx_i16mf2(_v32, 0, __RISCV_VXRM_RNU, vl); + __riscv_vse8_v_i8mf4(pp, __riscv_vnclip_wx_i8mf4(_v16, 0, __RISCV_VXRM_RNU, vl), vl); + pp += packn; + } + } + } +#endif + for (; ii + 1 < max_ii; ii += 2) + { + const int i0 = i + ii; + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + float absmax0 = 0.f; + float absmax1 = 0.f; + +#if __riscv_vector + int kk = 0; + while (kk < max_kk) + { + const size_t vl = __riscv_vsetvl_e32m8(max_kk - kk); + vfloat32m8_t _v0 = __riscv_vlse32_v_f32m8((const float*)A + (size_t)(k0 + kk) * A_hstep + i0, (ptrdiff_t)A_hstep * sizeof(float), vl); + vfloat32m8_t _v1 = __riscv_vlse32_v_f32m8((const float*)A + (size_t)(k0 + kk) * A_hstep + i0 + 1, (ptrdiff_t)A_hstep * sizeof(float), vl); + if (input_scale_ptr) + { + const vfloat32m8_t _s = __riscv_vle32_v_f32m8(input_scale_ptr + k0 + kk, vl); + _v0 = __riscv_vfmul_vv_f32m8(_v0, _s, vl); + _v1 = __riscv_vfmul_vv_f32m8(_v1, _s, vl); + } + _v0 = __riscv_vfabs_v_f32m8(_v0, vl); + _v1 = __riscv_vfabs_v_f32m8(_v1, vl); + absmax0 = __riscv_vfmv_f_s_f32m1_f32(__riscv_vfredmax_vs_f32m8_f32m1(_v0, __riscv_vfmv_s_f_f32m1(absmax0, 1), vl)); + absmax1 = __riscv_vfmv_f_s_f32m1_f32(__riscv_vfredmax_vs_f32m8_f32m1(_v1, __riscv_vfmv_s_f_f32m1(absmax1, 1), vl)); + kk += vl; + } +#else + for (int kk = 0; kk < max_kk; kk++) + { + const int k = k0 + kk; + const float* p0 = (const float*)A + (size_t)k * A_hstep + i0; + float v0 = p0[0]; + float v1 = p0[1]; + if (input_scale_ptr) + { + const float s = input_scale_ptr[k]; + v0 *= s; + v1 *= s; + } + absmax0 = std::max(absmax0, fabsf(v0)); + absmax1 = std::max(absmax1, fabsf(v1)); + } +#endif + + volatile double scale0_fp64 = absmax0 == 0.f ? 0.0 : 127.0 / (double)absmax0; + volatile double scale1_fp64 = absmax1 == 0.f ? 0.0 : 127.0 / (double)absmax1; + const float scale0 = (float)scale0_fp64; + const float scale1 = (float)scale1_fp64; + pd[0] = absmax0 / 127.f; + pd[1] = absmax1 / 127.f; + pd += 2; + +#if __riscv_vector + kk = 0; + while (kk < max_kk) + { + const size_t vl = __riscv_vsetvl_e32m8(max_kk - kk); + vfloat32m8_t _v0 = __riscv_vlse32_v_f32m8((const float*)A + (size_t)(k0 + kk) * A_hstep + i0, (ptrdiff_t)A_hstep * sizeof(float), vl); + vfloat32m8_t _v1 = __riscv_vlse32_v_f32m8((const float*)A + (size_t)(k0 + kk) * A_hstep + i0 + 1, (ptrdiff_t)A_hstep * sizeof(float), vl); + if (input_scale_ptr) + { + const vfloat32m8_t _s = __riscv_vle32_v_f32m8(input_scale_ptr + k0 + kk, vl); + _v0 = __riscv_vfmul_vv_f32m8(_v0, _s, vl); + _v1 = __riscv_vfmul_vv_f32m8(_v1, _s, vl); + } + const vint8m2_t _q0 = float2int8(__riscv_vfmul_vf_f32m8(_v0, scale0, vl), vl); + const vint8m2_t _q1 = float2int8(__riscv_vfmul_vf_f32m8(_v1, scale1, vl), vl); + __riscv_vsse8_v_i8m2(pp, 2, _q0, vl); + __riscv_vsse8_v_i8m2(pp + 1, 2, _q1, vl); + pp += vl * 2; + kk += vl; + } +#else + for (int kk = 0; kk < max_kk; kk++) + { + const int k = k0 + kk; + const float* p0 = (const float*)A + (size_t)k * A_hstep + i0; + float v0 = p0[0]; + float v1 = p0[1]; + if (input_scale_ptr) + { + const float s = input_scale_ptr[k]; + v0 *= s; + v1 *= s; + asm volatile("" : "+f"(v0), "+f"(v1)); + } + pp[0] = float2int8(v0 * scale0); + pp[1] = float2int8(v1 * scale1); + pp += 2; + } +#endif + } + } + for (; ii < max_ii; ii++) + { + const int i0 = i + ii; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + float absmax = 0.f; + +#if __riscv_vector + int kk = 0; + while (kk < max_kk) + { + const size_t vl = __riscv_vsetvl_e32m8(max_kk - kk); + vfloat32m8_t _v = __riscv_vlse32_v_f32m8((const float*)A + (size_t)(k0 + kk) * A_hstep + i0, (ptrdiff_t)A_hstep * sizeof(float), vl); + if (input_scale_ptr) + _v = __riscv_vfmul_vv_f32m8(_v, __riscv_vle32_v_f32m8(input_scale_ptr + k0 + kk, vl), vl); + _v = __riscv_vfabs_v_f32m8(_v, vl); + absmax = __riscv_vfmv_f_s_f32m1_f32(__riscv_vfredmax_vs_f32m8_f32m1(_v, __riscv_vfmv_s_f_f32m1(absmax, 1), vl)); + kk += vl; + } +#else + for (int kk = 0; kk < max_kk; kk++) + { + const int k = k0 + kk; + float v = ((const float*)A)[(size_t)k * A_hstep + i0]; + if (input_scale_ptr) + v *= input_scale_ptr[k]; + absmax = std::max(absmax, fabsf(v)); + } +#endif + + volatile double scale_fp64 = absmax == 0.f ? 0.0 : 127.0 / (double)absmax; + const float scale = (float)scale_fp64; + *pd++ = absmax / 127.f; + +#if __riscv_vector + kk = 0; + while (kk < max_kk) + { + const size_t vl = __riscv_vsetvl_e32m8(max_kk - kk); + vfloat32m8_t _v = __riscv_vlse32_v_f32m8((const float*)A + (size_t)(k0 + kk) * A_hstep + i0, (ptrdiff_t)A_hstep * sizeof(float), vl); + if (input_scale_ptr) + _v = __riscv_vfmul_vv_f32m8(_v, __riscv_vle32_v_f32m8(input_scale_ptr + k0 + kk, vl), vl); + __riscv_vse8_v_i8m2(pp, float2int8(__riscv_vfmul_vf_f32m8(_v, scale, vl), vl), vl); + pp += vl; + kk += vl; + } +#else + for (int kk = 0; kk < max_kk; kk++) + { + const int k = k0 + kk; + float v = ((const float*)A)[(size_t)k * A_hstep + i0]; + if (input_scale_ptr) + { + v *= input_scale_ptr[k]; + asm volatile("" : "+f"(v)); + } + *pp++ = float2int8(v * scale); + } +#endif + } + } +} + +static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int block_size) +{ + const signed char* pAT = AT_tile; + const float* pAT_descales = AT_descales_tile; + const signed char* pBT = BT_tile; + const float* pBT_descales = BT_descales_tile; + float* outptr = topT_tile; + const int K = AT_tile.w; + const int block_count = AT_descales_tile.w; + + int ii = 0; +#if __riscv_vector + const int packn = csrr_vlenb() / 4; + const size_t vl4 = __riscv_vsetvl_e8mf4(4); + for (; ii < max_ii;) + { + const signed char* pB = pBT; + const float* pB_descales = pBT_descales; + const int mr = ii + (packn - 1) < max_ii ? packn : ii + 1 < max_ii ? 2 : 1; + const size_t vl = __riscv_vsetvl_e32m1(mr); + + int jj = 0; + for (; jj + 3 < max_jj; jj += 4) + { + vfloat32m1_t _fsum0 = __riscv_vfmv_v_f_f32m1(0.f, vl); + vfloat32m1_t _fsum1 = __riscv_vfmv_v_f_f32m1(0.f, vl); + vfloat32m1_t _fsum2 = __riscv_vfmv_v_f_f32m1(0.f, vl); + vfloat32m1_t _fsum3 = __riscv_vfmv_v_f_f32m1(0.f, vl); + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + const int max_kk = std::min(K - k, block_size); + vint32m1_t _sum0 = __riscv_vmv_v_x_i32m1(0, vl); + vint32m1_t _sum1 = __riscv_vmv_v_x_i32m1(0, vl); + vint32m1_t _sum2 = __riscv_vmv_v_x_i32m1(0, vl); + vint32m1_t _sum3 = __riscv_vmv_v_x_i32m1(0, vl); + + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + vint16mf2_t _a = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pA, vl), vl); + uint32_t b = *(const uint32_t*)pB; + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, (signed char)b, _a, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, (signed char)(b >> 8), _a, vl); + _sum2 = __riscv_vwmacc_vx_i32m1(_sum2, (signed char)(b >> 16), _a, vl); + _sum3 = __riscv_vwmacc_vx_i32m1(_sum3, (signed char)(b >> 24), _a, vl); + _a = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pA + mr, vl), vl); + b = *(const uint32_t*)(pB + 4); + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, (signed char)b, _a, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, (signed char)(b >> 8), _a, vl); + _sum2 = __riscv_vwmacc_vx_i32m1(_sum2, (signed char)(b >> 16), _a, vl); + _sum3 = __riscv_vwmacc_vx_i32m1(_sum3, (signed char)(b >> 24), _a, vl); + _a = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pA + mr * 2, vl), vl); + b = *(const uint32_t*)(pB + 8); + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, (signed char)b, _a, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, (signed char)(b >> 8), _a, vl); + _sum2 = __riscv_vwmacc_vx_i32m1(_sum2, (signed char)(b >> 16), _a, vl); + _sum3 = __riscv_vwmacc_vx_i32m1(_sum3, (signed char)(b >> 24), _a, vl); + _a = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pA + mr * 3, vl), vl); + b = *(const uint32_t*)(pB + 12); + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, (signed char)b, _a, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, (signed char)(b >> 8), _a, vl); + _sum2 = __riscv_vwmacc_vx_i32m1(_sum2, (signed char)(b >> 16), _a, vl); + _sum3 = __riscv_vwmacc_vx_i32m1(_sum3, (signed char)(b >> 24), _a, vl); + pA += mr * 4; + pB += 16; + } + for (; kk + 1 < max_kk; kk += 2) + { + vint16mf2_t _a = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pA, vl), vl); + uint32_t b = *(const uint32_t*)pB; + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, (signed char)b, _a, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, (signed char)(b >> 8), _a, vl); + _sum2 = __riscv_vwmacc_vx_i32m1(_sum2, (signed char)(b >> 16), _a, vl); + _sum3 = __riscv_vwmacc_vx_i32m1(_sum3, (signed char)(b >> 24), _a, vl); + _a = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pA + mr, vl), vl); + b = *(const uint32_t*)(pB + 4); + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, (signed char)b, _a, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, (signed char)(b >> 8), _a, vl); + _sum2 = __riscv_vwmacc_vx_i32m1(_sum2, (signed char)(b >> 16), _a, vl); + _sum3 = __riscv_vwmacc_vx_i32m1(_sum3, (signed char)(b >> 24), _a, vl); + pA += mr * 2; + pB += 8; + } + for (; kk < max_kk; kk++) + { + const vint16mf2_t _a = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pA, vl), vl); + const uint32_t b = *(const uint32_t*)pB; + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, (signed char)b, _a, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, (signed char)(b >> 8), _a, vl); + _sum2 = __riscv_vwmacc_vx_i32m1(_sum2, (signed char)(b >> 16), _a, vl); + _sum3 = __riscv_vwmacc_vx_i32m1(_sum3, (signed char)(b >> 24), _a, vl); + pA += mr; + pB += 4; + } + + const vfloat32m1_t _ad = __riscv_vle32_v_f32m1(pA_descales, vl); + vfloat32m1_t _v = __riscv_vfmul_vv_f32m1(__riscv_vfcvt_f_x_v_f32m1(_sum0, vl), _ad, vl); + _fsum0 = __riscv_vfmacc_vf_f32m1(_fsum0, pB_descales[0], _v, vl); + _v = __riscv_vfmul_vv_f32m1(__riscv_vfcvt_f_x_v_f32m1(_sum1, vl), _ad, vl); + _fsum1 = __riscv_vfmacc_vf_f32m1(_fsum1, pB_descales[1], _v, vl); + _v = __riscv_vfmul_vv_f32m1(__riscv_vfcvt_f_x_v_f32m1(_sum2, vl), _ad, vl); + _fsum2 = __riscv_vfmacc_vf_f32m1(_fsum2, pB_descales[2], _v, vl); + _v = __riscv_vfmul_vv_f32m1(__riscv_vfcvt_f_x_v_f32m1(_sum3, vl), _ad, vl); + _fsum3 = __riscv_vfmacc_vf_f32m1(_fsum3, pB_descales[3], _v, vl); + pA_descales += mr; + pB_descales += 4; + } + + __riscv_vse32_v_f32m1(outptr, _fsum0, vl); + __riscv_vse32_v_f32m1(outptr + mr, _fsum1, vl); + __riscv_vse32_v_f32m1(outptr + mr * 2, _fsum2, vl); + __riscv_vse32_v_f32m1(outptr + mr * 3, _fsum3, vl); + outptr += mr * 4; + } + for (; jj + 1 < max_jj; jj += 2) + { + vfloat32m1_t _fsum0 = __riscv_vfmv_v_f_f32m1(0.f, vl); + vfloat32m1_t _fsum1 = __riscv_vfmv_v_f_f32m1(0.f, vl); + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + const int max_kk = std::min(K - k, block_size); + vint32m1_t _sum0 = __riscv_vmv_v_x_i32m1(0, vl); + vint32m1_t _sum1 = __riscv_vmv_v_x_i32m1(0, vl); + + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + vint16mf2_t _a = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pA, vl), vl); + uint16_t b = *(const uint16_t*)pB; + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, (signed char)b, _a, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, (signed char)(b >> 8), _a, vl); + _a = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pA + mr, vl), vl); + b = *(const uint16_t*)(pB + 2); + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, (signed char)b, _a, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, (signed char)(b >> 8), _a, vl); + _a = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pA + mr * 2, vl), vl); + b = *(const uint16_t*)(pB + 4); + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, (signed char)b, _a, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, (signed char)(b >> 8), _a, vl); + _a = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pA + mr * 3, vl), vl); + b = *(const uint16_t*)(pB + 6); + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, (signed char)b, _a, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, (signed char)(b >> 8), _a, vl); + pA += mr * 4; + pB += 8; + } + for (; kk + 1 < max_kk; kk += 2) + { + vint16mf2_t _a = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pA, vl), vl); + uint16_t b = *(const uint16_t*)pB; + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, (signed char)b, _a, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, (signed char)(b >> 8), _a, vl); + _a = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pA + mr, vl), vl); + b = *(const uint16_t*)(pB + 2); + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, (signed char)b, _a, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, (signed char)(b >> 8), _a, vl); + pA += mr * 2; + pB += 4; + } + for (; kk < max_kk; kk++) + { + const vint16mf2_t _a = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pA, vl), vl); + const uint16_t b = *(const uint16_t*)pB; + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, (signed char)b, _a, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, (signed char)(b >> 8), _a, vl); + pA += mr; + pB += 2; + } + + const vfloat32m1_t _ad = __riscv_vle32_v_f32m1(pA_descales, vl); + vfloat32m1_t _v = __riscv_vfmul_vv_f32m1(__riscv_vfcvt_f_x_v_f32m1(_sum0, vl), _ad, vl); + _fsum0 = __riscv_vfmacc_vf_f32m1(_fsum0, pB_descales[0], _v, vl); + _v = __riscv_vfmul_vv_f32m1(__riscv_vfcvt_f_x_v_f32m1(_sum1, vl), _ad, vl); + _fsum1 = __riscv_vfmacc_vf_f32m1(_fsum1, pB_descales[1], _v, vl); + pA_descales += mr; + pB_descales += 2; + } + + __riscv_vse32_v_f32m1(outptr, _fsum0, vl); + __riscv_vse32_v_f32m1(outptr + mr, _fsum1, vl); + outptr += mr * 2; + } + for (; jj < max_jj; jj++) + { + vfloat32m1_t _fsum = __riscv_vfmv_v_f_f32m1(0.f, vl); + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + const int max_kk = std::min(K - k, block_size); + vint32m1_t _sum = __riscv_vmv_v_x_i32m1(0, vl); + + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + vint8mf4_t _b = __riscv_vle8_v_i8mf4(pB, vl4); + const signed char b0 = __riscv_vmv_x_s_i8mf4_i8(_b); + _b = __riscv_vslidedown_vx_i8mf4(_b, 1, vl4); + const signed char b1 = __riscv_vmv_x_s_i8mf4_i8(_b); + _b = __riscv_vslidedown_vx_i8mf4(_b, 1, vl4); + const signed char b2 = __riscv_vmv_x_s_i8mf4_i8(_b); + _b = __riscv_vslidedown_vx_i8mf4(_b, 1, vl4); + const signed char b3 = __riscv_vmv_x_s_i8mf4_i8(_b); + vint16mf2_t _a = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pA, vl), vl); + _sum = __riscv_vwmacc_vx_i32m1(_sum, b0, _a, vl); + _a = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pA + mr, vl), vl); + _sum = __riscv_vwmacc_vx_i32m1(_sum, b1, _a, vl); + _a = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pA + mr * 2, vl), vl); + _sum = __riscv_vwmacc_vx_i32m1(_sum, b2, _a, vl); + _a = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pA + mr * 3, vl), vl); + _sum = __riscv_vwmacc_vx_i32m1(_sum, b3, _a, vl); + pA += mr * 4; + pB += 4; + } + for (; kk + 1 < max_kk; kk += 2) + { + const uint16_t b = *(const uint16_t*)pB; + vint16mf2_t _a = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pA, vl), vl); + _sum = __riscv_vwmacc_vx_i32m1(_sum, (signed char)b, _a, vl); + _a = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pA + mr, vl), vl); + _sum = __riscv_vwmacc_vx_i32m1(_sum, (signed char)(b >> 8), _a, vl); + pA += mr * 2; + pB += 2; + } + for (; kk < max_kk; kk++) + { + const vint16mf2_t _a = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pA, vl), vl); + _sum = __riscv_vwmacc_vx_i32m1(_sum, pB[0], _a, vl); + pA += mr; + pB++; + } + + const vfloat32m1_t _ad = __riscv_vle32_v_f32m1(pA_descales, vl); + const vfloat32m1_t _v = __riscv_vfmul_vv_f32m1(__riscv_vfcvt_f_x_v_f32m1(_sum, vl), _ad, vl); + _fsum = __riscv_vfmacc_vf_f32m1(_fsum, pB_descales[0], _v, vl); + pA_descales += mr; + pB_descales++; + } + + __riscv_vse32_v_f32m1(outptr, _fsum, vl); + outptr += mr; + } + + pAT += K * mr; + pAT_descales += block_count * mr; + ii += mr; + } +#else + for (; ii + 1 < max_ii; ii += 2) + { + const signed char* pB = pBT; + const float* pB_descales = pBT_descales; + int jj = 0; + for (; jj + 3 < max_jj; jj += 4) + { + float sum00 = 0.f; + float sum01 = 0.f; + float sum02 = 0.f; + float sum03 = 0.f; + float sum10 = 0.f; + float sum11 = 0.f; + float sum12 = 0.f; + float sum13 = 0.f; + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + const int max_kk = std::min(K - k, block_size); + int s00 = 0, s01 = 0, s02 = 0, s03 = 0; + int s10 = 0, s11 = 0, s12 = 0, s13 = 0; + + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + const signed char* b0 = pB; + const signed char* b1 = b0 + 4; + const signed char* b2 = b1 + 4; + const signed char* b3 = b2 + 4; + s00 += pA[0] * b0[0] + pA[2] * b0[1] + pA[4] * b0[2] + pA[6] * b0[3]; + s01 += pA[0] * b1[0] + pA[2] * b1[1] + pA[4] * b1[2] + pA[6] * b1[3]; + s02 += pA[0] * b2[0] + pA[2] * b2[1] + pA[4] * b2[2] + pA[6] * b2[3]; + s03 += pA[0] * b3[0] + pA[2] * b3[1] + pA[4] * b3[2] + pA[6] * b3[3]; + s10 += pA[1] * b0[0] + pA[3] * b0[1] + pA[5] * b0[2] + pA[7] * b0[3]; + s11 += pA[1] * b1[0] + pA[3] * b1[1] + pA[5] * b1[2] + pA[7] * b1[3]; + s12 += pA[1] * b2[0] + pA[3] * b2[1] + pA[5] * b2[2] + pA[7] * b2[3]; + s13 += pA[1] * b3[0] + pA[3] * b3[1] + pA[5] * b3[2] + pA[7] * b3[3]; + pA += 8; + pB += 16; + } + for (; kk + 1 < max_kk; kk += 2) + { + const signed char* b0 = pB; + const signed char* b1 = b0 + 2; + const signed char* b2 = b1 + 2; + const signed char* b3 = b2 + 2; + s00 += pA[0] * b0[0] + pA[2] * b0[1]; + s01 += pA[0] * b1[0] + pA[2] * b1[1]; + s02 += pA[0] * b2[0] + pA[2] * b2[1]; + s03 += pA[0] * b3[0] + pA[2] * b3[1]; + s10 += pA[1] * b0[0] + pA[3] * b0[1]; + s11 += pA[1] * b1[0] + pA[3] * b1[1]; + s12 += pA[1] * b2[0] + pA[3] * b2[1]; + s13 += pA[1] * b3[0] + pA[3] * b3[1]; + pA += 4; + pB += 8; + } + for (; kk < max_kk; kk++) + { + s00 += pA[0] * pB[0]; + s01 += pA[0] * pB[1]; + s02 += pA[0] * pB[2]; + s03 += pA[0] * pB[3]; + s10 += pA[1] * pB[0]; + s11 += pA[1] * pB[1]; + s12 += pA[1] * pB[2]; + s13 += pA[1] * pB[3]; + pA += 2; + pB += 4; + } + + const float ad0 = pA_descales[0]; + const float ad1 = pA_descales[1]; + const float* bd = pB_descales; + sum00 += s00 * ad0 * bd[0]; + sum01 += s01 * ad0 * bd[1]; + sum02 += s02 * ad0 * bd[2]; + sum03 += s03 * ad0 * bd[3]; + sum10 += s10 * ad1 * bd[0]; + sum11 += s11 * ad1 * bd[1]; + sum12 += s12 * ad1 * bd[2]; + sum13 += s13 * ad1 * bd[3]; + pA_descales += 2; + pB_descales += 4; + } + + outptr[0] = sum00; + outptr[1] = sum10; + outptr[2] = sum01; + outptr[3] = sum11; + outptr[4] = sum02; + outptr[5] = sum12; + outptr[6] = sum03; + outptr[7] = sum13; + outptr += 8; + } + for (; jj + 1 < max_jj; jj += 2) + { + float sum00 = 0.f, sum01 = 0.f, sum10 = 0.f, sum11 = 0.f; + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + const int max_kk = std::min(K - k, block_size); + int s00 = 0, s01 = 0, s10 = 0, s11 = 0; + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + const signed char* b0 = pB; + const signed char* b1 = b0 + 4; + s00 += pA[0] * b0[0] + pA[2] * b0[1] + pA[4] * b0[2] + pA[6] * b0[3]; + s01 += pA[0] * b1[0] + pA[2] * b1[1] + pA[4] * b1[2] + pA[6] * b1[3]; + s10 += pA[1] * b0[0] + pA[3] * b0[1] + pA[5] * b0[2] + pA[7] * b0[3]; + s11 += pA[1] * b1[0] + pA[3] * b1[1] + pA[5] * b1[2] + pA[7] * b1[3]; + pA += 8; + pB += 8; + } + for (; kk + 1 < max_kk; kk += 2) + { + const signed char* b0 = pB; + const signed char* b1 = b0 + 2; + s00 += pA[0] * b0[0] + pA[2] * b0[1]; + s01 += pA[0] * b1[0] + pA[2] * b1[1]; + s10 += pA[1] * b0[0] + pA[3] * b0[1]; + s11 += pA[1] * b1[0] + pA[3] * b1[1]; + pA += 4; + pB += 4; + } + for (; kk < max_kk; kk++) + { + s00 += pA[0] * pB[0]; + s01 += pA[0] * pB[1]; + s10 += pA[1] * pB[0]; + s11 += pA[1] * pB[1]; + pA += 2; + pB += 2; + } + const float ad0 = pA_descales[0]; + const float ad1 = pA_descales[1]; + const float* bd = pB_descales; + sum00 += s00 * ad0 * bd[0]; + sum01 += s01 * ad0 * bd[1]; + sum10 += s10 * ad1 * bd[0]; + sum11 += s11 * ad1 * bd[1]; + pA_descales += 2; + pB_descales += 2; + } + outptr[0] = sum00; + outptr[1] = sum10; + outptr[2] = sum01; + outptr[3] = sum11; + outptr += 4; + } + for (; jj < max_jj; jj++) + { + float sum0 = 0.f, sum1 = 0.f; + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + const int max_kk = std::min(K - k, block_size); + int s0 = 0, s1 = 0; + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + const signed char* b = pB; + s0 += pA[0] * b[0] + pA[2] * b[1] + pA[4] * b[2] + pA[6] * b[3]; + s1 += pA[1] * b[0] + pA[3] * b[1] + pA[5] * b[2] + pA[7] * b[3]; + pA += 8; + pB += 4; + } + for (; kk + 1 < max_kk; kk += 2) + { + const signed char* b = pB; + s0 += pA[0] * b[0] + pA[2] * b[1]; + s1 += pA[1] * b[0] + pA[3] * b[1]; + pA += 4; + pB += 2; + } + for (; kk < max_kk; kk++) + { + s0 += pA[0] * pB[0]; + s1 += pA[1] * pB[0]; + pA += 2; + pB++; + } + const float bd = pB_descales[0]; + sum0 += s0 * pA_descales[0] * bd; + sum1 += s1 * pA_descales[1] * bd; + pA_descales += 2; + pB_descales++; + } + outptr[0] = sum0; + outptr[1] = sum1; + outptr += 2; + } + pAT += K * 2; + pAT_descales += block_count * 2; + } + for (; ii < max_ii; ii++) + { + const signed char* pB = pBT; + const float* pB_descales = pBT_descales; + int jj = 0; + for (; jj + 3 < max_jj; jj += 4) + { + float sum0 = 0.f, sum1 = 0.f, sum2 = 0.f, sum3 = 0.f; + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + const int max_kk = std::min(K - k, block_size); + int s0 = 0, s1 = 0, s2 = 0, s3 = 0; + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + const signed char* b0 = pB; + const signed char* b1 = b0 + 4; + const signed char* b2 = b1 + 4; + const signed char* b3 = b2 + 4; + s0 += pA[0] * b0[0] + pA[1] * b0[1] + pA[2] * b0[2] + pA[3] * b0[3]; + s1 += pA[0] * b1[0] + pA[1] * b1[1] + pA[2] * b1[2] + pA[3] * b1[3]; + s2 += pA[0] * b2[0] + pA[1] * b2[1] + pA[2] * b2[2] + pA[3] * b2[3]; + s3 += pA[0] * b3[0] + pA[1] * b3[1] + pA[2] * b3[2] + pA[3] * b3[3]; + pA += 4; + pB += 16; + } + for (; kk + 1 < max_kk; kk += 2) + { + const signed char* b0 = pB; + const signed char* b1 = b0 + 2; + const signed char* b2 = b1 + 2; + const signed char* b3 = b2 + 2; + s0 += pA[0] * b0[0] + pA[1] * b0[1]; + s1 += pA[0] * b1[0] + pA[1] * b1[1]; + s2 += pA[0] * b2[0] + pA[1] * b2[1]; + s3 += pA[0] * b3[0] + pA[1] * b3[1]; + pA += 2; + pB += 8; + } + for (; kk < max_kk; kk++) + { + s0 += pA[0] * pB[0]; + s1 += pA[0] * pB[1]; + s2 += pA[0] * pB[2]; + s3 += pA[0] * pB[3]; + pA++; + pB += 4; + } + const float ad = pA_descales[0]; + const float* bd = pB_descales; + sum0 += s0 * ad * bd[0]; + sum1 += s1 * ad * bd[1]; + sum2 += s2 * ad * bd[2]; + sum3 += s3 * ad * bd[3]; + pA_descales++; + pB_descales += 4; + } + outptr[0] = sum0; + outptr[1] = sum1; + outptr[2] = sum2; + outptr[3] = sum3; + outptr += 4; + } + for (; jj + 1 < max_jj; jj += 2) + { + float sum0 = 0.f, sum1 = 0.f; + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + const int max_kk = std::min(K - k, block_size); + int s0 = 0, s1 = 0; + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + const signed char* b0 = pB; + const signed char* b1 = b0 + 4; + s0 += pA[0] * b0[0] + pA[1] * b0[1] + pA[2] * b0[2] + pA[3] * b0[3]; + s1 += pA[0] * b1[0] + pA[1] * b1[1] + pA[2] * b1[2] + pA[3] * b1[3]; + pA += 4; + pB += 8; + } + for (; kk + 1 < max_kk; kk += 2) + { + const signed char* b0 = pB; + const signed char* b1 = b0 + 2; + s0 += pA[0] * b0[0] + pA[1] * b0[1]; + s1 += pA[0] * b1[0] + pA[1] * b1[1]; + pA += 2; + pB += 4; + } + for (; kk < max_kk; kk++) + { + s0 += pA[0] * pB[0]; + s1 += pA[0] * pB[1]; + pA++; + pB += 2; + } + const float ad = pA_descales[0]; + const float* bd = pB_descales; + sum0 += s0 * ad * bd[0]; + sum1 += s1 * ad * bd[1]; + pA_descales++; + pB_descales += 2; + } + outptr[0] = sum0; + outptr[1] = sum1; + outptr += 2; + } + for (; jj < max_jj; jj++) + { + float sum = 0.f; + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + const int max_kk = std::min(K - k, block_size); + int s = 0; + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + const signed char* b = pB; + s += pA[0] * b[0] + pA[1] * b[1] + pA[2] * b[2] + pA[3] * b[3]; + pA += 4; + pB += 4; + } + for (; kk + 1 < max_kk; kk += 2) + { + const signed char* b = pB; + s += pA[0] * b[0] + pA[1] * b[1]; + pA += 2; + pB += 2; + } + for (; kk < max_kk; kk++) + { + s += pA[0] * pB[0]; + pA++; + pB++; + } + sum += s * pA_descales[0] * pB_descales[0]; + pA_descales++; + pB_descales++; + } + outptr[0] = sum; + outptr++; + } + pAT += K; + pAT_descales += block_count; + } +#endif +} + +static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta) +{ + beta *= alpha; + (void)N; + const size_t c_hstep = C.dims == 3 ? C.cstep : (size_t)C.w; + const size_t out_hstep = top_blob.dims == 3 ? top_blob.cstep : (size_t)top_blob.w; + float* outptr = (float*)top_blob + (size_t)i * out_hstep + j; + + int ii = 0; +#if __riscv_vector + const int packn = csrr_vlenb() / 4; + const size_t vl_packn = __riscv_vsetvl_e32m1(packn); + const ptrdiff_t c_stride = (ptrdiff_t)c_hstep * sizeof(float); + const ptrdiff_t out_stride = (ptrdiff_t)out_hstep * sizeof(float); + for (; ii + (packn - 1) < max_ii; ii += packn) + { + const float* pC = C; + if (pC) + { + if (broadcast_type_C == 1 || broadcast_type_C == 2) + pC += i + ii; + if (broadcast_type_C == 3) + pC += (size_t)(i + ii) * c_hstep + j; + if (broadcast_type_C == 4) + pC += j; + } + + float c0 = 0.f; + vfloat32m1_t _c = __riscv_vfmv_v_f_f32m1(0.f, vl_packn); + if (pC && broadcast_type_C == 0) + c0 = pC[0] * beta; + if (pC && (broadcast_type_C == 1 || broadcast_type_C == 2)) + _c = __riscv_vfmul_vf_f32m1(__riscv_vle32_v_f32m1(pC, vl_packn), beta, vl_packn); + + for (int jj = 0; jj < max_jj; jj++) + { + vfloat32m1_t _sum = __riscv_vle32_v_f32m1(pp, vl_packn); + if (alpha != 1.f) + _sum = __riscv_vfmul_vf_f32m1(_sum, alpha, vl_packn); + if (pC) + { + if (broadcast_type_C == 0) + _sum = __riscv_vfadd_vf_f32m1(_sum, c0, vl_packn); + if (broadcast_type_C == 1 || broadcast_type_C == 2) + _sum = __riscv_vfadd_vv_f32m1(_sum, _c, vl_packn); + if (broadcast_type_C == 3) + { + const vfloat32m1_t _c0 = __riscv_vlse32_v_f32m1(pC + jj, c_stride, vl_packn); + if (beta == 1.f) + _sum = __riscv_vfadd_vv_f32m1(_sum, _c0, vl_packn); + else + _sum = __riscv_vfmacc_vf_f32m1(_sum, beta, _c0, vl_packn); + } + if (broadcast_type_C == 4) + _sum = __riscv_vfadd_vf_f32m1(_sum, pC[jj] * beta, vl_packn); + } + __riscv_vsse32_v_f32m1(outptr + jj, out_stride, _sum, vl_packn); + pp += packn; + } + outptr += out_hstep * packn; + } + for (; ii + 1 < max_ii; ii += 2) + { + const float* pC = C; + if (pC) + { + if (broadcast_type_C == 1 || broadcast_type_C == 2) + pC += i + ii; + if (broadcast_type_C == 3) + pC += (size_t)(i + ii) * c_hstep + j; + if (broadcast_type_C == 4) + pC += j; + } + + float c0 = 0.f; + float c1 = 0.f; + if (pC && (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2)) + { + c0 = pC[0] * beta; + c1 = pC[broadcast_type_C == 0 ? 0 : 1] * beta; + } + + float* out0 = outptr; + float* out1 = out0 + out_hstep; + const size_t vl = __riscv_vsetvl_e32m4(max_jj); + const vfloat32m4x2_t _s = __riscv_vlseg2e32_v_f32m4x2(pp, vl); + vfloat32m4_t _sum0 = __riscv_vget_v_f32m4x2_f32m4(_s, 0); + vfloat32m4_t _sum1 = __riscv_vget_v_f32m4x2_f32m4(_s, 1); + if (alpha != 1.f) + { + _sum0 = __riscv_vfmul_vf_f32m4(_sum0, alpha, vl); + _sum1 = __riscv_vfmul_vf_f32m4(_sum1, alpha, vl); + } + + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _sum0 = __riscv_vfadd_vf_f32m4(_sum0, c0, vl); + _sum1 = __riscv_vfadd_vf_f32m4(_sum1, c1, vl); + } + if (broadcast_type_C == 3) + { + const vfloat32m4_t _c0 = __riscv_vle32_v_f32m4(pC, vl); + const vfloat32m4_t _c1 = __riscv_vle32_v_f32m4(pC + c_hstep, vl); + if (beta == 1.f) + { + _sum0 = __riscv_vfadd_vv_f32m4(_sum0, _c0, vl); + _sum1 = __riscv_vfadd_vv_f32m4(_sum1, _c1, vl); + } + else + { + _sum0 = __riscv_vfmacc_vf_f32m4(_sum0, beta, _c0, vl); + _sum1 = __riscv_vfmacc_vf_f32m4(_sum1, beta, _c1, vl); + } + } + if (broadcast_type_C == 4) + { + const vfloat32m4_t _c = __riscv_vle32_v_f32m4(pC, vl); + if (beta == 1.f) + { + _sum0 = __riscv_vfadd_vv_f32m4(_sum0, _c, vl); + _sum1 = __riscv_vfadd_vv_f32m4(_sum1, _c, vl); + } + else + { + _sum0 = __riscv_vfmacc_vf_f32m4(_sum0, beta, _c, vl); + _sum1 = __riscv_vfmacc_vf_f32m4(_sum1, beta, _c, vl); + } + } + } + + __riscv_vse32_v_f32m4(out0, _sum0, vl); + __riscv_vse32_v_f32m4(out1, _sum1, vl); + pp += vl * 2; + outptr += out_hstep * 2; + } + for (; ii < max_ii; ii++) + { + const float* pC = C; + if (pC) + { + if (broadcast_type_C == 1 || broadcast_type_C == 2) + pC += i + ii; + if (broadcast_type_C == 3) + pC += (size_t)(i + ii) * c_hstep + j; + if (broadcast_type_C == 4) + pC += j; + } + + float c0 = 0.f; + if (pC && (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2)) + c0 = pC[0] * beta; + + float* out0 = outptr; + const size_t vl = __riscv_vsetvl_e32m4(max_jj); + vfloat32m4_t _sum = __riscv_vle32_v_f32m4(pp, vl); + if (alpha != 1.f) + _sum = __riscv_vfmul_vf_f32m4(_sum, alpha, vl); + + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + _sum = __riscv_vfadd_vf_f32m4(_sum, c0, vl); + if (broadcast_type_C == 3) + { + const vfloat32m4_t _c = __riscv_vle32_v_f32m4(pC, vl); + if (beta == 1.f) + _sum = __riscv_vfadd_vv_f32m4(_sum, _c, vl); + else + _sum = __riscv_vfmacc_vf_f32m4(_sum, beta, _c, vl); + } + if (broadcast_type_C == 4) + { + const vfloat32m4_t _c = __riscv_vle32_v_f32m4(pC, vl); + if (beta == 1.f) + _sum = __riscv_vfadd_vv_f32m4(_sum, _c, vl); + else + _sum = __riscv_vfmacc_vf_f32m4(_sum, beta, _c, vl); + } + } + + __riscv_vse32_v_f32m4(out0, _sum, vl); + pp += vl; + outptr += out_hstep; + } +#else + for (; ii + 1 < max_ii; ii += 2) + { + float* out0 = outptr; + float* out1 = out0 + out_hstep; + const float* pC = C; + if (pC) + { + if (broadcast_type_C == 1 || broadcast_type_C == 2) + pC += i + ii; + if (broadcast_type_C == 3) + pC += (size_t)(i + ii) * c_hstep + j; + if (broadcast_type_C == 4) + pC += j; + } + + float c0 = 0.f; + float c1 = 0.f; + if (pC && (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2)) + { + c0 = pC[0] * beta; + c1 = pC[broadcast_type_C == 0 ? 0 : 1] * beta; + } + int jj = 0; + for (; jj + 3 < max_jj; jj += 4) + { + float sum00 = pp[0] * alpha; + float sum10 = pp[1] * alpha; + float sum01 = pp[2] * alpha; + float sum11 = pp[3] * alpha; + float sum02 = pp[4] * alpha; + float sum12 = pp[5] * alpha; + float sum03 = pp[6] * alpha; + float sum13 = pp[7] * alpha; + pp += 8; + + if (pC) + { + if (broadcast_type_C == 0) + { + sum00 += c0; + sum01 += c0; + sum02 += c0; + sum03 += c0; + sum10 += c1; + sum11 += c1; + sum12 += c1; + sum13 += c1; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + sum00 += c0; + sum01 += c0; + sum02 += c0; + sum03 += c0; + sum10 += c1; + sum11 += c1; + sum12 += c1; + sum13 += c1; + } + if (broadcast_type_C == 3) + { + sum00 += pC[0] * beta; + sum01 += pC[1] * beta; + sum02 += pC[2] * beta; + sum03 += pC[3] * beta; + sum10 += pC[c_hstep] * beta; + sum11 += pC[c_hstep + 1] * beta; + sum12 += pC[c_hstep + 2] * beta; + sum13 += pC[c_hstep + 3] * beta; + } + if (broadcast_type_C == 4) + { + sum00 += pC[0] * beta; + sum01 += pC[1] * beta; + sum02 += pC[2] * beta; + sum03 += pC[3] * beta; + sum10 += pC[0] * beta; + sum11 += pC[1] * beta; + sum12 += pC[2] * beta; + sum13 += pC[3] * beta; + } + } + + out0[0] = sum00; + out0[1] = sum01; + out0[2] = sum02; + out0[3] = sum03; + out1[0] = sum10; + out1[1] = sum11; + out1[2] = sum12; + out1[3] = sum13; + out0 += 4; + out1 += 4; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 4; + } + for (; jj + 1 < max_jj; jj += 2) + { + float sum00 = pp[0] * alpha; + float sum10 = pp[1] * alpha; + float sum01 = pp[2] * alpha; + float sum11 = pp[3] * alpha; + pp += 4; + + if (pC) + { + if (broadcast_type_C == 0) + { + sum00 += c0; + sum01 += c0; + sum10 += c1; + sum11 += c1; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + sum00 += c0; + sum01 += c0; + sum10 += c1; + sum11 += c1; + } + if (broadcast_type_C == 3) + { + sum00 += pC[0] * beta; + sum01 += pC[1] * beta; + sum10 += pC[c_hstep] * beta; + sum11 += pC[c_hstep + 1] * beta; + } + if (broadcast_type_C == 4) + { + sum00 += pC[0] * beta; + sum01 += pC[1] * beta; + sum10 += pC[0] * beta; + sum11 += pC[1] * beta; + } + } + + out0[0] = sum00; + out0[1] = sum01; + out1[0] = sum10; + out1[1] = sum11; + out0 += 2; + out1 += 2; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 2; + } + for (; jj < max_jj; jj++) + { + float sum0 = pp[0] * alpha; + float sum1 = pp[1] * alpha; + pp += 2; + if (pC) + { + if (broadcast_type_C == 0) + { + sum0 += c0; + sum1 += c1; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + sum0 += c0; + sum1 += c1; + } + if (broadcast_type_C == 3) + { + sum0 += pC[0] * beta; + sum1 += pC[c_hstep] * beta; + } + if (broadcast_type_C == 4) + { + sum0 += pC[0] * beta; + sum1 += pC[0] * beta; + } + } + out0[0] = sum0; + out1[0] = sum1; + out0++; + out1++; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC++; + } + outptr += out_hstep * 2; + } + for (; ii < max_ii; ii++) + { + float* out0 = outptr; + const float* pC = C; + if (pC) + { + if (broadcast_type_C == 1 || broadcast_type_C == 2) + pC += i + ii; + if (broadcast_type_C == 3) + pC += (size_t)(i + ii) * c_hstep + j; + if (broadcast_type_C == 4) + pC += j; + } + + float c0 = 0.f; + if (pC && (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2)) + c0 = pC[0] * beta; + int jj = 0; + for (; jj + 3 < max_jj; jj += 4) + { + float sum0 = pp[0] * alpha; + float sum1 = pp[1] * alpha; + float sum2 = pp[2] * alpha; + float sum3 = pp[3] * alpha; + pp += 4; + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + sum0 += c0; + sum1 += c0; + sum2 += c0; + sum3 += c0; + } + if (broadcast_type_C == 3) + { + sum0 += pC[0] * beta; + sum1 += pC[1] * beta; + sum2 += pC[2] * beta; + sum3 += pC[3] * beta; + } + if (broadcast_type_C == 4) + { + sum0 += pC[0] * beta; + sum1 += pC[1] * beta; + sum2 += pC[2] * beta; + sum3 += pC[3] * beta; + } + } + out0[0] = sum0; + out0[1] = sum1; + out0[2] = sum2; + out0[3] = sum3; + out0 += 4; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 4; + } + for (; jj + 1 < max_jj; jj += 2) + { + float sum0 = pp[0] * alpha; + float sum1 = pp[1] * alpha; + pp += 2; + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + sum0 += c0; + sum1 += c0; + } + if (broadcast_type_C == 3) + { + sum0 += pC[0] * beta; + sum1 += pC[1] * beta; + } + if (broadcast_type_C == 4) + { + sum0 += pC[0] * beta; + sum1 += pC[1] * beta; + } + } + out0[0] = sum0; + out0[1] = sum1; + out0 += 2; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 2; + } + for (; jj < max_jj; jj++) + { + float sum = *pp++ * alpha; + if (pC) + { + if (broadcast_type_C == 0) + sum += c0; + if (broadcast_type_C == 1 || broadcast_type_C == 2) + sum += c0; + if (broadcast_type_C == 3) + sum += pC[0] * beta; + if (broadcast_type_C == 4) + sum += pC[0] * beta; + } + out0[0] = sum; + out0++; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC++; + } + outptr += out_hstep; + } +#endif +} + +static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta) +{ + beta *= alpha; + (void)N; + const size_t c_hstep = C.dims == 3 ? C.cstep : (size_t)C.w; + const size_t out_hstep = top_blob.dims == 3 ? top_blob.cstep : (size_t)top_blob.w; + float* outptr = (float*)top_blob + (size_t)j * out_hstep + i; +#if __riscv_vector + const ptrdiff_t out_stride = (ptrdiff_t)out_hstep * sizeof(float); +#endif + + int ii = 0; +#if __riscv_vector + const int packn = csrr_vlenb() / 4; + const size_t vl_packn = __riscv_vsetvl_e32m1(packn); + const ptrdiff_t c_stride = (ptrdiff_t)c_hstep * sizeof(float); + for (; ii + (packn - 1) < max_ii; ii += packn) + { + const float* pC = C; + if (pC) + { + if (broadcast_type_C == 1 || broadcast_type_C == 2) + pC += i + ii; + if (broadcast_type_C == 3) + pC += (size_t)(i + ii) * c_hstep + j; + if (broadcast_type_C == 4) + pC += j; + } + + float c0 = 0.f; + vfloat32m1_t _c = __riscv_vfmv_v_f_f32m1(0.f, vl_packn); + if (pC && broadcast_type_C == 0) + c0 = pC[0] * beta; + if (pC && (broadcast_type_C == 1 || broadcast_type_C == 2)) + _c = __riscv_vfmul_vf_f32m1(__riscv_vle32_v_f32m1(pC, vl_packn), beta, vl_packn); + + for (int jj = 0; jj < max_jj; jj++) + { + vfloat32m1_t _sum = __riscv_vle32_v_f32m1(pp, vl_packn); + if (alpha != 1.f) + _sum = __riscv_vfmul_vf_f32m1(_sum, alpha, vl_packn); + if (pC) + { + if (broadcast_type_C == 0) + _sum = __riscv_vfadd_vf_f32m1(_sum, c0, vl_packn); + if (broadcast_type_C == 1 || broadcast_type_C == 2) + _sum = __riscv_vfadd_vv_f32m1(_sum, _c, vl_packn); + if (broadcast_type_C == 3) + { + const vfloat32m1_t _c0 = __riscv_vlse32_v_f32m1(pC + jj, c_stride, vl_packn); + if (beta == 1.f) + _sum = __riscv_vfadd_vv_f32m1(_sum, _c0, vl_packn); + else + _sum = __riscv_vfmacc_vf_f32m1(_sum, beta, _c0, vl_packn); + } + if (broadcast_type_C == 4) + _sum = __riscv_vfadd_vf_f32m1(_sum, pC[jj] * beta, vl_packn); + } + __riscv_vse32_v_f32m1(outptr + (size_t)jj * out_hstep, _sum, vl_packn); + pp += packn; + } + outptr += packn; + } + for (; ii + 1 < max_ii; ii += 2) + { + float* out0 = outptr; + const float* pC = C; + if (pC) + { + if (broadcast_type_C == 1 || broadcast_type_C == 2) + pC += i + ii; + if (broadcast_type_C == 3) + pC += (size_t)(i + ii) * c_hstep + j; + if (broadcast_type_C == 4) + pC += j; + } + + float c0 = 0.f; + float c1 = 0.f; + if (pC && (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2)) + { + c0 = pC[0] * beta; + c1 = pC[broadcast_type_C == 0 ? 0 : 1] * beta; + } + + const size_t vl = __riscv_vsetvl_e32m4(max_jj); + const vfloat32m4x2_t _s = __riscv_vlseg2e32_v_f32m4x2(pp, vl); + vfloat32m4_t _sum0 = __riscv_vget_v_f32m4x2_f32m4(_s, 0); + vfloat32m4_t _sum1 = __riscv_vget_v_f32m4x2_f32m4(_s, 1); + if (alpha != 1.f) + { + _sum0 = __riscv_vfmul_vf_f32m4(_sum0, alpha, vl); + _sum1 = __riscv_vfmul_vf_f32m4(_sum1, alpha, vl); + } + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _sum0 = __riscv_vfadd_vf_f32m4(_sum0, c0, vl); + _sum1 = __riscv_vfadd_vf_f32m4(_sum1, c1, vl); + } + if (broadcast_type_C == 3) + { + const vfloat32m4_t _c0 = __riscv_vle32_v_f32m4(pC, vl); + const vfloat32m4_t _c1 = __riscv_vle32_v_f32m4(pC + c_hstep, vl); + if (beta == 1.f) + { + _sum0 = __riscv_vfadd_vv_f32m4(_sum0, _c0, vl); + _sum1 = __riscv_vfadd_vv_f32m4(_sum1, _c1, vl); + } + else + { + _sum0 = __riscv_vfmacc_vf_f32m4(_sum0, beta, _c0, vl); + _sum1 = __riscv_vfmacc_vf_f32m4(_sum1, beta, _c1, vl); + } + } + if (broadcast_type_C == 4) + { + const vfloat32m4_t _c = __riscv_vle32_v_f32m4(pC, vl); + if (beta == 1.f) + { + _sum0 = __riscv_vfadd_vv_f32m4(_sum0, _c, vl); + _sum1 = __riscv_vfadd_vv_f32m4(_sum1, _c, vl); + } + else + { + _sum0 = __riscv_vfmacc_vf_f32m4(_sum0, beta, _c, vl); + _sum1 = __riscv_vfmacc_vf_f32m4(_sum1, beta, _c, vl); + } + } + } + const vfloat32m4x2_t _sum = __riscv_vcreate_v_f32m4x2(_sum0, _sum1); + if (out_hstep == 2) + __riscv_vsseg2e32_v_f32m4x2(out0, _sum, vl); + else + __riscv_vssseg2e32_v_f32m4x2(out0, out_stride, _sum, vl); + pp += vl * 2; + outptr += 2; + } + for (; ii < max_ii; ii++) + { + float* out0 = outptr; + const float* pC = C; + if (pC) + { + if (broadcast_type_C == 1 || broadcast_type_C == 2) + pC += i + ii; + if (broadcast_type_C == 3) + pC += (size_t)(i + ii) * c_hstep + j; + if (broadcast_type_C == 4) + pC += j; + } + + float c0 = 0.f; + if (pC && (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2)) + c0 = pC[0] * beta; + + const size_t vl = __riscv_vsetvl_e32m4(max_jj); + vfloat32m4_t _sum = __riscv_vle32_v_f32m4(pp, vl); + if (alpha != 1.f) + _sum = __riscv_vfmul_vf_f32m4(_sum, alpha, vl); + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + _sum = __riscv_vfadd_vf_f32m4(_sum, c0, vl); + if (broadcast_type_C == 3) + { + const vfloat32m4_t _c = __riscv_vle32_v_f32m4(pC, vl); + if (beta == 1.f) + _sum = __riscv_vfadd_vv_f32m4(_sum, _c, vl); + else + _sum = __riscv_vfmacc_vf_f32m4(_sum, beta, _c, vl); + } + if (broadcast_type_C == 4) + { + const vfloat32m4_t _c = __riscv_vle32_v_f32m4(pC, vl); + if (beta == 1.f) + _sum = __riscv_vfadd_vv_f32m4(_sum, _c, vl); + else + _sum = __riscv_vfmacc_vf_f32m4(_sum, beta, _c, vl); + } + } + if (out_hstep == 1) + __riscv_vse32_v_f32m4(out0, _sum, vl); + else + __riscv_vsse32_v_f32m4(out0, out_stride, _sum, vl); + pp += vl; + outptr++; + } +#else + for (; ii + 1 < max_ii; ii += 2) + { + float* out0 = outptr; + const float* pC = C; + if (pC) + { + if (broadcast_type_C == 1 || broadcast_type_C == 2) + pC += i + ii; + if (broadcast_type_C == 3) + pC += (size_t)(i + ii) * c_hstep + j; + if (broadcast_type_C == 4) + pC += j; + } + + float c0 = 0.f; + float c1 = 0.f; + if (pC && (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2)) + { + c0 = pC[0] * beta; + c1 = pC[broadcast_type_C == 0 ? 0 : 1] * beta; + } + int jj = 0; + for (; jj + 3 < max_jj; jj += 4) + { + float sum00 = pp[0] * alpha; + float sum10 = pp[1] * alpha; + float sum01 = pp[2] * alpha; + float sum11 = pp[3] * alpha; + float sum02 = pp[4] * alpha; + float sum12 = pp[5] * alpha; + float sum03 = pp[6] * alpha; + float sum13 = pp[7] * alpha; + pp += 8; + if (pC) + { + if (broadcast_type_C == 0) + { + sum00 += c0; + sum01 += c0; + sum02 += c0; + sum03 += c0; + sum10 += c1; + sum11 += c1; + sum12 += c1; + sum13 += c1; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + sum00 += c0; + sum01 += c0; + sum02 += c0; + sum03 += c0; + sum10 += c1; + sum11 += c1; + sum12 += c1; + sum13 += c1; + } + if (broadcast_type_C == 3) + { + sum00 += pC[0] * beta; + sum01 += pC[1] * beta; + sum02 += pC[2] * beta; + sum03 += pC[3] * beta; + sum10 += pC[c_hstep] * beta; + sum11 += pC[c_hstep + 1] * beta; + sum12 += pC[c_hstep + 2] * beta; + sum13 += pC[c_hstep + 3] * beta; + } + if (broadcast_type_C == 4) + { + sum00 += pC[0] * beta; + sum01 += pC[1] * beta; + sum02 += pC[2] * beta; + sum03 += pC[3] * beta; + sum10 += pC[0] * beta; + sum11 += pC[1] * beta; + sum12 += pC[2] * beta; + sum13 += pC[3] * beta; + } + } + out0[0] = sum00; + out0[1] = sum10; + out0[out_hstep] = sum01; + out0[out_hstep + 1] = sum11; + out0[out_hstep * 2] = sum02; + out0[out_hstep * 2 + 1] = sum12; + out0[out_hstep * 3] = sum03; + out0[out_hstep * 3 + 1] = sum13; + out0 += out_hstep * 4; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 4; + } + for (; jj + 1 < max_jj; jj += 2) + { + float sum00 = pp[0] * alpha; + float sum10 = pp[1] * alpha; + float sum01 = pp[2] * alpha; + float sum11 = pp[3] * alpha; + pp += 4; + if (pC) + { + if (broadcast_type_C == 0) + { + sum00 += c0; + sum01 += c0; + sum10 += c1; + sum11 += c1; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + sum00 += c0; + sum01 += c0; + sum10 += c1; + sum11 += c1; + } + if (broadcast_type_C == 3) + { + sum00 += pC[0] * beta; + sum01 += pC[1] * beta; + sum10 += pC[c_hstep] * beta; + sum11 += pC[c_hstep + 1] * beta; + } + if (broadcast_type_C == 4) + { + sum00 += pC[0] * beta; + sum01 += pC[1] * beta; + sum10 += pC[0] * beta; + sum11 += pC[1] * beta; + } + } + out0[0] = sum00; + out0[1] = sum10; + out0[out_hstep] = sum01; + out0[out_hstep + 1] = sum11; + out0 += out_hstep * 2; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 2; + } + for (; jj < max_jj; jj++) + { + float sum0 = pp[0] * alpha; + float sum1 = pp[1] * alpha; + pp += 2; + if (pC) + { + if (broadcast_type_C == 0) + { + sum0 += c0; + sum1 += c1; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + sum0 += c0; + sum1 += c1; + } + if (broadcast_type_C == 3) + { + sum0 += pC[0] * beta; + sum1 += pC[c_hstep] * beta; + } + if (broadcast_type_C == 4) + { + sum0 += pC[0] * beta; + sum1 += pC[0] * beta; + } + } + out0[0] = sum0; + out0[1] = sum1; + out0 += out_hstep; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC++; + } + outptr += 2; + } + for (; ii < max_ii; ii++) + { + float* out0 = outptr; + const float* pC = C; + if (pC) + { + if (broadcast_type_C == 1 || broadcast_type_C == 2) + pC += i + ii; + if (broadcast_type_C == 3) + pC += (size_t)(i + ii) * c_hstep + j; + if (broadcast_type_C == 4) + pC += j; + } + + float c0 = 0.f; + if (pC && (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2)) + c0 = pC[0] * beta; + int jj = 0; + for (; jj + 3 < max_jj; jj += 4) + { + float sum0 = pp[0] * alpha; + float sum1 = pp[1] * alpha; + float sum2 = pp[2] * alpha; + float sum3 = pp[3] * alpha; + pp += 4; + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + sum0 += c0; + sum1 += c0; + sum2 += c0; + sum3 += c0; + } + if (broadcast_type_C == 3) + { + sum0 += pC[0] * beta; + sum1 += pC[1] * beta; + sum2 += pC[2] * beta; + sum3 += pC[3] * beta; + } + if (broadcast_type_C == 4) + { + sum0 += pC[0] * beta; + sum1 += pC[1] * beta; + sum2 += pC[2] * beta; + sum3 += pC[3] * beta; + } + } + out0[0] = sum0; + out0[out_hstep] = sum1; + out0[out_hstep * 2] = sum2; + out0[out_hstep * 3] = sum3; + out0 += out_hstep * 4; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 4; + } + for (; jj + 1 < max_jj; jj += 2) + { + float sum0 = pp[0] * alpha; + float sum1 = pp[1] * alpha; + pp += 2; + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + sum0 += c0; + sum1 += c0; + } + if (broadcast_type_C == 3) + { + sum0 += pC[0] * beta; + sum1 += pC[1] * beta; + } + if (broadcast_type_C == 4) + { + sum0 += pC[0] * beta; + sum1 += pC[1] * beta; + } + } + out0[0] = sum0; + out0[out_hstep] = sum1; + out0 += out_hstep * 2; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += 2; + } + for (; jj < max_jj; jj++) + { + float sum = *pp++ * alpha; + if (pC) + { + if (broadcast_type_C == 0) + sum += c0; + if (broadcast_type_C == 1 || broadcast_type_C == 2) + sum += c0; + if (broadcast_type_C == 3) + sum += pC[0] * beta; + if (broadcast_type_C == 4) + sum += pC[0] * beta; + } + out0[0] = sum; + out0 += out_hstep; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC++; + } + outptr++; + } +#endif +} + +static void get_optimal_tile_mnk_wq_int8(int M, int N, int K, int constant_TILE_M, int constant_TILE_N, int constant_TILE_K, int& TILE_M, int& TILE_N, int& TILE_K, int nT) +{ +#if __riscv_vector + const int packm = std::max(8, csrr_vlenb() / 4); + const int packn = csrr_vlenb(); +#else + const int packm = 8; + const int packn = 4; +#endif + + TILE_M = packm; + TILE_N = packn; + TILE_K = K; + + // take constant TILE_M/N value when provided + if (constant_TILE_M > 0) + { + TILE_M = (constant_TILE_M + (packm - 1)) / packm * packm; + } + + if (constant_TILE_N > 0) + { + TILE_N = (constant_TILE_N + (packn - 1)) / packn * packn; + } + + // one driver tile follows the natural producer slab + TILE_M = std::min(TILE_M, packm); + TILE_N = std::min(TILE_N, packn); + + (void)M; + (void)N; + (void)constant_TILE_K; + (void)nT; +} diff --git a/src/layer/riscv/multiheadattention_riscv.cpp b/src/layer/riscv/multiheadattention_riscv.cpp new file mode 100644 index 000000000000..af9d641f66fb --- /dev/null +++ b/src/layer/riscv/multiheadattention_riscv.cpp @@ -0,0 +1,921 @@ +// Copyright 2026 Tencent +// SPDX-License-Identifier: BSD-3-Clause + +#include "multiheadattention_riscv.h" + +#include "layer_type.h" + +namespace ncnn { + +MultiHeadAttention_riscv::MultiHeadAttention_riscv() +{ +#if __riscv_vector + support_packing = true; +#endif // __riscv_vector + q_gemm = 0; + k_gemm = 0; + v_gemm = 0; + + qk_gemm = 0; + qkv_gemm = 0; + + qk_softmax = 0; + + o_gemm = 0; +} + +#if NCNN_WEIGHT_QUANT +int MultiHeadAttention_riscv::create_pipeline_wq_int8(const Option& _opt) +{ + if (q_gemm) + return 0; + + Option opt = _opt; + Option opt_wq = opt; + opt_wq.use_packing_layout = false; + opt_wq.use_fp16_packed = false; + opt_wq.use_fp16_storage = false; + opt_wq.use_fp16_arithmetic = false; + opt_wq.use_bf16_packed = false; + opt_wq.use_bf16_storage = false; + + { + qk_softmax = ncnn::create_layer_cpu(ncnn::LayerType::Softmax); + if (!qk_softmax) + return -100; + ncnn::ParamDict pd; + pd.set(0, -1); + pd.set(1, 1); + int ret = qk_softmax->load_param(pd); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = qk_softmax->load_model(ModelBinFromMatArray(0)); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = qk_softmax->create_pipeline(opt); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + } + + const int qdim = weight_data_size / embed_dim; + + { + q_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!q_gemm) + { + destroy_pipeline(_opt); + return -100; + } + ncnn::ParamDict pd; + pd.set(0, scale); + pd.set(1, 1.f); + pd.set(2, 0); // transA + pd.set(3, 1); // transB + pd.set(4, 0); // constantA + pd.set(5, 1); // constantB + pd.set(6, 1); // constantC + pd.set(7, 0); // M + pd.set(8, embed_dim); // N + pd.set(9, qdim); // K + pd.set(10, 4); // constant_broadcast_type_C + pd.set(11, 0); // output_N1M + pd.set(12, 0); // output_elempack + pd.set(14, 1); // output_transpose + pd.set(18, quantize_term); + int ret = q_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + Mat weights[4]; + weights[0] = q_weight_data; + weights[1] = q_bias_data; + weights[2] = q_weight_data_quantize_scales; + weights[3] = q_weight_data_input_scales; + ret = q_gemm->load_model(ModelBinFromMatArray(weights)); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = q_gemm->create_pipeline(opt_wq); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + } + + { + k_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!k_gemm) + { + destroy_pipeline(_opt); + return -100; + } + ncnn::ParamDict pd; + pd.set(2, 0); // transA + pd.set(3, 1); // transB + pd.set(4, 0); // constantA + pd.set(5, 1); // constantB + pd.set(6, 1); // constantC + pd.set(7, 0); // M + pd.set(8, embed_dim); // N + pd.set(9, kdim); // K + pd.set(10, 4); // constant_broadcast_type_C + pd.set(11, 0); // output_N1M + pd.set(12, 0); // output_elempack + pd.set(14, 1); // output_transpose + pd.set(18, quantize_term); + int ret = k_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + Mat weights[4]; + weights[0] = k_weight_data; + weights[1] = k_bias_data; + weights[2] = k_weight_data_quantize_scales; + weights[3] = k_weight_data_input_scales; + ret = k_gemm->load_model(ModelBinFromMatArray(weights)); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = k_gemm->create_pipeline(opt_wq); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + } + + { + v_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!v_gemm) + { + destroy_pipeline(_opt); + return -100; + } + ncnn::ParamDict pd; + pd.set(2, 0); // transA + pd.set(3, 1); // transB + pd.set(4, 0); // constantA + pd.set(5, 1); // constantB + pd.set(6, 1); // constantC + pd.set(7, 0); // M + pd.set(8, embed_dim); // N + pd.set(9, vdim); // K + pd.set(10, 4); // constant_broadcast_type_C + pd.set(11, 0); // output_N1M + pd.set(12, 0); // output_elempack + pd.set(14, 1); // output_transpose + pd.set(18, quantize_term); + int ret = v_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + Mat weights[4]; + weights[0] = v_weight_data; + weights[1] = v_bias_data; + weights[2] = v_weight_data_quantize_scales; + weights[3] = v_weight_data_input_scales; + ret = v_gemm->load_model(ModelBinFromMatArray(weights)); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = v_gemm->create_pipeline(opt_wq); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + } + + { + o_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!o_gemm) + { + destroy_pipeline(_opt); + return -100; + } + ncnn::ParamDict pd; + pd.set(2, 1); // transA + pd.set(3, 1); // transB + pd.set(4, 0); // constantA + pd.set(5, 1); // constantB + pd.set(6, 1); // constantC + pd.set(7, 0); // M = outch + pd.set(8, qdim); // N = size + pd.set(9, embed_dim); // K = maxk*inch + pd.set(10, 4); // constant_broadcast_type_C + pd.set(11, 0); // output_N1M + pd.set(18, quantize_term); + int ret = o_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + Mat weights[4]; + weights[0] = out_weight_data; + weights[1] = out_bias_data; + weights[2] = out_weight_data_quantize_scales; + weights[3] = out_weight_data_input_scales; + ret = o_gemm->load_model(ModelBinFromMatArray(weights)); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = o_gemm->create_pipeline(opt_wq); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + } + + { + qk_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!qk_gemm) + { + destroy_pipeline(_opt); + return -100; + } + ncnn::ParamDict pd; + pd.set(2, 1); // transA + pd.set(3, 0); // transB + pd.set(4, 0); // constantA + pd.set(5, 0); // constantB + pd.set(6, attn_mask ? 0 : 1); // constantC + pd.set(7, 0); // M + pd.set(8, 0); // N + pd.set(9, 0); // K + pd.set(10, attn_mask ? 3 : -1); // constant_broadcast_type_C + pd.set(11, 0); // output_N1M + pd.set(12, 1); // output_elempack + pd.set(13, 1); // output_elemtype = fp32 + int ret = qk_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = qk_gemm->load_model(ModelBinFromMatArray(0)); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + Option opt1 = opt; + opt1.use_bf16_packed = false; + opt1.use_bf16_storage = false; + opt1.num_threads = 1; + ret = qk_gemm->create_pipeline(opt1); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + } + + { + qkv_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!qkv_gemm) + { + destroy_pipeline(_opt); + return -100; + } + ncnn::ParamDict pd; + pd.set(2, 0); // transA + pd.set(3, 1); // transB + pd.set(4, 0); // constantA + pd.set(5, 0); // constantB + pd.set(6, 1); // constantC + pd.set(7, 0); // M + pd.set(8, 0); // N + pd.set(9, 0); // K + pd.set(10, -1); // constant_broadcast_type_C + pd.set(11, 0); // output_N1M + pd.set(12, 1); // output_elempack + pd.set(13, 1); // output_elemtype = fp32 + pd.set(14, 1); // output_transpose + int ret = qkv_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = qkv_gemm->load_model(ModelBinFromMatArray(0)); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + Option opt1 = opt; + opt1.use_bf16_packed = false; + opt1.use_bf16_storage = false; + opt1.num_threads = 1; + ret = qkv_gemm->create_pipeline(opt1); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + } + + q_weight_data.release(); + q_bias_data.release(); + k_weight_data.release(); + k_bias_data.release(); + v_weight_data.release(); + v_bias_data.release(); + out_weight_data.release(); + out_bias_data.release(); + q_weight_data_quantize_scales.release(); + k_weight_data_quantize_scales.release(); + v_weight_data_quantize_scales.release(); + out_weight_data_quantize_scales.release(); + q_weight_data_input_scales.release(); + k_weight_data_input_scales.release(); + v_weight_data_input_scales.release(); + out_weight_data_input_scales.release(); + + return 0; +} +#endif // NCNN_WEIGHT_QUANT + +int MultiHeadAttention_riscv::create_pipeline(const Option& _opt) +{ +#if NCNN_WEIGHT_QUANT + if (weight_block_quantize) + { + if (quantize_term / 100 != 8) + return MultiHeadAttention::create_pipeline(_opt); + + return create_pipeline_wq_int8(_opt); + } +#endif + + Option opt = _opt; + if (int8_scale_term) + { + support_packing = false; + support_bf16_storage = false; + + opt.use_packing_layout = false; // TODO enable packing + } + + { + qk_softmax = ncnn::create_layer_cpu(ncnn::LayerType::Softmax); + ncnn::ParamDict pd; + pd.set(0, -1); + pd.set(1, 1); + qk_softmax->load_param(pd); + qk_softmax->load_model(ModelBinFromMatArray(0)); + qk_softmax->create_pipeline(opt); + } + + const int qdim = weight_data_size / embed_dim; + + { + q_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + ncnn::ParamDict pd; + pd.set(0, scale); + pd.set(1, 1.f); + pd.set(2, 0); // transA + pd.set(3, 1); // transB + pd.set(4, 1); // constantA + pd.set(5, 0); // constantB + pd.set(6, 1); // constantC + pd.set(7, embed_dim); // M + pd.set(8, 0); // N + pd.set(9, qdim); // K + pd.set(10, 1); // constant_broadcast_type_C + pd.set(11, 0); // output_N1M + pd.set(12, 1); // output_elempack + pd.set(14, 0); // output_transpose +#if NCNN_INT8 + pd.set(18, int8_scale_term); +#endif + q_gemm->load_param(pd); + Mat weights[3]; + weights[0] = q_weight_data; + weights[1] = q_bias_data; +#if NCNN_INT8 + weights[2] = q_weight_data_int8_scales; +#endif + q_gemm->load_model(ModelBinFromMatArray(weights)); + q_gemm->create_pipeline(opt); + + if (opt.lightmode) + { + q_weight_data.release(); + q_bias_data.release(); + } + } + + { + k_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + ncnn::ParamDict pd; + pd.set(2, 0); // transA + pd.set(3, 1); // transB + pd.set(4, 1); // constantA + pd.set(5, 0); // constantB + pd.set(6, 1); // constantC + pd.set(7, embed_dim); // M + pd.set(8, 0); // N + pd.set(9, kdim); // K + pd.set(10, 1); // constant_broadcast_type_C + pd.set(11, 0); // output_N1M + pd.set(12, 1); // output_elempack + pd.set(14, 0); // output_transpose +#if NCNN_INT8 + pd.set(18, int8_scale_term); +#endif + k_gemm->load_param(pd); + Mat weights[3]; + weights[0] = k_weight_data; + weights[1] = k_bias_data; +#if NCNN_INT8 + weights[2] = k_weight_data_int8_scales; +#endif + k_gemm->load_model(ModelBinFromMatArray(weights)); + k_gemm->create_pipeline(opt); + + if (opt.lightmode) + { + k_weight_data.release(); + k_bias_data.release(); + } + } + + { + v_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + ncnn::ParamDict pd; + pd.set(2, 0); // transA + pd.set(3, 1); // transB + pd.set(4, 1); // constantA + pd.set(5, 0); // constantB + pd.set(6, 1); // constantC + pd.set(7, embed_dim); // M + pd.set(8, 0); // N + pd.set(9, vdim); // K + pd.set(10, 1); // constant_broadcast_type_C + pd.set(11, 0); // output_N1M + pd.set(12, 1); // output_elempack + pd.set(14, 0); // output_transpose +#if NCNN_INT8 + pd.set(18, int8_scale_term); +#endif + v_gemm->load_param(pd); + Mat weights[3]; + weights[0] = v_weight_data; + weights[1] = v_bias_data; +#if NCNN_INT8 + weights[2] = v_weight_data_int8_scales; +#endif + v_gemm->load_model(ModelBinFromMatArray(weights)); + v_gemm->create_pipeline(opt); + + if (opt.lightmode) + { + v_weight_data.release(); + v_bias_data.release(); + } + } + + { + o_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + ncnn::ParamDict pd; + pd.set(2, 1); // transA + pd.set(3, 1); // transB + pd.set(4, 0); // constantA + pd.set(5, 1); // constantB + pd.set(6, 1); // constantC + pd.set(7, 0); // M = outch + pd.set(8, qdim); // N = size + pd.set(9, embed_dim); // K = maxk*inch + pd.set(10, 4); // constant_broadcast_type_C + pd.set(11, 0); // output_N1M +#if NCNN_INT8 + pd.set(18, int8_scale_term); +#endif + o_gemm->load_param(pd); + Mat weights[3]; + weights[0] = out_weight_data; + weights[1] = out_bias_data; +#if NCNN_INT8 + Mat out_weight_data_int8_scales(1); + out_weight_data_int8_scales[0] = out_weight_data_int8_scale; + weights[2] = out_weight_data_int8_scales; +#endif + o_gemm->load_model(ModelBinFromMatArray(weights)); + Option opt_fp32 = opt; + opt_fp32.use_bf16_packed = false; + opt_fp32.use_bf16_storage = false; + o_gemm->create_pipeline(opt_fp32); + + if (opt.lightmode) + { + out_weight_data.release(); + out_bias_data.release(); + } + } + + { + qk_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + ncnn::ParamDict pd; + pd.set(2, 1); // transA + pd.set(3, 0); // transB + pd.set(4, 0); // constantA + pd.set(5, 0); // constantB + pd.set(6, attn_mask ? 0 : 1); // constantC + pd.set(7, 0); // M + pd.set(8, 0); // N + pd.set(9, 0); // K + pd.set(10, attn_mask ? 3 : -1); // constant_broadcast_type_C + pd.set(11, 0); // output_N1M + pd.set(12, 1); // output_elempack + pd.set(13, 1); // output_elemtype = fp32 +#if NCNN_INT8 + pd.set(18, int8_scale_term); +#endif + qk_gemm->load_param(pd); + qk_gemm->load_model(ModelBinFromMatArray(0)); + Option opt1 = opt; + opt1.use_bf16_packed = false; + opt1.use_bf16_storage = false; + opt1.num_threads = 1; + qk_gemm->create_pipeline(opt1); + } + + { + qkv_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + ncnn::ParamDict pd; + pd.set(2, 0); // transA + pd.set(3, 1); // transB + pd.set(4, 0); // constantA + pd.set(5, 0); // constantB + pd.set(6, 1); // constantC + pd.set(7, 0); // M + pd.set(8, 0); // N + pd.set(9, 0); // K + pd.set(10, -1); // constant_broadcast_type_C + pd.set(11, 0); // output_N1M + pd.set(12, 1); // output_elempack + pd.set(13, 1); // output_elemtype = fp32 + pd.set(14, 1); // output_transpose +#if NCNN_INT8 + pd.set(18, int8_scale_term); +#endif + qkv_gemm->load_param(pd); + qkv_gemm->load_model(ModelBinFromMatArray(0)); + Option opt1 = opt; + opt1.use_bf16_packed = false; + opt1.use_bf16_storage = false; + opt1.num_threads = 1; + qkv_gemm->create_pipeline(opt1); + } + + return 0; +} + +int MultiHeadAttention_riscv::destroy_pipeline(const Option& _opt) +{ + if (weight_block_quantize && quantize_term / 100 != 8) + return MultiHeadAttention::destroy_pipeline(_opt); + + Option opt = _opt; + if (int8_scale_term && !weight_block_quantize) + { + opt.use_packing_layout = false; // TODO enable packing + } + Option opt_wq = opt; + if (weight_block_quantize) + { + opt_wq.use_packing_layout = false; + opt_wq.use_fp16_packed = false; + opt_wq.use_fp16_storage = false; + opt_wq.use_fp16_arithmetic = false; + opt_wq.use_bf16_packed = false; + opt_wq.use_bf16_storage = false; + } + + if (qk_softmax) + { + qk_softmax->destroy_pipeline(opt); + delete qk_softmax; + qk_softmax = 0; + } + + if (q_gemm) + { + q_gemm->destroy_pipeline(opt_wq); + delete q_gemm; + q_gemm = 0; + } + + if (k_gemm) + { + k_gemm->destroy_pipeline(opt_wq); + delete k_gemm; + k_gemm = 0; + } + + if (v_gemm) + { + v_gemm->destroy_pipeline(opt_wq); + delete v_gemm; + v_gemm = 0; + } + + if (o_gemm) + { + o_gemm->destroy_pipeline(opt_wq); + delete o_gemm; + o_gemm = 0; + } + + if (qk_gemm) + { + qk_gemm->destroy_pipeline(opt); + delete qk_gemm; + qk_gemm = 0; + } + if (qkv_gemm) + { + qkv_gemm->destroy_pipeline(opt); + delete qkv_gemm; + qkv_gemm = 0; + } + + return 0; +} + +int MultiHeadAttention_riscv::forward(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& _opt) const +{ + if (weight_block_quantize && quantize_term / 100 != 8) + return MultiHeadAttention::forward(bottom_blobs, top_blobs, _opt); + + int q_blob_i = 0; + int k_blob_i = 0; + int v_blob_i = 0; + int attn_mask_i = 0; + int cached_xk_i = 0; + int cached_xv_i = 0; + resolve_bottom_blob_index((int)bottom_blobs.size(), q_blob_i, k_blob_i, v_blob_i, attn_mask_i, cached_xk_i, cached_xv_i); + + const Mat& q_blob = bottom_blobs[q_blob_i]; + const Mat& k_blob = bottom_blobs[k_blob_i]; + const Mat& v_blob = bottom_blobs[v_blob_i]; + const Mat& attn_mask_blob = attn_mask ? bottom_blobs[attn_mask_i] : Mat(); + const Mat& cached_xk_blob = kv_cache ? bottom_blobs[cached_xk_i] : Mat(); + const Mat& cached_xv_blob = kv_cache ? bottom_blobs[cached_xv_i] : Mat(); + + Option opt = _opt; + if (int8_scale_term && !weight_block_quantize) + { + opt.use_packing_layout = false; // TODO enable packing + } + Option opt_wq = opt; + if (weight_block_quantize) + { + opt_wq.use_packing_layout = false; + opt_wq.use_fp16_packed = false; + opt_wq.use_fp16_storage = false; + opt_wq.use_fp16_arithmetic = false; + opt_wq.use_bf16_packed = false; + opt_wq.use_bf16_storage = false; + } + + Mat attn_mask_blob_unpacked; + if (attn_mask && attn_mask_blob.elempack != 1) + { + convert_packing(attn_mask_blob, attn_mask_blob_unpacked, 1, opt); + if (attn_mask_blob_unpacked.empty()) + return -100; + } + else + { + attn_mask_blob_unpacked = attn_mask_blob; + } + + Mat cached_xk_blob_unpacked; + if (kv_cache && !cached_xk_blob.empty() && cached_xk_blob.elempack != 1) + { + convert_packing(cached_xk_blob, cached_xk_blob_unpacked, 1, opt); + if (cached_xk_blob_unpacked.empty()) + return -100; + } + else + { + cached_xk_blob_unpacked = cached_xk_blob; + } + + Mat cached_xv_blob_unpacked; + if (kv_cache && !cached_xv_blob.empty() && cached_xv_blob.elempack != 1) + { + convert_packing(cached_xv_blob, cached_xv_blob_unpacked, 1, opt); + if (cached_xv_blob_unpacked.empty()) + return -100; + } + else + { + cached_xv_blob_unpacked = cached_xv_blob; + } + + const int embed_dim_per_head = embed_dim / num_heads; + const int src_seqlen = q_blob.h * q_blob.elempack; + const int cur_seqlen = k_blob.h * k_blob.elempack; + const int past_seqlen = kv_cache && !cached_xk_blob_unpacked.empty() ? cached_xk_blob_unpacked.w : 0; + const int dst_seqlen = past_seqlen > 0 ? (q_blob_i == k_blob_i ? (past_seqlen + cur_seqlen) : past_seqlen) : cur_seqlen; + + Mat q_affine; + int retq = q_gemm->forward(q_blob, q_affine, opt_wq); + if (retq != 0) + return retq; + + Mat k_affine; + if (past_seqlen > 0) + { + if (q_blob_i == k_blob_i) + { + Mat k_affine_q; + int retk = k_gemm->forward(q_blob, k_affine_q, opt_wq); + if (retk != 0) + return retk; + + // assert dst_seqlen == cached_xk_blob_unpacked.w + k_affine_q.w + + // merge cached_xk_blob_unpacked and k_affine_q + k_affine.create(dst_seqlen, embed_dim, k_affine_q.elemsize); + if (k_affine.empty()) + return -100; + + for (int i = 0; i < embed_dim; i++) + { + const unsigned char* ptr = cached_xk_blob_unpacked.row(i); + const unsigned char* ptrq = k_affine_q.row(i); + unsigned char* outptr = k_affine.row(i); + + memcpy(outptr, ptr, past_seqlen * k_affine.elemsize); + memcpy(outptr + past_seqlen * k_affine.elemsize, ptrq, cur_seqlen * k_affine.elemsize); + } + } + else + { + k_affine = cached_xk_blob_unpacked; + } + } + else + { + int retk = k_gemm->forward(k_blob, k_affine, opt_wq); + if (retk != 0) + return retk; + } + + Mat qk_cross(dst_seqlen, src_seqlen * num_heads, 4u, opt.blob_allocator); + if (qk_cross.empty()) + return -100; + + std::vector retqks; + retqks.resize(num_heads); + #pragma omp parallel for num_threads(opt.num_threads) + for (int i = 0; i < num_heads; i++) + { + std::vector qk_bottom_blobs(2); + qk_bottom_blobs[0] = q_affine.row_range(i * embed_dim_per_head, embed_dim_per_head); + qk_bottom_blobs[1] = k_affine.row_range(i * embed_dim_per_head, embed_dim_per_head); + if (attn_mask) + { + const Mat& maskm = attn_mask_blob_unpacked.dims == 3 ? attn_mask_blob_unpacked.channel(i) : attn_mask_blob_unpacked; + qk_bottom_blobs.push_back(maskm); + } + std::vector qk_top_blobs(1); + qk_top_blobs[0] = qk_cross.row_range(i * src_seqlen, src_seqlen); + Option opt1 = opt; + opt1.num_threads = 1; + retqks[i] = qk_gemm->forward(qk_bottom_blobs, qk_top_blobs, opt1); + } + for (int i = 0; i < num_heads; i++) + { + if (retqks[i] != 0) + return retqks[i]; + } + + q_affine.release(); + + if (!kv_cache) + { + k_affine.release(); + } + + int retqk = qk_softmax->forward_inplace(qk_cross, opt); + if (retqk != 0) + return retqk; + + Mat v_affine; + if (past_seqlen > 0) + { + if (q_blob_i == v_blob_i) + { + Mat v_affine_q; + int retk = v_gemm->forward(v_blob, v_affine_q, opt_wq); + if (retk != 0) + return retk; + + // assert dst_seqlen == cached_xv_blob_unpacked.w + v_affine_q.w + + // merge cached_xv_blob_unpacked and v_affine_q + v_affine.create(dst_seqlen, embed_dim, v_affine_q.elemsize); + if (v_affine.empty()) + return -100; + + for (int i = 0; i < embed_dim; i++) + { + const unsigned char* ptr = cached_xv_blob_unpacked.row(i); + const unsigned char* ptrq = v_affine_q.row(i); + unsigned char* outptr = v_affine.row(i); + + memcpy(outptr, ptr, past_seqlen * v_affine.elemsize); + memcpy(outptr + past_seqlen * v_affine.elemsize, ptrq, cur_seqlen * v_affine.elemsize); + } + } + else + { + v_affine = cached_xv_blob_unpacked; + } + } + else + { + int retv = v_gemm->forward(v_blob, v_affine, opt_wq); + if (retv != 0) + return retv; + } + + Mat v_affine_fp32 = v_affine; + Mat qkv_cross(src_seqlen, embed_dim_per_head * num_heads, 4u, opt.blob_allocator); + if (qkv_cross.empty()) + return -100; + + std::vector retqkvs; + retqkvs.resize(num_heads); + #pragma omp parallel for num_threads(opt.num_threads) + for (int i = 0; i < num_heads; i++) + { + std::vector qkv_bottom_blobs(2); + qkv_bottom_blobs[0] = qk_cross.row_range(i * src_seqlen, src_seqlen); + qkv_bottom_blobs[1] = v_affine_fp32.row_range(i * embed_dim_per_head, embed_dim_per_head); + std::vector qkv_top_blobs(1); + qkv_top_blobs[0] = qkv_cross.row_range(i * embed_dim_per_head, embed_dim_per_head); + Option opt1 = opt; + opt1.num_threads = 1; + retqkvs[i] = qkv_gemm->forward(qkv_bottom_blobs, qkv_top_blobs, opt1); + } + for (int i = 0; i < num_heads; i++) + { + if (retqkvs[i] != 0) + return retqkvs[i]; + } + + v_affine_fp32.release(); + + if (!kv_cache) + { + v_affine.release(); + } + + int reto = o_gemm->forward(qkv_cross, top_blobs[0], opt_wq); + if (reto != 0) + return reto; + + if (kv_cache) + { + // assert top_blobs.size() == 3 + top_blobs[1] = k_affine; + top_blobs[2] = v_affine; + } + + return 0; +} + +} // namespace ncnn + diff --git a/src/layer/riscv/multiheadattention_riscv.h b/src/layer/riscv/multiheadattention_riscv.h new file mode 100644 index 000000000000..728331d735d7 --- /dev/null +++ b/src/layer/riscv/multiheadattention_riscv.h @@ -0,0 +1,40 @@ +// Copyright 2026 Tencent +// SPDX-License-Identifier: BSD-3-Clause + +#ifndef LAYER_MULTIHEADATTENTION_RISCV_H +#define LAYER_MULTIHEADATTENTION_RISCV_H + +#include "multiheadattention.h" + +namespace ncnn { + +class MultiHeadAttention_riscv : public MultiHeadAttention +{ +public: + MultiHeadAttention_riscv(); + + virtual int create_pipeline(const Option& opt); + virtual int destroy_pipeline(const Option& opt); + + virtual int forward(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& opt) const; + +protected: +#if NCNN_WEIGHT_QUANT + int create_pipeline_wq_int8(const Option& opt); +#endif + +public: + Layer* q_gemm; + Layer* k_gemm; + Layer* v_gemm; + Layer* o_gemm; + + Layer* qk_gemm; + Layer* qkv_gemm; + + Layer* qk_softmax; +}; + +} // namespace ncnn + +#endif // LAYER_MULTIHEADATTENTION_RISCV_H diff --git a/src/layer/x86/gemm_wq_int8.h b/src/layer/x86/gemm_wq_int8.h new file mode 100644 index 000000000000..1a4e78d0e6cc --- /dev/null +++ b/src/layer/x86/gemm_wq_int8.h @@ -0,0 +1,8579 @@ +// Copyright 2026 Tencent +// SPDX-License-Identifier: BSD-3-Clause + +#if NCNN_RUNTIME_CPU && NCNN_AVX512VNNI && __AVX512F__ && !__AVX512VNNI__ +void pack_B_tile_wq_int8_avx512vnni(const Mat& B, const Mat& B_scales, unsigned char* pp, float* pd, int j, int max_jj, int K, int block_size); +void quantize_A_tile_wq_int8_avx512vnni(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr); +void transpose_quantize_A_tile_wq_int8_avx512vnni(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr); +void gemm_transB_packed_tile_wq_int8_avx512vnni(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int block_size); +void unpack_output_tile_wq_int8_avx512vnni(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta); +void transpose_unpack_output_tile_wq_int8_avx512vnni(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta); +#endif + +#if NCNN_RUNTIME_CPU && NCNN_AVXVNNIINT8 && __AVX__ && !__AVXVNNIINT8__ && !__AVX512VNNI__ +void pack_B_tile_wq_int8_avxvnniint8(const Mat& B, const Mat& B_scales, unsigned char* pp, float* pd, int j, int max_jj, int K, int block_size); +void quantize_A_tile_wq_int8_avxvnniint8(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr); +void transpose_quantize_A_tile_wq_int8_avxvnniint8(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr); +void gemm_transB_packed_tile_wq_int8_avxvnniint8(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int block_size); +void unpack_output_tile_wq_int8_avxvnniint8(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta); +void transpose_unpack_output_tile_wq_int8_avxvnniint8(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta); +#endif + +#if NCNN_RUNTIME_CPU && NCNN_AVXVNNI && __AVX__ && !__AVXVNNI__ && !__AVXVNNIINT8__ && !__AVX512VNNI__ +void pack_B_tile_wq_int8_avxvnni(const Mat& B, const Mat& B_scales, unsigned char* pp, float* pd, int j, int max_jj, int K, int block_size); +void quantize_A_tile_wq_int8_avxvnni(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr); +void transpose_quantize_A_tile_wq_int8_avxvnni(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr); +void gemm_transB_packed_tile_wq_int8_avxvnni(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int block_size); +void unpack_output_tile_wq_int8_avxvnni(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta); +void transpose_unpack_output_tile_wq_int8_avxvnni(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta); +#endif + +#if NCNN_RUNTIME_CPU && NCNN_AVX2 && __AVX__ && !__AVX2__ && !__AVXVNNI__ && !__AVXVNNIINT8__ && !__AVX512VNNI__ +void pack_B_tile_wq_int8_avx2(const Mat& B, const Mat& B_scales, unsigned char* pp, float* pd, int j, int max_jj, int K, int block_size); +void quantize_A_tile_wq_int8_avx2(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr); +void transpose_quantize_A_tile_wq_int8_avx2(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr); +void gemm_transB_packed_tile_wq_int8_avx2(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int block_size); +void unpack_output_tile_wq_int8_avx2(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta); +void transpose_unpack_output_tile_wq_int8_avx2(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta); +#endif + +#if NCNN_RUNTIME_CPU && NCNN_XOP && __SSE2__ && !__XOP__ && !__AVX2__ && !__AVXVNNI__ && !__AVXVNNIINT8__ && !__AVX512VNNI__ +void gemm_transB_packed_tile_wq_int8_xop(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int block_size); +#endif + +static void pack_B_tile_wq_int8(const Mat& B, const Mat& B_scales, unsigned char* pp, float* pd, int j, int max_jj, int K, int block_size) +{ +#if NCNN_RUNTIME_CPU && NCNN_AVX512VNNI && __AVX512F__ && !__AVX512VNNI__ + if (ncnn::cpu_support_x86_avx512_vnni()) + { + pack_B_tile_wq_int8_avx512vnni(B, B_scales, pp, pd, j, max_jj, K, block_size); + return; + } +#endif +#if NCNN_RUNTIME_CPU && NCNN_AVXVNNIINT8 && __AVX__ && !__AVXVNNIINT8__ && !__AVX512VNNI__ + if (ncnn::cpu_support_x86_avx_vnni_int8()) + { + pack_B_tile_wq_int8_avxvnniint8(B, B_scales, pp, pd, j, max_jj, K, block_size); + return; + } +#endif +#if NCNN_RUNTIME_CPU && NCNN_AVXVNNI && __AVX__ && !__AVXVNNI__ && !__AVXVNNIINT8__ && !__AVX512VNNI__ + if (ncnn::cpu_support_x86_avx_vnni()) + { + pack_B_tile_wq_int8_avxvnni(B, B_scales, pp, pd, j, max_jj, K, block_size); + return; + } +#endif +#if NCNN_RUNTIME_CPU && NCNN_AVX2 && __AVX__ && !__AVX2__ && !__AVXVNNI__ && !__AVXVNNIINT8__ && !__AVX512VNNI__ + if (ncnn::cpu_support_x86_avx2()) + { + pack_B_tile_wq_int8_avx2(B, B_scales, pp, pd, j, max_jj, K, block_size); + return; + } +#endif + + const int block_count = (K + block_size - 1) / block_size; + + int jj = 0; +#if __SSE2__ +#if defined(__x86_64__) || defined(_M_X64) +#if __AVX512F__ + for (; jj + 7 < max_jj; jj += 8) + { + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + + const signed char* p0 = B.row(j + jj) + k0; + __m256i _vindex = _mm256_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7); + _vindex = _mm256_mullo_epi32(_vindex, _mm256_set1_epi32(B.w)); + int kk = 0; +#if __AVX512VNNI__ || __AVXVNNI__ +#if __AVXVNNIINT8__ + for (; kk + 3 < max_kk; kk += 4) + { + const __m256i _p = _mm256_i32gather_epi32((const int*)p0, _vindex, sizeof(signed char)); + _mm256_storeu_si256((__m256i*)pp, _p); + pp += 32; + p0 += 4; + } +#else // __AVXVNNIINT8__ + const __m256i _v127 = _mm256_set1_epi8(127); + for (; kk + 3 < max_kk; kk += 4) + { + __m256i _p = _mm256_i32gather_epi32((const int*)p0, _vindex, sizeof(signed char)); + _p = _mm256_add_epi8(_p, _v127); + _mm256_storeu_si256((__m256i*)pp, _p); + pp += 32; + p0 += 4; + } +#endif // __AVXVNNIINT8__ +#endif // __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 1 < max_kk; kk += 2) + { + const __m128i _p = _mm256_comp_cvtepi32_epi16(_mm256_i32gather_epi32((const int*)p0, _vindex, sizeof(signed char))); + _mm_storeu_si128((__m128i*)pp, _p); + pp += 16; + p0 += 2; + } + for (; kk < max_kk; kk++) + { + const __m128i _p = _mm256_comp_cvtepi32_epi8(_mm256_i32gather_epi32((const int*)p0, _vindex, sizeof(signed char))); + _mm_storel_epi64((__m128i*)pp, _p); + pp += 8; + p0++; + } + + __m256i _sindex = _mm256_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7); + _sindex = _mm256_mullo_epi32(_sindex, _mm256_set1_epi32(B_scales.w)); + const __m256 _scale = _mm256_i32gather_ps(B_scales.row(j + jj) + g, _sindex, sizeof(float)); + _mm256_storeu_ps(pd, _mm256_div_ps(_mm256_set1_ps(1.f), _scale)); + pd += 8; + } + } +#endif // __AVX512F__ + for (; jj + 3 < max_jj; jj += 4) + { + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + + int kk = 0; +#if __AVX512VNNI__ || __AVXVNNI__ + // VNNI consumes one contiguous K4 dword per output lane. + for (; kk + 3 < max_kk; kk += 4) + { + const signed char* p0 = B.row(j + jj) + k0 + kk; + const signed char* p1 = B.row(j + jj + 1) + k0 + kk; + const signed char* p2 = B.row(j + jj + 2) + k0 + kk; + const signed char* p3 = B.row(j + jj + 3) + k0 + kk; + __m128i _p = _mm_setr_epi32(*(const int*)p0, *(const int*)p1, *(const int*)p2, *(const int*)p3); +#if !__AVXVNNIINT8__ + _p = _mm_add_epi8(_p, _mm_set1_epi8(127)); +#endif // __AVXVNNIINT8__ + _mm_storeu_si128((__m128i*)pp, _p); + pp += 16; + } +#else + // AVX2/SSE2 consumes two K2 vectors for each real K4 region. + for (; kk + 3 < max_kk; kk += 4) + { + const signed char* p0 = B.row(j + jj) + k0 + kk; + const signed char* p1 = B.row(j + jj + 1) + k0 + kk; + const signed char* p2 = B.row(j + jj + 2) + k0 + kk; + const signed char* p3 = B.row(j + jj + 3) + k0 + kk; + const __m128i _p01 = _mm_setr_epi16((short)*(const unsigned short*)p0, (short)*(const unsigned short*)p1, (short)*(const unsigned short*)p2, (short)*(const unsigned short*)p3, 0, 0, 0, 0); + const __m128i _p23 = _mm_setr_epi16((short)*(const unsigned short*)(p0 + 2), (short)*(const unsigned short*)(p1 + 2), (short)*(const unsigned short*)(p2 + 2), (short)*(const unsigned short*)(p3 + 2), 0, 0, 0, 0); + _mm_storel_epi64((__m128i*)pp, _p01); + _mm_storel_epi64((__m128i*)(pp + 8), _p23); + pp += 16; + } +#endif // __AVX512VNNI__ || __AVXVNNI__ + // K2/K1 are always signed and compact, including classic VNNI. + for (; kk + 1 < max_kk; kk += 2) + { + const signed char* p0 = B.row(j + jj) + k0 + kk; + const signed char* p1 = B.row(j + jj + 1) + k0 + kk; + const signed char* p2 = B.row(j + jj + 2) + k0 + kk; + const signed char* p3 = B.row(j + jj + 3) + k0 + kk; + const __m128i _p = _mm_setr_epi16((short)*(const unsigned short*)p0, (short)*(const unsigned short*)p1, (short)*(const unsigned short*)p2, (short)*(const unsigned short*)p3, 0, 0, 0, 0); + _mm_storel_epi64((__m128i*)pp, _p); + pp += 8; + } + for (; kk < max_kk; kk++) + { + pp[0] = (unsigned char)B.row(j + jj)[k0 + kk]; + pp[1] = (unsigned char)B.row(j + jj + 1)[k0 + kk]; + pp[2] = (unsigned char)B.row(j + jj + 2)[k0 + kk]; + pp[3] = (unsigned char)B.row(j + jj + 3)[k0 + kk]; + pp += 4; + } + + pd[0] = 1.f / B_scales.row(j + jj)[g]; + pd[1] = 1.f / B_scales.row(j + jj + 1)[g]; + pd[2] = 1.f / B_scales.row(j + jj + 2)[g]; + pd[3] = 1.f / B_scales.row(j + jj + 3)[g]; + pd += 4; + } + } +#endif // defined(__x86_64__) || defined(_M_X64) +#endif // __SSE2__ + for (; jj + 1 < max_jj; jj += 2) + { + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + + int kk = 0; +#if __AVX512VNNI__ || __AVXVNNI__ + // VNNI consumes one contiguous K4 dword per output lane. + for (; kk + 3 < max_kk; kk += 4) + { + const signed char* p0 = B.row(j + jj) + k0 + kk; + const signed char* p1 = B.row(j + jj + 1) + k0 + kk; + __m128i _p = _mm_setr_epi32(*(const int*)p0, *(const int*)p1, 0, 0); +#if !__AVXVNNIINT8__ + _p = _mm_add_epi8(_p, _mm_set1_epi8(127)); +#endif // __AVXVNNIINT8__ + _mm_storel_epi64((__m128i*)pp, _p); + pp += 8; + } +#else +#if __SSE2__ + // AVX2/SSE2 consumes two K2 vectors for each real K4 region. + for (; kk + 3 < max_kk; kk += 4) + { + const signed char* p0 = B.row(j + jj) + k0 + kk; + const signed char* p1 = B.row(j + jj + 1) + k0 + kk; + *(unsigned short*)pp = *(const unsigned short*)p0; + *(unsigned short*)(pp + 2) = *(const unsigned short*)p1; + *(unsigned short*)(pp + 4) = *(const unsigned short*)(p0 + 2); + *(unsigned short*)(pp + 6) = *(const unsigned short*)(p1 + 2); + pp += 8; + } +#endif // __SSE2__ +#endif // __AVX512VNNI__ || __AVXVNNI__ + // K2/K1 are always signed and compact, including classic VNNI. +#if __SSE2__ + for (; kk + 1 < max_kk; kk += 2) + { + const signed char* p0 = B.row(j + jj) + k0 + kk; + const signed char* p1 = B.row(j + jj + 1) + k0 + kk; + *(unsigned short*)pp = *(const unsigned short*)p0; + *(unsigned short*)(pp + 2) = *(const unsigned short*)p1; + pp += 4; + } +#endif // __SSE2__ + for (; kk < max_kk; kk++) + { + pp[0] = (unsigned char)B.row(j + jj)[k0 + kk]; + pp[1] = (unsigned char)B.row(j + jj + 1)[k0 + kk]; + pp += 2; + } + + pd[0] = 1.f / B_scales.row(j + jj)[g]; + pd[1] = 1.f / B_scales.row(j + jj + 1)[g]; + pd += 2; + } + } + for (; jj < max_jj; jj++) + { + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + + int kk = 0; +#if __AVX512VNNI__ || __AVXVNNI__ + // VNNI consumes one contiguous K4 dword per output lane. + for (; kk + 3 < max_kk; kk += 4) + { + const signed char* p0 = B.row(j + jj) + k0 + kk; +#if !__AVXVNNIINT8__ + __m128i _p = _mm_cvtsi32_si128(*(const int*)p0); + _p = _mm_add_epi8(_p, _mm_set1_epi8(127)); + *(int*)pp = _mm_cvtsi128_si32(_p); +#else // __AVXVNNIINT8__ + *(int*)pp = *(const int*)p0; +#endif // __AVXVNNIINT8__ + pp += 4; + } +#else + // AVX2/SSE2 consumes two K2 vectors for each real K4 region. + for (; kk + 3 < max_kk; kk += 4) + { + const signed char* p0 = B.row(j + jj) + k0 + kk; + *(int*)pp = *(const int*)p0; + pp += 4; + } +#endif // __AVX512VNNI__ || __AVXVNNI__ + // K2/K1 are always signed and compact, including classic VNNI. + for (; kk + 1 < max_kk; kk += 2) + { + const signed char* p0 = B.row(j + jj) + k0 + kk; + *(unsigned short*)pp = *(const unsigned short*)p0; + pp += 2; + } + for (; kk < max_kk; kk++) + { + *pp++ = (unsigned char)B.row(j + jj)[k0 + kk]; + } + + pd[0] = 1.f / B_scales.row(j + jj)[g]; + pd += 1; + } + } +} + +static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr) +{ +#if NCNN_RUNTIME_CPU && NCNN_AVX512VNNI && __AVX512F__ && !__AVX512VNNI__ + if (ncnn::cpu_support_x86_avx512_vnni()) + { + quantize_A_tile_wq_int8_avx512vnni(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); + return; + } +#endif +#if NCNN_RUNTIME_CPU && NCNN_AVXVNNIINT8 && __AVX__ && !__AVXVNNIINT8__ && !__AVX512VNNI__ + if (ncnn::cpu_support_x86_avx_vnni_int8()) + { + quantize_A_tile_wq_int8_avxvnniint8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); + return; + } +#endif +#if NCNN_RUNTIME_CPU && NCNN_AVXVNNI && __AVX__ && !__AVXVNNI__ && !__AVXVNNIINT8__ && !__AVX512VNNI__ + if (ncnn::cpu_support_x86_avx_vnni()) + { + quantize_A_tile_wq_int8_avxvnni(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); + return; + } +#endif +#if NCNN_RUNTIME_CPU && NCNN_AVX2 && __AVX__ && !__AVX2__ && !__AVXVNNI__ && !__AVXVNNIINT8__ && !__AVX512VNNI__ + if (ncnn::cpu_support_x86_avx2()) + { + quantize_A_tile_wq_int8_avx2(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); + return; + } +#endif + + signed char* outptr = AT_tile; + const int out_hstep = AT_tile.w; + float* descale_ptr = AT_descales_tile; + const size_t A_hstep = A.dims == 3 ? A.cstep : (size_t)A.w; + const int block_count = (K + block_size - 1) / block_size; + + int ii = 0; +#if __SSE2__ +#if __AVX2__ +#if __AVX512F__ + __m512i _vindex = _mm512_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15); + _vindex = _mm512_mullo_epi32(_vindex, _mm512_set1_epi32((int)A_hstep)); + for (; ii + 15 < max_ii; ii += 16) + { + signed char* outptr0 = outptr + ii * out_hstep; + float* descale_ptr0 = descale_ptr + ii * block_count; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + + const float* p0 = (const float*)A + (i + ii) * A_hstep + k0; + __m512 _absmax0 = _mm512_setzero_ps(); + __m512 _absmax1 = _mm512_setzero_ps(); + __m512 _absmax2 = _mm512_setzero_ps(); + __m512 _absmax3 = _mm512_setzero_ps(); + __m512 _absmax4 = _mm512_setzero_ps(); + __m512 _absmax5 = _mm512_setzero_ps(); + __m512 _absmax6 = _mm512_setzero_ps(); + __m512 _absmax7 = _mm512_setzero_ps(); + __m512 _absmax8 = _mm512_setzero_ps(); + __m512 _absmax9 = _mm512_setzero_ps(); + __m512 _absmaxa = _mm512_setzero_ps(); + __m512 _absmaxb = _mm512_setzero_ps(); + __m512 _absmaxc = _mm512_setzero_ps(); + __m512 _absmaxd = _mm512_setzero_ps(); + __m512 _absmaxe = _mm512_setzero_ps(); + __m512 _absmaxf = _mm512_setzero_ps(); + int kk_absmax = 0; + for (; kk_absmax + 15 < max_kk; kk_absmax += 16) + { + __m512 _p0 = _mm512_loadu_ps(p0 + kk_absmax); + __m512 _p1 = _mm512_loadu_ps(p0 + A_hstep + kk_absmax); + __m512 _p2 = _mm512_loadu_ps(p0 + A_hstep * 2 + kk_absmax); + __m512 _p3 = _mm512_loadu_ps(p0 + A_hstep * 3 + kk_absmax); + __m512 _p4 = _mm512_loadu_ps(p0 + A_hstep * 4 + kk_absmax); + __m512 _p5 = _mm512_loadu_ps(p0 + A_hstep * 5 + kk_absmax); + __m512 _p6 = _mm512_loadu_ps(p0 + A_hstep * 6 + kk_absmax); + __m512 _p7 = _mm512_loadu_ps(p0 + A_hstep * 7 + kk_absmax); + __m512 _p8 = _mm512_loadu_ps(p0 + A_hstep * 8 + kk_absmax); + __m512 _p9 = _mm512_loadu_ps(p0 + A_hstep * 9 + kk_absmax); + __m512 _pa = _mm512_loadu_ps(p0 + A_hstep * 10 + kk_absmax); + __m512 _pb = _mm512_loadu_ps(p0 + A_hstep * 11 + kk_absmax); + __m512 _pc = _mm512_loadu_ps(p0 + A_hstep * 12 + kk_absmax); + __m512 _pd = _mm512_loadu_ps(p0 + A_hstep * 13 + kk_absmax); + __m512 _pe = _mm512_loadu_ps(p0 + A_hstep * 14 + kk_absmax); + __m512 _pf = _mm512_loadu_ps(p0 + A_hstep * 15 + kk_absmax); + if (input_scale_ptr) + { + const __m512 _s = _mm512_loadu_ps(input_scale_ptr + k0 + kk_absmax); + _p0 = _mm512_mul_ps(_p0, _s); + _p1 = _mm512_mul_ps(_p1, _s); + _p2 = _mm512_mul_ps(_p2, _s); + _p3 = _mm512_mul_ps(_p3, _s); + _p4 = _mm512_mul_ps(_p4, _s); + _p5 = _mm512_mul_ps(_p5, _s); + _p6 = _mm512_mul_ps(_p6, _s); + _p7 = _mm512_mul_ps(_p7, _s); + _p8 = _mm512_mul_ps(_p8, _s); + _p9 = _mm512_mul_ps(_p9, _s); + _pa = _mm512_mul_ps(_pa, _s); + _pb = _mm512_mul_ps(_pb, _s); + _pc = _mm512_mul_ps(_pc, _s); + _pd = _mm512_mul_ps(_pd, _s); + _pe = _mm512_mul_ps(_pe, _s); + _pf = _mm512_mul_ps(_pf, _s); + } + _absmax0 = _mm512_max_ps(_absmax0, abs512_ps(_p0)); + _absmax1 = _mm512_max_ps(_absmax1, abs512_ps(_p1)); + _absmax2 = _mm512_max_ps(_absmax2, abs512_ps(_p2)); + _absmax3 = _mm512_max_ps(_absmax3, abs512_ps(_p3)); + _absmax4 = _mm512_max_ps(_absmax4, abs512_ps(_p4)); + _absmax5 = _mm512_max_ps(_absmax5, abs512_ps(_p5)); + _absmax6 = _mm512_max_ps(_absmax6, abs512_ps(_p6)); + _absmax7 = _mm512_max_ps(_absmax7, abs512_ps(_p7)); + _absmax8 = _mm512_max_ps(_absmax8, abs512_ps(_p8)); + _absmax9 = _mm512_max_ps(_absmax9, abs512_ps(_p9)); + _absmaxa = _mm512_max_ps(_absmaxa, abs512_ps(_pa)); + _absmaxb = _mm512_max_ps(_absmaxb, abs512_ps(_pb)); + _absmaxc = _mm512_max_ps(_absmaxc, abs512_ps(_pc)); + _absmaxd = _mm512_max_ps(_absmaxd, abs512_ps(_pd)); + _absmaxe = _mm512_max_ps(_absmaxe, abs512_ps(_pe)); + _absmaxf = _mm512_max_ps(_absmaxf, abs512_ps(_pf)); + } + + float absmax0 = _mm512_reduce_max_ps(_absmax0); + float absmax1 = _mm512_reduce_max_ps(_absmax1); + float absmax2 = _mm512_reduce_max_ps(_absmax2); + float absmax3 = _mm512_reduce_max_ps(_absmax3); + float absmax4 = _mm512_reduce_max_ps(_absmax4); + float absmax5 = _mm512_reduce_max_ps(_absmax5); + float absmax6 = _mm512_reduce_max_ps(_absmax6); + float absmax7 = _mm512_reduce_max_ps(_absmax7); + float absmax8 = _mm512_reduce_max_ps(_absmax8); + float absmax9 = _mm512_reduce_max_ps(_absmax9); + float absmaxa = _mm512_reduce_max_ps(_absmaxa); + float absmaxb = _mm512_reduce_max_ps(_absmaxb); + float absmaxc = _mm512_reduce_max_ps(_absmaxc); + float absmaxd = _mm512_reduce_max_ps(_absmaxd); + float absmaxe = _mm512_reduce_max_ps(_absmaxe); + float absmaxf = _mm512_reduce_max_ps(_absmaxf); + for (; kk_absmax + 3 < max_kk; kk_absmax += 4) + { + const __m128 _s = input_scale_ptr ? _mm_loadu_ps(input_scale_ptr + k0 + kk_absmax) : _mm_set1_ps(1.f); + absmax0 = std::max(absmax0, _mm_reduce_max_ps(abs_ps(_mm_mul_ps(_mm_loadu_ps(p0 + kk_absmax), _s)))); + absmax1 = std::max(absmax1, _mm_reduce_max_ps(abs_ps(_mm_mul_ps(_mm_loadu_ps(p0 + A_hstep + kk_absmax), _s)))); + absmax2 = std::max(absmax2, _mm_reduce_max_ps(abs_ps(_mm_mul_ps(_mm_loadu_ps(p0 + A_hstep * 2 + kk_absmax), _s)))); + absmax3 = std::max(absmax3, _mm_reduce_max_ps(abs_ps(_mm_mul_ps(_mm_loadu_ps(p0 + A_hstep * 3 + kk_absmax), _s)))); + absmax4 = std::max(absmax4, _mm_reduce_max_ps(abs_ps(_mm_mul_ps(_mm_loadu_ps(p0 + A_hstep * 4 + kk_absmax), _s)))); + absmax5 = std::max(absmax5, _mm_reduce_max_ps(abs_ps(_mm_mul_ps(_mm_loadu_ps(p0 + A_hstep * 5 + kk_absmax), _s)))); + absmax6 = std::max(absmax6, _mm_reduce_max_ps(abs_ps(_mm_mul_ps(_mm_loadu_ps(p0 + A_hstep * 6 + kk_absmax), _s)))); + absmax7 = std::max(absmax7, _mm_reduce_max_ps(abs_ps(_mm_mul_ps(_mm_loadu_ps(p0 + A_hstep * 7 + kk_absmax), _s)))); + absmax8 = std::max(absmax8, _mm_reduce_max_ps(abs_ps(_mm_mul_ps(_mm_loadu_ps(p0 + A_hstep * 8 + kk_absmax), _s)))); + absmax9 = std::max(absmax9, _mm_reduce_max_ps(abs_ps(_mm_mul_ps(_mm_loadu_ps(p0 + A_hstep * 9 + kk_absmax), _s)))); + absmaxa = std::max(absmaxa, _mm_reduce_max_ps(abs_ps(_mm_mul_ps(_mm_loadu_ps(p0 + A_hstep * 10 + kk_absmax), _s)))); + absmaxb = std::max(absmaxb, _mm_reduce_max_ps(abs_ps(_mm_mul_ps(_mm_loadu_ps(p0 + A_hstep * 11 + kk_absmax), _s)))); + absmaxc = std::max(absmaxc, _mm_reduce_max_ps(abs_ps(_mm_mul_ps(_mm_loadu_ps(p0 + A_hstep * 12 + kk_absmax), _s)))); + absmaxd = std::max(absmaxd, _mm_reduce_max_ps(abs_ps(_mm_mul_ps(_mm_loadu_ps(p0 + A_hstep * 13 + kk_absmax), _s)))); + absmaxe = std::max(absmaxe, _mm_reduce_max_ps(abs_ps(_mm_mul_ps(_mm_loadu_ps(p0 + A_hstep * 14 + kk_absmax), _s)))); + absmaxf = std::max(absmaxf, _mm_reduce_max_ps(abs_ps(_mm_mul_ps(_mm_loadu_ps(p0 + A_hstep * 15 + kk_absmax), _s)))); + } + for (; kk_absmax < max_kk; kk_absmax++) + { + const float s = input_scale_ptr ? input_scale_ptr[k0 + kk_absmax] : 1.f; + absmax0 = std::max(absmax0, fabsf(p0[kk_absmax] * s)); + absmax1 = std::max(absmax1, fabsf(p0[A_hstep + kk_absmax] * s)); + absmax2 = std::max(absmax2, fabsf(p0[A_hstep * 2 + kk_absmax] * s)); + absmax3 = std::max(absmax3, fabsf(p0[A_hstep * 3 + kk_absmax] * s)); + absmax4 = std::max(absmax4, fabsf(p0[A_hstep * 4 + kk_absmax] * s)); + absmax5 = std::max(absmax5, fabsf(p0[A_hstep * 5 + kk_absmax] * s)); + absmax6 = std::max(absmax6, fabsf(p0[A_hstep * 6 + kk_absmax] * s)); + absmax7 = std::max(absmax7, fabsf(p0[A_hstep * 7 + kk_absmax] * s)); + absmax8 = std::max(absmax8, fabsf(p0[A_hstep * 8 + kk_absmax] * s)); + absmax9 = std::max(absmax9, fabsf(p0[A_hstep * 9 + kk_absmax] * s)); + absmaxa = std::max(absmaxa, fabsf(p0[A_hstep * 10 + kk_absmax] * s)); + absmaxb = std::max(absmaxb, fabsf(p0[A_hstep * 11 + kk_absmax] * s)); + absmaxc = std::max(absmaxc, fabsf(p0[A_hstep * 12 + kk_absmax] * s)); + absmaxd = std::max(absmaxd, fabsf(p0[A_hstep * 13 + kk_absmax] * s)); + absmaxe = std::max(absmaxe, fabsf(p0[A_hstep * 14 + kk_absmax] * s)); + absmaxf = std::max(absmaxf, fabsf(p0[A_hstep * 15 + kk_absmax] * s)); + } + const __m512 _absmax = _mm512_setr_ps(absmax0, absmax1, absmax2, absmax3, absmax4, absmax5, absmax6, absmax7, absmax8, absmax9, absmaxa, absmaxb, absmaxc, absmaxd, absmaxe, absmaxf); + + const __m512 _descale = _mm512_div_ps(_absmax, _mm512_set1_ps(127.f)); + const __m256 _absmax0_fp32 = _mm512_castps512_ps256(_absmax); + const __m256 _absmax1_fp32 = _mm256_castpd_ps(_mm512_extractf64x4_pd(_mm512_castps_pd(_absmax), 1)); + const __m512d _absmax0_fp64 = _mm512_cvtps_pd(_absmax0_fp32); + const __m512d _absmax1_fp64 = _mm512_cvtps_pd(_absmax1_fp32); + const __mmask8 _nonzero0 = _mm512_cmp_pd_mask(_absmax0_fp64, _mm512_setzero_pd(), _CMP_NEQ_OQ); + const __mmask8 _nonzero1 = _mm512_cmp_pd_mask(_absmax1_fp64, _mm512_setzero_pd(), _CMP_NEQ_OQ); + const __m256 _scale0 = _mm512_cvtpd_ps(_mm512_maskz_div_pd(_nonzero0, _mm512_set1_pd(127.0), _absmax0_fp64)); + const __m256 _scale1 = _mm512_cvtpd_ps(_mm512_maskz_div_pd(_nonzero1, _mm512_set1_pd(127.0), _absmax1_fp64)); + const __m512 _scale = combine8x2_ps(_scale0, _scale1); + _mm512_storeu_ps(descale_ptr0 + g * 16, _descale); + +#if __AVX512VNNI__ + __m512i _w_shift = _mm512_setzero_si512(); +#endif +#if __AVX512VNNI__ + signed char* pp = outptr0 + (k0 + g * 4) * 16; +#else + signed char* pp = outptr0 + k0 * 16; +#endif + int kk = 0; + for (; kk + 15 < max_kk; kk += 16) + { + __m512 _p0 = _mm512_loadu_ps(p0 + kk); + __m512 _p1 = _mm512_loadu_ps(p0 + A_hstep + kk); + __m512 _p2 = _mm512_loadu_ps(p0 + A_hstep * 2 + kk); + __m512 _p3 = _mm512_loadu_ps(p0 + A_hstep * 3 + kk); + __m512 _p4 = _mm512_loadu_ps(p0 + A_hstep * 4 + kk); + __m512 _p5 = _mm512_loadu_ps(p0 + A_hstep * 5 + kk); + __m512 _p6 = _mm512_loadu_ps(p0 + A_hstep * 6 + kk); + __m512 _p7 = _mm512_loadu_ps(p0 + A_hstep * 7 + kk); + __m512 _p8 = _mm512_loadu_ps(p0 + A_hstep * 8 + kk); + __m512 _p9 = _mm512_loadu_ps(p0 + A_hstep * 9 + kk); + __m512 _pa = _mm512_loadu_ps(p0 + A_hstep * 10 + kk); + __m512 _pb = _mm512_loadu_ps(p0 + A_hstep * 11 + kk); + __m512 _pc = _mm512_loadu_ps(p0 + A_hstep * 12 + kk); + __m512 _pd = _mm512_loadu_ps(p0 + A_hstep * 13 + kk); + __m512 _pe = _mm512_loadu_ps(p0 + A_hstep * 14 + kk); + __m512 _pf = _mm512_loadu_ps(p0 + A_hstep * 15 + kk); + if (input_scale_ptr) + { + const __m512 _s = _mm512_loadu_ps(input_scale_ptr + k0 + kk); + _p0 = _mm512_mul_ps(_p0, _s); + _p1 = _mm512_mul_ps(_p1, _s); + _p2 = _mm512_mul_ps(_p2, _s); + _p3 = _mm512_mul_ps(_p3, _s); + _p4 = _mm512_mul_ps(_p4, _s); + _p5 = _mm512_mul_ps(_p5, _s); + _p6 = _mm512_mul_ps(_p6, _s); + _p7 = _mm512_mul_ps(_p7, _s); + _p8 = _mm512_mul_ps(_p8, _s); + _p9 = _mm512_mul_ps(_p9, _s); + _pa = _mm512_mul_ps(_pa, _s); + _pb = _mm512_mul_ps(_pb, _s); + _pc = _mm512_mul_ps(_pc, _s); + _pd = _mm512_mul_ps(_pd, _s); + _pe = _mm512_mul_ps(_pe, _s); + _pf = _mm512_mul_ps(_pf, _s); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+x"(_p0), "+x"(_p1), "+x"(_p2), "+x"(_p3), "+x"(_p4), "+x"(_p5), "+x"(_p6), "+x"(_p7)); + asm volatile("" : "+x"(_p8), "+x"(_p9), "+x"(_pa), "+x"(_pb), "+x"(_pc), "+x"(_pd), "+x"(_pe), "+x"(_pf)); +#else + volatile __m512 _p0_ordered = _p0; + volatile __m512 _p1_ordered = _p1; + volatile __m512 _p2_ordered = _p2; + volatile __m512 _p3_ordered = _p3; + volatile __m512 _p4_ordered = _p4; + volatile __m512 _p5_ordered = _p5; + volatile __m512 _p6_ordered = _p6; + volatile __m512 _p7_ordered = _p7; + volatile __m512 _p8_ordered = _p8; + volatile __m512 _p9_ordered = _p9; + volatile __m512 _pa_ordered = _pa; + volatile __m512 _pb_ordered = _pb; + volatile __m512 _pc_ordered = _pc; + volatile __m512 _pd_ordered = _pd; + volatile __m512 _pe_ordered = _pe; + volatile __m512 _pf_ordered = _pf; + _p0 = _p0_ordered; + _p1 = _p1_ordered; + _p2 = _p2_ordered; + _p3 = _p3_ordered; + _p4 = _p4_ordered; + _p5 = _p5_ordered; + _p6 = _p6_ordered; + _p7 = _p7_ordered; + _p8 = _p8_ordered; + _p9 = _p9_ordered; + _pa = _pa_ordered; + _pb = _pb_ordered; + _pc = _pc_ordered; + _pd = _pd_ordered; + _pe = _pe_ordered; + _pf = _pf_ordered; +#endif + } + transpose16x16_ps(_p0, _p1, _p2, _p3, _p4, _p5, _p6, _p7, _p8, _p9, _pa, _pb, _pc, _pd, _pe, _pf); + __m128i _q0 = float2int8_avx512(_mm512_mul_ps(_p0, _scale)); + __m128i _q1 = float2int8_avx512(_mm512_mul_ps(_p1, _scale)); + __m128i _q2 = float2int8_avx512(_mm512_mul_ps(_p2, _scale)); + __m128i _q3 = float2int8_avx512(_mm512_mul_ps(_p3, _scale)); +#if __AVX512VNNI__ + transpose16x4_epi8(_q0, _q1, _q2, _q3); + __m512i _q = combine4x4_epi32(_q0, _q1, _q2, _q3); + _mm512_storeu_si512((__m512i*)pp, _q); + _w_shift = _mm512_dpbusd_epi32(_w_shift, _mm512_set1_epi8(127), _q); +#else + _mm_storeu_si128((__m128i*)pp, _mm_unpacklo_epi8(_q0, _q1)); + _mm_storeu_si128((__m128i*)(pp + 16), _mm_unpackhi_epi8(_q0, _q1)); + _mm_storeu_si128((__m128i*)(pp + 32), _mm_unpacklo_epi8(_q2, _q3)); + _mm_storeu_si128((__m128i*)(pp + 48), _mm_unpackhi_epi8(_q2, _q3)); +#endif + pp += 64; + + _q0 = float2int8_avx512(_mm512_mul_ps(_p4, _scale)); + _q1 = float2int8_avx512(_mm512_mul_ps(_p5, _scale)); + _q2 = float2int8_avx512(_mm512_mul_ps(_p6, _scale)); + _q3 = float2int8_avx512(_mm512_mul_ps(_p7, _scale)); +#if __AVX512VNNI__ + transpose16x4_epi8(_q0, _q1, _q2, _q3); + _q = combine4x4_epi32(_q0, _q1, _q2, _q3); + _mm512_storeu_si512((__m512i*)pp, _q); + _w_shift = _mm512_dpbusd_epi32(_w_shift, _mm512_set1_epi8(127), _q); +#else + _mm_storeu_si128((__m128i*)pp, _mm_unpacklo_epi8(_q0, _q1)); + _mm_storeu_si128((__m128i*)(pp + 16), _mm_unpackhi_epi8(_q0, _q1)); + _mm_storeu_si128((__m128i*)(pp + 32), _mm_unpacklo_epi8(_q2, _q3)); + _mm_storeu_si128((__m128i*)(pp + 48), _mm_unpackhi_epi8(_q2, _q3)); +#endif + pp += 64; + + _q0 = float2int8_avx512(_mm512_mul_ps(_p8, _scale)); + _q1 = float2int8_avx512(_mm512_mul_ps(_p9, _scale)); + _q2 = float2int8_avx512(_mm512_mul_ps(_pa, _scale)); + _q3 = float2int8_avx512(_mm512_mul_ps(_pb, _scale)); +#if __AVX512VNNI__ + transpose16x4_epi8(_q0, _q1, _q2, _q3); + _q = combine4x4_epi32(_q0, _q1, _q2, _q3); + _mm512_storeu_si512((__m512i*)pp, _q); + _w_shift = _mm512_dpbusd_epi32(_w_shift, _mm512_set1_epi8(127), _q); +#else + _mm_storeu_si128((__m128i*)pp, _mm_unpacklo_epi8(_q0, _q1)); + _mm_storeu_si128((__m128i*)(pp + 16), _mm_unpackhi_epi8(_q0, _q1)); + _mm_storeu_si128((__m128i*)(pp + 32), _mm_unpacklo_epi8(_q2, _q3)); + _mm_storeu_si128((__m128i*)(pp + 48), _mm_unpackhi_epi8(_q2, _q3)); +#endif + pp += 64; + + _q0 = float2int8_avx512(_mm512_mul_ps(_pc, _scale)); + _q1 = float2int8_avx512(_mm512_mul_ps(_pd, _scale)); + _q2 = float2int8_avx512(_mm512_mul_ps(_pe, _scale)); + _q3 = float2int8_avx512(_mm512_mul_ps(_pf, _scale)); +#if __AVX512VNNI__ + transpose16x4_epi8(_q0, _q1, _q2, _q3); + _q = combine4x4_epi32(_q0, _q1, _q2, _q3); + _mm512_storeu_si512((__m512i*)pp, _q); + _w_shift = _mm512_dpbusd_epi32(_w_shift, _mm512_set1_epi8(127), _q); +#else + _mm_storeu_si128((__m128i*)pp, _mm_unpacklo_epi8(_q0, _q1)); + _mm_storeu_si128((__m128i*)(pp + 16), _mm_unpackhi_epi8(_q0, _q1)); + _mm_storeu_si128((__m128i*)(pp + 32), _mm_unpacklo_epi8(_q2, _q3)); + _mm_storeu_si128((__m128i*)(pp + 48), _mm_unpackhi_epi8(_q2, _q3)); +#endif + pp += 64; + } + for (; kk + 3 < max_kk; kk += 4) + { + __m128 _p0 = _mm_loadu_ps(p0 + kk); + __m128 _p1 = _mm_loadu_ps(p0 + A_hstep + kk); + __m128 _p2 = _mm_loadu_ps(p0 + A_hstep * 2 + kk); + __m128 _p3 = _mm_loadu_ps(p0 + A_hstep * 3 + kk); + __m128 _p4 = _mm_loadu_ps(p0 + A_hstep * 4 + kk); + __m128 _p5 = _mm_loadu_ps(p0 + A_hstep * 5 + kk); + __m128 _p6 = _mm_loadu_ps(p0 + A_hstep * 6 + kk); + __m128 _p7 = _mm_loadu_ps(p0 + A_hstep * 7 + kk); + __m128 _p8 = _mm_loadu_ps(p0 + A_hstep * 8 + kk); + __m128 _p9 = _mm_loadu_ps(p0 + A_hstep * 9 + kk); + __m128 _pa = _mm_loadu_ps(p0 + A_hstep * 10 + kk); + __m128 _pb = _mm_loadu_ps(p0 + A_hstep * 11 + kk); + __m128 _pc = _mm_loadu_ps(p0 + A_hstep * 12 + kk); + __m128 _pd = _mm_loadu_ps(p0 + A_hstep * 13 + kk); + __m128 _pe = _mm_loadu_ps(p0 + A_hstep * 14 + kk); + __m128 _pf = _mm_loadu_ps(p0 + A_hstep * 15 + kk); + if (input_scale_ptr) + { + const __m128 _s = _mm_loadu_ps(input_scale_ptr + k0 + kk); + _p0 = _mm_mul_ps(_p0, _s); + _p1 = _mm_mul_ps(_p1, _s); + _p2 = _mm_mul_ps(_p2, _s); + _p3 = _mm_mul_ps(_p3, _s); + _p4 = _mm_mul_ps(_p4, _s); + _p5 = _mm_mul_ps(_p5, _s); + _p6 = _mm_mul_ps(_p6, _s); + _p7 = _mm_mul_ps(_p7, _s); + _p8 = _mm_mul_ps(_p8, _s); + _p9 = _mm_mul_ps(_p9, _s); + _pa = _mm_mul_ps(_pa, _s); + _pb = _mm_mul_ps(_pb, _s); + _pc = _mm_mul_ps(_pc, _s); + _pd = _mm_mul_ps(_pd, _s); + _pe = _mm_mul_ps(_pe, _s); + _pf = _mm_mul_ps(_pf, _s); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+x"(_p0), "+x"(_p1), "+x"(_p2), "+x"(_p3), "+x"(_p4), "+x"(_p5), "+x"(_p6), "+x"(_p7)); + asm volatile("" : "+x"(_p8), "+x"(_p9), "+x"(_pa), "+x"(_pb), "+x"(_pc), "+x"(_pd), "+x"(_pe), "+x"(_pf)); +#else + volatile __m128 _p0_ordered = _p0; + volatile __m128 _p1_ordered = _p1; + volatile __m128 _p2_ordered = _p2; + volatile __m128 _p3_ordered = _p3; + volatile __m128 _p4_ordered = _p4; + volatile __m128 _p5_ordered = _p5; + volatile __m128 _p6_ordered = _p6; + volatile __m128 _p7_ordered = _p7; + volatile __m128 _p8_ordered = _p8; + volatile __m128 _p9_ordered = _p9; + volatile __m128 _pa_ordered = _pa; + volatile __m128 _pb_ordered = _pb; + volatile __m128 _pc_ordered = _pc; + volatile __m128 _pd_ordered = _pd; + volatile __m128 _pe_ordered = _pe; + volatile __m128 _pf_ordered = _pf; + _p0 = _p0_ordered; + _p1 = _p1_ordered; + _p2 = _p2_ordered; + _p3 = _p3_ordered; + _p4 = _p4_ordered; + _p5 = _p5_ordered; + _p6 = _p6_ordered; + _p7 = _p7_ordered; + _p8 = _p8_ordered; + _p9 = _p9_ordered; + _pa = _pa_ordered; + _pb = _pb_ordered; + _pc = _pc_ordered; + _pd = _pd_ordered; + _pe = _pe_ordered; + _pf = _pf_ordered; +#endif + } + __m512 _t0 = combine4x4_ps(_p0, _p4, _p8, _pc); + __m512 _t1 = combine4x4_ps(_p1, _p5, _p9, _pd); + __m512 _t2 = combine4x4_ps(_p2, _p6, _pa, _pe); + __m512 _t3 = combine4x4_ps(_p3, _p7, _pb, _pf); + __m512 _t4 = _mm512_unpacklo_ps(_t0, _t1); + __m512 _t5 = _mm512_unpackhi_ps(_t0, _t1); + __m512 _t6 = _mm512_unpacklo_ps(_t2, _t3); + __m512 _t7 = _mm512_unpackhi_ps(_t2, _t3); + _t0 = _mm512_castpd_ps(_mm512_unpacklo_pd(_mm512_castps_pd(_t4), _mm512_castps_pd(_t6))); + _t1 = _mm512_castpd_ps(_mm512_unpackhi_pd(_mm512_castps_pd(_t4), _mm512_castps_pd(_t6))); + _t2 = _mm512_castpd_ps(_mm512_unpacklo_pd(_mm512_castps_pd(_t5), _mm512_castps_pd(_t7))); + _t3 = _mm512_castpd_ps(_mm512_unpackhi_pd(_mm512_castps_pd(_t5), _mm512_castps_pd(_t7))); + __m128i _q0 = float2int8_avx512(_mm512_mul_ps(_t0, _scale)); + __m128i _q1 = float2int8_avx512(_mm512_mul_ps(_t1, _scale)); + __m128i _q2 = float2int8_avx512(_mm512_mul_ps(_t2, _scale)); + __m128i _q3 = float2int8_avx512(_mm512_mul_ps(_t3, _scale)); +#if __AVX512VNNI__ + transpose16x4_epi8(_q0, _q1, _q2, _q3); + const __m512i _q = combine4x4_epi32(_q0, _q1, _q2, _q3); + _mm512_storeu_si512((__m512i*)pp, _q); + _w_shift = _mm512_dpbusd_epi32(_w_shift, _mm512_set1_epi8(127), _q); +#else + _mm_storeu_si128((__m128i*)pp, _mm_unpacklo_epi8(_q0, _q1)); + _mm_storeu_si128((__m128i*)(pp + 16), _mm_unpackhi_epi8(_q0, _q1)); + _mm_storeu_si128((__m128i*)(pp + 32), _mm_unpacklo_epi8(_q2, _q3)); + _mm_storeu_si128((__m128i*)(pp + 48), _mm_unpackhi_epi8(_q2, _q3)); +#endif + pp += 64; + } +#if __AVX512VNNI__ + if (max_kk >= 4) + { + _mm512_storeu_si512((__m512i*)pp, _w_shift); + pp += 64; + } +#endif // __AVX512VNNI__ + for (; kk + 1 < max_kk; kk += 2) + { + __m512 _p0 = _mm512_i32gather_ps(_vindex, (const float*)A + (i + ii) * A_hstep + k0 + kk, sizeof(float)); + __m512 _p1 = _mm512_i32gather_ps(_vindex, (const float*)A + (i + ii) * A_hstep + k0 + kk + 1, sizeof(float)); + if (input_scale_ptr) + { + _p0 = _mm512_mul_ps(_p0, _mm512_set1_ps(input_scale_ptr[k0 + kk])); + _p1 = _mm512_mul_ps(_p1, _mm512_set1_ps(input_scale_ptr[k0 + kk + 1])); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+x"(_p0), "+x"(_p1)); +#else + volatile __m512 _p0_ordered = _p0; + volatile __m512 _p1_ordered = _p1; + _p0 = _p0_ordered; + _p1 = _p1_ordered; +#endif + } + const __m128i _q0 = float2int8_avx512(_mm512_mul_ps(_p0, _scale)); + const __m128i _q1 = float2int8_avx512(_mm512_mul_ps(_p1, _scale)); + _mm_storeu_si128((__m128i*)pp, _mm_unpacklo_epi8(_q0, _q1)); + _mm_storeu_si128((__m128i*)(pp + 16), _mm_unpackhi_epi8(_q0, _q1)); + pp += 32; + } + if (kk < max_kk) + { + __m512 _p = _mm512_i32gather_ps(_vindex, (const float*)A + (i + ii) * A_hstep + k0 + kk, sizeof(float)); + if (input_scale_ptr) + { + _p = _mm512_mul_ps(_p, _mm512_set1_ps(input_scale_ptr[k0 + kk])); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+x"(_p)); +#else + volatile __m512 _p_ordered = _p; + _p = _p_ordered; +#endif + } + _mm_storeu_si128((__m128i*)pp, float2int8_avx512(_mm512_mul_ps(_p, _scale))); + } + } + } +#endif // __AVX512F__ + for (; ii + 7 < max_ii; ii += 8) + { + signed char* outptr0 = outptr + ii * out_hstep; + float* descale_ptr0 = descale_ptr + ii * block_count; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + const float* p0 = (const float*)A + (i + ii) * A_hstep + k0; + + __m256 _absmax0 = _mm256_setzero_ps(); + __m256 _absmax1 = _mm256_setzero_ps(); + __m256 _absmax2 = _mm256_setzero_ps(); + __m256 _absmax3 = _mm256_setzero_ps(); + __m256 _absmax4 = _mm256_setzero_ps(); + __m256 _absmax5 = _mm256_setzero_ps(); + __m256 _absmax6 = _mm256_setzero_ps(); + __m256 _absmax7 = _mm256_setzero_ps(); + int kk = 0; + for (; kk + 7 < max_kk; kk += 8) + { + __m256 _p0 = _mm256_loadu_ps(p0 + kk); + __m256 _p1 = _mm256_loadu_ps(p0 + A_hstep + kk); + __m256 _p2 = _mm256_loadu_ps(p0 + A_hstep * 2 + kk); + __m256 _p3 = _mm256_loadu_ps(p0 + A_hstep * 3 + kk); + __m256 _p4 = _mm256_loadu_ps(p0 + A_hstep * 4 + kk); + __m256 _p5 = _mm256_loadu_ps(p0 + A_hstep * 5 + kk); + __m256 _p6 = _mm256_loadu_ps(p0 + A_hstep * 6 + kk); + __m256 _p7 = _mm256_loadu_ps(p0 + A_hstep * 7 + kk); + if (input_scale_ptr) + { + const __m256 _s = _mm256_loadu_ps(input_scale_ptr + k0 + kk); + _p0 = _mm256_mul_ps(_p0, _s); + _p1 = _mm256_mul_ps(_p1, _s); + _p2 = _mm256_mul_ps(_p2, _s); + _p3 = _mm256_mul_ps(_p3, _s); + _p4 = _mm256_mul_ps(_p4, _s); + _p5 = _mm256_mul_ps(_p5, _s); + _p6 = _mm256_mul_ps(_p6, _s); + _p7 = _mm256_mul_ps(_p7, _s); + } + _absmax0 = _mm256_max_ps(_absmax0, abs256_ps(_p0)); + _absmax1 = _mm256_max_ps(_absmax1, abs256_ps(_p1)); + _absmax2 = _mm256_max_ps(_absmax2, abs256_ps(_p2)); + _absmax3 = _mm256_max_ps(_absmax3, abs256_ps(_p3)); + _absmax4 = _mm256_max_ps(_absmax4, abs256_ps(_p4)); + _absmax5 = _mm256_max_ps(_absmax5, abs256_ps(_p5)); + _absmax6 = _mm256_max_ps(_absmax6, abs256_ps(_p6)); + _absmax7 = _mm256_max_ps(_absmax7, abs256_ps(_p7)); + } + + float absmax0 = _mm256_reduce_max_ps(_absmax0); + float absmax1 = _mm256_reduce_max_ps(_absmax1); + float absmax2 = _mm256_reduce_max_ps(_absmax2); + float absmax3 = _mm256_reduce_max_ps(_absmax3); + float absmax4 = _mm256_reduce_max_ps(_absmax4); + float absmax5 = _mm256_reduce_max_ps(_absmax5); + float absmax6 = _mm256_reduce_max_ps(_absmax6); + float absmax7 = _mm256_reduce_max_ps(_absmax7); + for (; kk < max_kk; kk++) + { + const float s = input_scale_ptr ? input_scale_ptr[k0 + kk] : 1.f; + absmax0 = std::max(absmax0, fabsf(p0[kk] * s)); + absmax1 = std::max(absmax1, fabsf(p0[A_hstep + kk] * s)); + absmax2 = std::max(absmax2, fabsf(p0[A_hstep * 2 + kk] * s)); + absmax3 = std::max(absmax3, fabsf(p0[A_hstep * 3 + kk] * s)); + absmax4 = std::max(absmax4, fabsf(p0[A_hstep * 4 + kk] * s)); + absmax5 = std::max(absmax5, fabsf(p0[A_hstep * 5 + kk] * s)); + absmax6 = std::max(absmax6, fabsf(p0[A_hstep * 6 + kk] * s)); + absmax7 = std::max(absmax7, fabsf(p0[A_hstep * 7 + kk] * s)); + } + + const __m256 _absmax = _mm256_setr_ps(absmax0, absmax1, absmax2, absmax3, absmax4, absmax5, absmax6, absmax7); + const __m256 _descale = _mm256_div_ps(_absmax, _mm256_set1_ps(127.f)); + const __m256 _nonzero = _mm256_cmp_ps(_absmax, _mm256_setzero_ps(), _CMP_NEQ_OQ); + const __m256 _absmax_nonzero = _mm256_blendv_ps(_mm256_set1_ps(1.f), _absmax, _nonzero); + const __m256d _absmax0_fp64 = _mm256_cvtps_pd(_mm256_castps256_ps128(_absmax_nonzero)); + const __m256d _absmax1_fp64 = _mm256_cvtps_pd(_mm256_extractf128_ps(_absmax_nonzero, 1)); + const __m128 _scale0 = _mm256_cvtpd_ps(_mm256_div_pd(_mm256_set1_pd(127.0), _absmax0_fp64)); + const __m128 _scale1 = _mm256_cvtpd_ps(_mm256_div_pd(_mm256_set1_pd(127.0), _absmax1_fp64)); + const __m256 _scale = _mm256_and_ps(combine4x2_ps(_scale0, _scale1), _nonzero); + _mm256_storeu_ps(descale_ptr0 + g * 8, _descale); + +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + __m256i _w_shift = _mm256_setzero_si256(); + signed char* pp = outptr0 + (k0 + g * 4) * 8; +#else + signed char* pp = outptr0 + k0 * 8; +#endif + kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __m128 _p0 = _mm_loadu_ps((const float*)A + (i + ii) * A_hstep + k0 + kk); + __m128 _p1 = _mm_loadu_ps((const float*)A + (i + ii + 1) * A_hstep + k0 + kk); + __m128 _p2 = _mm_loadu_ps((const float*)A + (i + ii + 2) * A_hstep + k0 + kk); + __m128 _p3 = _mm_loadu_ps((const float*)A + (i + ii + 3) * A_hstep + k0 + kk); + __m128 _p4 = _mm_loadu_ps((const float*)A + (i + ii + 4) * A_hstep + k0 + kk); + __m128 _p5 = _mm_loadu_ps((const float*)A + (i + ii + 5) * A_hstep + k0 + kk); + __m128 _p6 = _mm_loadu_ps((const float*)A + (i + ii + 6) * A_hstep + k0 + kk); + __m128 _p7 = _mm_loadu_ps((const float*)A + (i + ii + 7) * A_hstep + k0 + kk); + if (input_scale_ptr) + { + const __m128 _s = _mm_loadu_ps(input_scale_ptr + k0 + kk); + _p0 = _mm_mul_ps(_p0, _s); + _p1 = _mm_mul_ps(_p1, _s); + _p2 = _mm_mul_ps(_p2, _s); + _p3 = _mm_mul_ps(_p3, _s); + _p4 = _mm_mul_ps(_p4, _s); + _p5 = _mm_mul_ps(_p5, _s); + _p6 = _mm_mul_ps(_p6, _s); + _p7 = _mm_mul_ps(_p7, _s); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+x"(_p0), "+x"(_p1), "+x"(_p2), "+x"(_p3), "+x"(_p4), "+x"(_p5), "+x"(_p6), "+x"(_p7)); +#else + volatile __m128 _p0_ordered = _p0; + volatile __m128 _p1_ordered = _p1; + volatile __m128 _p2_ordered = _p2; + volatile __m128 _p3_ordered = _p3; + volatile __m128 _p4_ordered = _p4; + volatile __m128 _p5_ordered = _p5; + volatile __m128 _p6_ordered = _p6; + volatile __m128 _p7_ordered = _p7; + _p0 = _p0_ordered; + _p1 = _p1_ordered; + _p2 = _p2_ordered; + _p3 = _p3_ordered; + _p4 = _p4_ordered; + _p5 = _p5_ordered; + _p6 = _p6_ordered; + _p7 = _p7_ordered; +#endif + } + + __m256 _t0 = combine4x2_ps(_p0, _p4); + __m256 _t1 = combine4x2_ps(_p1, _p5); + __m256 _t2 = combine4x2_ps(_p2, _p6); + __m256 _t3 = combine4x2_ps(_p3, _p7); + __m256 _t4 = _mm256_unpacklo_ps(_t0, _t1); + __m256 _t5 = _mm256_unpackhi_ps(_t0, _t1); + __m256 _t6 = _mm256_unpacklo_ps(_t2, _t3); + __m256 _t7 = _mm256_unpackhi_ps(_t2, _t3); + _t0 = _mm256_castpd_ps(_mm256_unpacklo_pd(_mm256_castps_pd(_t4), _mm256_castps_pd(_t6))); + _t1 = _mm256_castpd_ps(_mm256_unpackhi_pd(_mm256_castps_pd(_t4), _mm256_castps_pd(_t6))); + _t2 = _mm256_castpd_ps(_mm256_unpacklo_pd(_mm256_castps_pd(_t5), _mm256_castps_pd(_t7))); + _t3 = _mm256_castpd_ps(_mm256_unpackhi_pd(_mm256_castps_pd(_t5), _mm256_castps_pd(_t7))); + _t0 = _mm256_mul_ps(_t0, _scale); + _t1 = _mm256_mul_ps(_t1, _scale); + _t2 = _mm256_mul_ps(_t2, _scale); + _t3 = _mm256_mul_ps(_t3, _scale); + + __m128i _q0 = float2int8_avx(_t0, _t2); + __m128i _q1 = float2int8_avx(_t1, _t3); + __m128i _q01 = _mm_unpacklo_epi8(_q0, _q1); + __m128i _q23 = _mm_unpackhi_epi8(_q0, _q1); +#if __AVX512VNNI__ || __AVXVNNI__ + _q0 = _mm_unpacklo_epi16(_q01, _q23); + _q1 = _mm_unpackhi_epi16(_q01, _q23); + const __m256i _q = combine4x2_epi32(_q0, _q1); + _mm256_storeu_si256((__m256i*)pp, _q); +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + _w_shift = _mm256_comp_dpbusd_epi32(_w_shift, _mm256_set1_epi8(127), _q); +#endif +#else + _mm_storeu_si128((__m128i*)pp, _q01); + _mm_storeu_si128((__m128i*)(pp + 16), _q23); +#endif // __AVX512VNNI__ || __AVXVNNI__ + pp += 32; + } +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + if (max_kk >= 4) + { + _mm256_storeu_si256((__m256i*)pp, _w_shift); + pp += 32; + } +#endif + for (; kk + 1 < max_kk; kk += 2) + { + __m256i _vindex = _mm256_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7); + _vindex = _mm256_mullo_epi32(_vindex, _mm256_set1_epi32((int)A_hstep)); + __m256 _p0 = _mm256_i32gather_ps((const float*)A + (i + ii) * A_hstep + k0 + kk, _vindex, sizeof(float)); + __m256 _p1 = _mm256_i32gather_ps((const float*)A + (i + ii) * A_hstep + k0 + kk + 1, _vindex, sizeof(float)); + if (input_scale_ptr) + { + _p0 = _mm256_mul_ps(_p0, _mm256_set1_ps(input_scale_ptr[k0 + kk])); + _p1 = _mm256_mul_ps(_p1, _mm256_set1_ps(input_scale_ptr[k0 + kk + 1])); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+x"(_p0), "+x"(_p1)); +#else + volatile __m256 _p0_ordered = _p0; + volatile __m256 _p1_ordered = _p1; + _p0 = _p0_ordered; + _p1 = _p1_ordered; +#endif + } + _p0 = _mm256_mul_ps(_p0, _scale); + _p1 = _mm256_mul_ps(_p1, _scale); + __m128i _q = float2int8_avx(_p0, _p1); + const __m128i _si = _mm_setr_epi8(0, 8, 1, 9, 2, 10, 3, 11, 4, 12, 5, 13, 6, 14, 7, 15); + _q = _mm_shuffle_epi8(_q, _si); + _mm_storeu_si128((__m128i*)pp, _q); + pp += 16; + } + for (; kk < max_kk; kk++) + { + __m256i _vindex = _mm256_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7); + _vindex = _mm256_mullo_epi32(_vindex, _mm256_set1_epi32((int)A_hstep)); + __m256 _p = _mm256_i32gather_ps((const float*)A + (i + ii) * A_hstep + k0 + kk, _vindex, sizeof(float)); + if (input_scale_ptr) + { + _p = _mm256_mul_ps(_p, _mm256_set1_ps(input_scale_ptr[k0 + kk])); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+x"(_p)); +#else + volatile __m256 _p_ordered = _p; + _p = _p_ordered; +#endif + } + *(int64_t*)pp = float2int8_avx(_mm256_mul_ps(_p, _scale)); + pp += 8; + } + } + } +#endif // __AVX2__ + for (; ii + 3 < max_ii; ii += 4) + { + signed char* outptr0 = outptr + ii * out_hstep; + float* descale_ptr0 = descale_ptr + ii * block_count; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + const float* p0 = (const float*)A + (i + ii) * A_hstep + k0; + + __m128 _absmax0 = _mm_setzero_ps(); + __m128 _absmax1 = _mm_setzero_ps(); + __m128 _absmax2 = _mm_setzero_ps(); + __m128 _absmax3 = _mm_setzero_ps(); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __m128 _p0 = _mm_loadu_ps(p0 + kk); + __m128 _p1 = _mm_loadu_ps(p0 + A_hstep + kk); + __m128 _p2 = _mm_loadu_ps(p0 + A_hstep * 2 + kk); + __m128 _p3 = _mm_loadu_ps(p0 + A_hstep * 3 + kk); + if (input_scale_ptr) + { + const __m128 _s = _mm_loadu_ps(input_scale_ptr + k0 + kk); + _p0 = _mm_mul_ps(_p0, _s); + _p1 = _mm_mul_ps(_p1, _s); + _p2 = _mm_mul_ps(_p2, _s); + _p3 = _mm_mul_ps(_p3, _s); + } + _absmax0 = _mm_max_ps(_absmax0, abs_ps(_p0)); + _absmax1 = _mm_max_ps(_absmax1, abs_ps(_p1)); + _absmax2 = _mm_max_ps(_absmax2, abs_ps(_p2)); + _absmax3 = _mm_max_ps(_absmax3, abs_ps(_p3)); + } + + float absmax0 = _mm_reduce_max_ps(_absmax0); + float absmax1 = _mm_reduce_max_ps(_absmax1); + float absmax2 = _mm_reduce_max_ps(_absmax2); + float absmax3 = _mm_reduce_max_ps(_absmax3); + for (; kk < max_kk; kk++) + { + const float s = input_scale_ptr ? input_scale_ptr[k0 + kk] : 1.f; + absmax0 = std::max(absmax0, fabsf(p0[kk] * s)); + absmax1 = std::max(absmax1, fabsf(p0[A_hstep + kk] * s)); + absmax2 = std::max(absmax2, fabsf(p0[A_hstep * 2 + kk] * s)); + absmax3 = std::max(absmax3, fabsf(p0[A_hstep * 3 + kk] * s)); + } + + const __m128 _absmax = _mm_setr_ps(absmax0, absmax1, absmax2, absmax3); + const __m128 _descale = _mm_div_ps(_absmax, _mm_set1_ps(127.f)); + const __m128 _nonzero = _mm_cmpneq_ps(_absmax, _mm_setzero_ps()); + const __m128 _absmax_nonzero = _mm_or_ps(_mm_and_ps(_absmax, _nonzero), _mm_andnot_ps(_nonzero, _mm_set1_ps(1.f))); + const __m128d _absmax01_fp64 = _mm_cvtps_pd(_absmax_nonzero); + const __m128d _absmax23_fp64 = _mm_cvtps_pd(_mm_movehl_ps(_absmax_nonzero, _absmax_nonzero)); + const __m128 _scale01 = _mm_cvtpd_ps(_mm_div_pd(_mm_set1_pd(127.0), _absmax01_fp64)); + const __m128 _scale23 = _mm_cvtpd_ps(_mm_div_pd(_mm_set1_pd(127.0), _absmax23_fp64)); + const __m128 _scale = _mm_and_ps(_mm_movelh_ps(_scale01, _scale23), _nonzero); + _mm_storeu_ps(descale_ptr0 + g * 4, _descale); + +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + __m128i _w_shift = _mm_setzero_si128(); + signed char* pp = outptr0 + (k0 + g * 4) * 4; +#else + signed char* pp = outptr0 + k0 * 4; +#endif + kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __m128 _p0 = _mm_loadu_ps(p0 + kk); + __m128 _p1 = _mm_loadu_ps(p0 + A_hstep + kk); + __m128 _p2 = _mm_loadu_ps(p0 + A_hstep * 2 + kk); + __m128 _p3 = _mm_loadu_ps(p0 + A_hstep * 3 + kk); + if (input_scale_ptr) + { + const __m128 _s = _mm_loadu_ps(input_scale_ptr + k0 + kk); + _p0 = _mm_mul_ps(_p0, _s); + _p1 = _mm_mul_ps(_p1, _s); + _p2 = _mm_mul_ps(_p2, _s); + _p3 = _mm_mul_ps(_p3, _s); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+x"(_p0), "+x"(_p1), "+x"(_p2), "+x"(_p3)); +#else + volatile __m128 _p0_ordered = _p0; + volatile __m128 _p1_ordered = _p1; + volatile __m128 _p2_ordered = _p2; + volatile __m128 _p3_ordered = _p3; + _p0 = _p0_ordered; + _p1 = _p1_ordered; + _p2 = _p2_ordered; + _p3 = _p3_ordered; +#endif + } + const __m128i _q0 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p0, _mm_shuffle_ps(_scale, _scale, _MM_SHUFFLE(0, 0, 0, 0))))); + const __m128i _q1 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p1, _mm_shuffle_ps(_scale, _scale, _MM_SHUFFLE(1, 1, 1, 1))))); + const __m128i _q2 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p2, _mm_shuffle_ps(_scale, _scale, _MM_SHUFFLE(2, 2, 2, 2))))); + const __m128i _q3 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p3, _mm_shuffle_ps(_scale, _scale, _MM_SHUFFLE(3, 3, 3, 3))))); +#if __AVX512VNNI__ || __AVXVNNI__ + const __m128i _q = _mm_unpacklo_epi64(_mm_unpacklo_epi32(_q0, _q1), _mm_unpacklo_epi32(_q2, _q3)); + _mm_storeu_si128((__m128i*)pp, _q); +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + _w_shift = _mm_comp_dpbusd_epi32(_w_shift, _mm_set1_epi8(127), _q); +#endif +#else + const __m128i _q01 = _mm_unpacklo_epi16(_q0, _q1); + const __m128i _q23 = _mm_unpacklo_epi16(_q2, _q3); + _mm_storeu_si128((__m128i*)pp, _mm_unpacklo_epi32(_q01, _q23)); +#endif // __AVX512VNNI__ || __AVXVNNI__ + pp += 16; + } +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + if (max_kk >= 4) + { + _mm_storeu_si128((__m128i*)pp, _w_shift); + pp += 16; + } +#endif + for (; kk + 1 < max_kk; kk += 2) + { + __m128 _p0 = _mm_setr_ps(p0[kk], p0[A_hstep + kk], p0[A_hstep * 2 + kk], p0[A_hstep * 3 + kk]); + __m128 _p1 = _mm_setr_ps(p0[kk + 1], p0[A_hstep + kk + 1], p0[A_hstep * 2 + kk + 1], p0[A_hstep * 3 + kk + 1]); + if (input_scale_ptr) + { + _p0 = _mm_mul_ps(_p0, _mm_set1_ps(input_scale_ptr[k0 + kk])); + _p1 = _mm_mul_ps(_p1, _mm_set1_ps(input_scale_ptr[k0 + kk + 1])); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+x"(_p0), "+x"(_p1)); +#else + volatile __m128 _p0_ordered = _p0; + volatile __m128 _p1_ordered = _p1; + _p0 = _p0_ordered; + _p1 = _p1_ordered; +#endif + } + const __m128i _q0 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p0, _scale))); + const __m128i _q1 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p1, _scale))); + _mm_storel_epi64((__m128i*)pp, _mm_unpacklo_epi8(_q0, _q1)); + pp += 8; + } + for (; kk < max_kk; kk++) + { + __m128 _p = _mm_setr_ps(p0[kk], p0[A_hstep + kk], p0[A_hstep * 2 + kk], p0[A_hstep * 3 + kk]); + if (input_scale_ptr) + { + _p = _mm_mul_ps(_p, _mm_set1_ps(input_scale_ptr[k0 + kk])); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+x"(_p)); +#else + volatile __m128 _p_ordered = _p; + _p = _p_ordered; +#endif + } + *(int*)pp = float2int8_sse(_mm_mul_ps(_p, _scale)); + pp += 4; + } + } + } +#endif // __SSE2__ + for (; ii + 1 < max_ii; ii += 2) + { +#if __SSE2__ + signed char* outptr0 = outptr + ii * out_hstep; + float* descale_ptr0 = descale_ptr + ii * block_count; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + const float* p0 = (const float*)A + (i + ii) * A_hstep + k0; + + __m128 _absmax0 = _mm_setzero_ps(); + __m128 _absmax1 = _mm_setzero_ps(); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __m128 _p0 = _mm_loadu_ps(p0 + kk); + __m128 _p1 = _mm_loadu_ps(p0 + A_hstep + kk); + if (input_scale_ptr) + { + const __m128 _s = _mm_loadu_ps(input_scale_ptr + k0 + kk); + _p0 = _mm_mul_ps(_p0, _s); + _p1 = _mm_mul_ps(_p1, _s); + } + _absmax0 = _mm_max_ps(_absmax0, abs_ps(_p0)); + _absmax1 = _mm_max_ps(_absmax1, abs_ps(_p1)); + } + + float absmax0 = _mm_reduce_max_ps(_absmax0); + float absmax1 = _mm_reduce_max_ps(_absmax1); + for (; kk < max_kk; kk++) + { + const float s = input_scale_ptr ? input_scale_ptr[k0 + kk] : 1.f; + absmax0 = std::max(absmax0, fabsf(p0[kk] * s)); + absmax1 = std::max(absmax1, fabsf(p0[A_hstep + kk] * s)); + } + + const __m128 _absmax = _mm_setr_ps(absmax0, absmax1, 0.f, 0.f); + const __m128 _descale = _mm_div_ps(_absmax, _mm_set1_ps(127.f)); + const __m128 _nonzero = _mm_cmpneq_ps(_absmax, _mm_setzero_ps()); + const __m128 _absmax_nonzero = _mm_or_ps(_mm_and_ps(_absmax, _nonzero), _mm_andnot_ps(_nonzero, _mm_set1_ps(1.f))); + const __m128 _scale = _mm_and_ps(_mm_cvtpd_ps(_mm_div_pd(_mm_set1_pd(127.0), _mm_cvtps_pd(_absmax_nonzero))), _nonzero); + _mm_storel_pi((__m64*)(descale_ptr0 + g * 2), _descale); + +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + __m128i _w_shift = _mm_setzero_si128(); + signed char* pp = outptr0 + (k0 + g * 4) * 2; +#else + signed char* pp = outptr0 + k0 * 2; +#endif + kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __m128 _p0 = _mm_loadu_ps(p0 + kk); + __m128 _p1 = _mm_loadu_ps(p0 + A_hstep + kk); + if (input_scale_ptr) + { + const __m128 _s = _mm_loadu_ps(input_scale_ptr + k0 + kk); + _p0 = _mm_mul_ps(_p0, _s); + _p1 = _mm_mul_ps(_p1, _s); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+x"(_p0), "+x"(_p1)); +#else + volatile __m128 _p0_ordered = _p0; + volatile __m128 _p1_ordered = _p1; + _p0 = _p0_ordered; + _p1 = _p1_ordered; +#endif + } + const __m128i _q0 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p0, _mm_shuffle_ps(_scale, _scale, _MM_SHUFFLE(0, 0, 0, 0))))); + const __m128i _q1 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p1, _mm_shuffle_ps(_scale, _scale, _MM_SHUFFLE(1, 1, 1, 1))))); +#if __AVX512VNNI__ || __AVXVNNI__ + const __m128i _q = _mm_unpacklo_epi32(_q0, _q1); + _mm_storel_epi64((__m128i*)pp, _q); +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + _w_shift = _mm_comp_dpbusd_epi32(_w_shift, _mm_set1_epi8(127), _q); +#endif +#else + _mm_storel_epi64((__m128i*)pp, _mm_unpacklo_epi16(_q0, _q1)); +#endif // __AVX512VNNI__ || __AVXVNNI__ + pp += 8; + } +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + if (max_kk >= 4) + { + _mm_storel_epi64((__m128i*)pp, _w_shift); + pp += 8; + } +#endif + for (; kk + 1 < max_kk; kk += 2) + { + __m128 _p0 = _mm_setr_ps(p0[kk], p0[A_hstep + kk], 0.f, 0.f); + __m128 _p1 = _mm_setr_ps(p0[kk + 1], p0[A_hstep + kk + 1], 0.f, 0.f); + if (input_scale_ptr) + { + _p0 = _mm_mul_ps(_p0, _mm_set1_ps(input_scale_ptr[k0 + kk])); + _p1 = _mm_mul_ps(_p1, _mm_set1_ps(input_scale_ptr[k0 + kk + 1])); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+x"(_p0), "+x"(_p1)); +#else + volatile __m128 _p0_ordered = _p0; + volatile __m128 _p1_ordered = _p1; + _p0 = _p0_ordered; + _p1 = _p1_ordered; +#endif + } + const __m128i _q0 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p0, _scale))); + const __m128i _q1 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p1, _scale))); + *(int*)pp = _mm_cvtsi128_si32(_mm_unpacklo_epi8(_q0, _q1)); + pp += 4; + } + for (; kk < max_kk; kk++) + { + __m128 _p = _mm_setr_ps(p0[kk], p0[A_hstep + kk], 0.f, 0.f); + if (input_scale_ptr) + { + _p = _mm_mul_ps(_p, _mm_set1_ps(input_scale_ptr[k0 + kk])); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+x"(_p)); +#else + volatile __m128 _p_ordered = _p; + _p = _p_ordered; +#endif + } + *(unsigned short*)pp = (unsigned short)float2int8_sse(_mm_mul_ps(_p, _scale)); + pp += 2; + } + } +#else + const float* p0 = (const float*)A + (i + ii) * A_hstep; + signed char* outptr0 = outptr + ii * out_hstep; + float* descale_ptr0 = descale_ptr + ii * block_count; + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + float absmax0 = 0.f; + float absmax1 = 0.f; + for (int kk = 0; kk < max_kk; kk++) + { + float v0 = p0[k0 + kk]; + float v1 = p0[A_hstep + k0 + kk]; + if (input_scale_ptr) + { + const float s = input_scale_ptr[k0 + kk]; + v0 *= s; + v1 *= s; + } + absmax0 = std::max(absmax0, (float)fabsf(v0)); + absmax1 = std::max(absmax1, (float)fabsf(v1)); + } + + float scale0 = 0.f; + float scale1 = 0.f; + if (absmax0 != 0.f) + { + volatile double scale_fp64 = 127.0 / (double)absmax0; + scale0 = (float)scale_fp64; + } + if (absmax1 != 0.f) + { + volatile double scale_fp64 = 127.0 / (double)absmax1; + scale1 = (float)scale_fp64; + } + descale_ptr0[g * 2] = absmax0 / 127.f; + descale_ptr0[g * 2 + 1] = absmax1 / 127.f; + + signed char* pp = outptr0 + k0 * 2; + for (int kk = 0; kk < max_kk; kk++) + { + float v0 = p0[k0 + kk]; + float v1 = p0[A_hstep + k0 + kk]; + if (input_scale_ptr) + { + const float s = input_scale_ptr[k0 + kk]; + v0 *= s; + v1 *= s; + volatile float v0_ordered = v0; + volatile float v1_ordered = v1; + v0 = v0_ordered; + v1 = v1_ordered; + } + pp[0] = float2int8(v0 * scale0); + pp[1] = float2int8(v1 * scale1); + pp += 2; + } + } +#endif // __SSE2__ + } + for (; ii < max_ii; ii++) + { + const float* ptrA = (const float*)A + (i + ii) * A_hstep; + signed char* outptr0 = outptr + ii * out_hstep; + float* descale_ptr0 = descale_ptr + ii * block_count; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + signed char* pp = outptr0 + k0 + g * 4; +#else + signed char* pp = outptr0 + k0; +#endif + float absmax = 0.f; + + int kk = 0; +#if __SSE2__ +#if __AVX__ +#if __AVX512F__ + __m512 _absmax512 = _mm512_setzero_ps(); + for (; kk + 15 < max_kk; kk += 16) + { + __m512 _p = _mm512_loadu_ps(ptrA + k0 + kk); + if (input_scale_ptr) + _p = _mm512_mul_ps(_p, _mm512_loadu_ps(input_scale_ptr + k0 + kk)); + _absmax512 = _mm512_max_ps(_absmax512, abs512_ps(_p)); + } + absmax = std::max(absmax, _mm512_comp_reduce_max_ps(_absmax512)); +#endif // __AVX512F__ + __m256 _absmax256 = _mm256_setzero_ps(); + for (; kk + 7 < max_kk; kk += 8) + { + __m256 _p = _mm256_loadu_ps(ptrA + k0 + kk); + if (input_scale_ptr) + _p = _mm256_mul_ps(_p, _mm256_loadu_ps(input_scale_ptr + k0 + kk)); + _absmax256 = _mm256_max_ps(_absmax256, abs256_ps(_p)); + } + absmax = std::max(absmax, _mm256_reduce_max_ps(_absmax256)); +#endif // __AVX__ + __m128 _absmax128 = _mm_setzero_ps(); + for (; kk + 3 < max_kk; kk += 4) + { + __m128 _p = _mm_loadu_ps(ptrA + k0 + kk); + if (input_scale_ptr) + _p = _mm_mul_ps(_p, _mm_loadu_ps(input_scale_ptr + k0 + kk)); + _absmax128 = _mm_max_ps(_absmax128, abs_ps(_p)); + } + absmax = std::max(absmax, _mm_reduce_max_ps(_absmax128)); +#endif // __SSE2__ + for (; kk < max_kk; kk++) + { + float v = ptrA[k0 + kk]; + if (input_scale_ptr) + v *= input_scale_ptr[k0 + kk]; + absmax = std::max(absmax, (float)fabsf(v)); + } + + if (absmax == 0.f) + { + descale_ptr0[g] = 0.f; +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + memset(pp, 0, max_kk >= 4 ? max_kk + 4 : max_kk); +#else + memset(pp, 0, max_kk); +#endif + continue; + } + +#if __SSE2__ + const float scale = (float)(127.0 / (double)absmax); +#else + volatile double scale_fp64 = 127.0 / (double)absmax; + const float scale = (float)scale_fp64; +#endif + descale_ptr0[g] = absmax / 127.f; +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + int w_shift = 0; +#endif + kk = 0; +#if __SSE2__ +#if __AVX__ +#if __AVX512F__ + const __m512 _scale512 = _mm512_set1_ps(scale); + for (; kk + 15 < max_kk; kk += 16) + { + __m512 _p = _mm512_loadu_ps(ptrA + k0 + kk); + if (input_scale_ptr) + { + _p = _mm512_mul_ps(_p, _mm512_loadu_ps(input_scale_ptr + k0 + kk)); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+x"(_p)); +#else + volatile __m512 _p_ordered = _p; + _p = _p_ordered; +#endif + } + const __m128i _q = float2int8_avx512(_mm512_mul_ps(_p, _scale512)); + _mm_storeu_si128((__m128i*)pp, _q); + pp += 16; +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + const __m256i _q16 = _mm256_cvtepi8_epi16(_q); + const __m256i _q32 = _mm256_madd_epi16(_q16, _mm256_set1_epi16(1)); + w_shift += _mm_reduce_add_epi32(_mm256_castsi256_si128(_q32)); + w_shift += _mm_reduce_add_epi32(_mm256_extracti128_si256(_q32, 1)); +#endif + } +#endif // __AVX512F__ + const __m256 _scale256 = _mm256_set1_ps(scale); + for (; kk + 7 < max_kk; kk += 8) + { + __m256 _p = _mm256_loadu_ps(ptrA + k0 + kk); + if (input_scale_ptr) + { + _p = _mm256_mul_ps(_p, _mm256_loadu_ps(input_scale_ptr + k0 + kk)); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+x"(_p)); +#else + volatile __m256 _p_ordered = _p; + _p = _p_ordered; +#endif + } + const int64_t q = float2int8_avx(_mm256_mul_ps(_p, _scale256)); + *(int64_t*)pp = q; + pp += 8; +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) +#if defined(__x86_64__) || defined(_M_X64) + const __m128i _q8 = _mm_cvtsi64_si128(q); +#else + const __m128i _q8 = _mm_loadl_epi64((const __m128i*)(pp - 8)); +#endif + const __m128i _q16 = _mm_unpacklo_epi8(_q8, _mm_cmpgt_epi8(_mm_setzero_si128(), _q8)); + w_shift += _mm_reduce_add_epi32(_mm_madd_epi16(_q16, _mm_set1_epi16(1))); +#endif + } +#endif // __AVX__ + const __m128 _scale128 = _mm_set1_ps(scale); + for (; kk + 3 < max_kk; kk += 4) + { + __m128 _p = _mm_loadu_ps(ptrA + k0 + kk); + if (input_scale_ptr) + { + _p = _mm_mul_ps(_p, _mm_loadu_ps(input_scale_ptr + k0 + kk)); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+x"(_p)); +#else + volatile __m128 _p_ordered = _p; + _p = _p_ordered; +#endif + } + const int32_t q = float2int8_sse(_mm_mul_ps(_p, _scale128)); + *(int32_t*)pp = q; + pp += 4; +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + const __m128i _q8 = _mm_cvtsi32_si128(q); + const __m128i _q16 = _mm_unpacklo_epi8(_q8, _mm_cmpgt_epi8(_mm_setzero_si128(), _q8)); + w_shift += _mm_reduce_add_epi32(_mm_madd_epi16(_q16, _mm_set1_epi16(1))); +#endif + } +#endif // __SSE2__ +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + if (max_kk >= 4) + { + ((int*)pp)[0] = w_shift * 127; + pp += 4; + } +#endif + for (; kk < max_kk; kk++) + { + float v = ptrA[k0 + kk]; + if (input_scale_ptr) + { + v *= input_scale_ptr[k0 + kk]; +#if NCNN_GNU_INLINE_ASM && __SSE2__ + asm volatile("" : "+x"(v)); +#else + volatile float v_ordered = v; + v = v_ordered; +#endif + } + *pp++ = float2int8(v * scale); + } + } + } +} + +static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr) +{ +#if NCNN_RUNTIME_CPU && NCNN_AVX512VNNI && __AVX512F__ && !__AVX512VNNI__ + if (ncnn::cpu_support_x86_avx512_vnni()) + { + transpose_quantize_A_tile_wq_int8_avx512vnni(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); + return; + } +#endif +#if NCNN_RUNTIME_CPU && NCNN_AVXVNNIINT8 && __AVX__ && !__AVXVNNIINT8__ && !__AVX512VNNI__ + if (ncnn::cpu_support_x86_avx_vnni_int8()) + { + transpose_quantize_A_tile_wq_int8_avxvnniint8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); + return; + } +#endif +#if NCNN_RUNTIME_CPU && NCNN_AVXVNNI && __AVX__ && !__AVXVNNI__ && !__AVXVNNIINT8__ && !__AVX512VNNI__ + if (ncnn::cpu_support_x86_avx_vnni()) + { + transpose_quantize_A_tile_wq_int8_avxvnni(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); + return; + } +#endif +#if NCNN_RUNTIME_CPU && NCNN_AVX2 && __AVX__ && !__AVX2__ && !__AVXVNNI__ && !__AVXVNNIINT8__ && !__AVX512VNNI__ + if (ncnn::cpu_support_x86_avx2()) + { + transpose_quantize_A_tile_wq_int8_avx2(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); + return; + } +#endif + + signed char* outptr = AT_tile; + const int out_hstep = AT_tile.w; + float* descale_ptr = AT_descales_tile; + const size_t A_hstep = A.dims == 3 ? A.cstep : (size_t)A.w; + const int block_count = (K + block_size - 1) / block_size; + + int ii = 0; +#if __SSE2__ +#if __AVX2__ +#if __AVX512F__ + for (; ii + 15 < max_ii; ii += 16) + { + signed char* outptr0 = outptr + ii * out_hstep; + float* descale_ptr0 = descale_ptr + ii * block_count; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + + __m512 _absmax = _mm512_setzero_ps(); + for (int kk = 0; kk < max_kk; kk++) + { + __m512 _p = _mm512_loadu_ps((const float*)A + (k0 + kk) * A_hstep + i + ii); + if (input_scale_ptr) + _p = _mm512_mul_ps(_p, _mm512_set1_ps(input_scale_ptr[k0 + kk])); + _absmax = _mm512_max_ps(_absmax, abs512_ps(_p)); + } + + const __m512 _descale = _mm512_div_ps(_absmax, _mm512_set1_ps(127.f)); + const __m256 _absmax0_fp32 = _mm512_castps512_ps256(_absmax); + const __m256 _absmax1_fp32 = _mm256_castpd_ps(_mm512_extractf64x4_pd(_mm512_castps_pd(_absmax), 1)); + const __m512d _absmax0_fp64 = _mm512_cvtps_pd(_absmax0_fp32); + const __m512d _absmax1_fp64 = _mm512_cvtps_pd(_absmax1_fp32); + const __mmask8 _nonzero0 = _mm512_cmp_pd_mask(_absmax0_fp64, _mm512_setzero_pd(), _CMP_NEQ_OQ); + const __mmask8 _nonzero1 = _mm512_cmp_pd_mask(_absmax1_fp64, _mm512_setzero_pd(), _CMP_NEQ_OQ); + const __m256 _scale0 = _mm512_cvtpd_ps(_mm512_maskz_div_pd(_nonzero0, _mm512_set1_pd(127.0), _absmax0_fp64)); + const __m256 _scale1 = _mm512_cvtpd_ps(_mm512_maskz_div_pd(_nonzero1, _mm512_set1_pd(127.0), _absmax1_fp64)); + const __m512 _scale = combine8x2_ps(_scale0, _scale1); + _mm512_storeu_ps(descale_ptr0 + g * 16, _descale); + +#if __AVX512VNNI__ + __m512i _w_shift = _mm512_setzero_si512(); +#endif +#if __AVX512VNNI__ + signed char* pp = outptr0 + (k0 + g * 4) * 16; +#else + signed char* pp = outptr0 + k0 * 16; +#endif + int kk = 0; +#if __AVX512VNNI__ + for (; kk + 3 < max_kk; kk += 4) + { + __m512 _p0 = _mm512_loadu_ps((const float*)A + (k0 + kk) * A_hstep + i + ii); + __m512 _p1 = _mm512_loadu_ps((const float*)A + (k0 + kk + 1) * A_hstep + i + ii); + __m512 _p2 = _mm512_loadu_ps((const float*)A + (k0 + kk + 2) * A_hstep + i + ii); + __m512 _p3 = _mm512_loadu_ps((const float*)A + (k0 + kk + 3) * A_hstep + i + ii); + if (input_scale_ptr) + { + _p0 = _mm512_mul_ps(_p0, _mm512_set1_ps(input_scale_ptr[k0 + kk])); + _p1 = _mm512_mul_ps(_p1, _mm512_set1_ps(input_scale_ptr[k0 + kk + 1])); + _p2 = _mm512_mul_ps(_p2, _mm512_set1_ps(input_scale_ptr[k0 + kk + 2])); + _p3 = _mm512_mul_ps(_p3, _mm512_set1_ps(input_scale_ptr[k0 + kk + 3])); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+x"(_p0), "+x"(_p1), "+x"(_p2), "+x"(_p3)); +#else + volatile __m512 _p0_ordered = _p0; + volatile __m512 _p1_ordered = _p1; + volatile __m512 _p2_ordered = _p2; + volatile __m512 _p3_ordered = _p3; + _p0 = _p0_ordered; + _p1 = _p1_ordered; + _p2 = _p2_ordered; + _p3 = _p3_ordered; +#endif + } + __m128i _q0 = float2int8_avx512(_mm512_mul_ps(_p0, _scale)); + __m128i _q1 = float2int8_avx512(_mm512_mul_ps(_p1, _scale)); + __m128i _q2 = float2int8_avx512(_mm512_mul_ps(_p2, _scale)); + __m128i _q3 = float2int8_avx512(_mm512_mul_ps(_p3, _scale)); + transpose16x4_epi8(_q0, _q1, _q2, _q3); + const __m512i _q = combine4x4_epi32(_q0, _q1, _q2, _q3); + _mm512_storeu_si512((__m512i*)pp, _q); + _w_shift = _mm512_dpbusd_epi32(_w_shift, _mm512_set1_epi8(127), _q); + pp += 64; + } +#endif // __AVX512VNNI__ +#if __AVX512VNNI__ + if (max_kk >= 4) + { + _mm512_storeu_si512((__m512i*)pp, _w_shift); + pp += 64; + } +#endif // __AVX512VNNI__ + for (; kk + 1 < max_kk; kk += 2) + { + __m512 _p0 = _mm512_loadu_ps((const float*)A + (k0 + kk) * A_hstep + i + ii); + __m512 _p1 = _mm512_loadu_ps((const float*)A + (k0 + kk + 1) * A_hstep + i + ii); + if (input_scale_ptr) + { + _p0 = _mm512_mul_ps(_p0, _mm512_set1_ps(input_scale_ptr[k0 + kk])); + _p1 = _mm512_mul_ps(_p1, _mm512_set1_ps(input_scale_ptr[k0 + kk + 1])); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+x"(_p0), "+x"(_p1)); +#else + volatile __m512 _p0_ordered = _p0; + volatile __m512 _p1_ordered = _p1; + _p0 = _p0_ordered; + _p1 = _p1_ordered; +#endif + } + const __m128i _q0 = float2int8_avx512(_mm512_mul_ps(_p0, _scale)); + const __m128i _q1 = float2int8_avx512(_mm512_mul_ps(_p1, _scale)); + _mm_storeu_si128((__m128i*)pp, _mm_unpacklo_epi8(_q0, _q1)); + _mm_storeu_si128((__m128i*)(pp + 16), _mm_unpackhi_epi8(_q0, _q1)); + pp += 32; + } + if (kk < max_kk) + { + __m512 _p = _mm512_loadu_ps((const float*)A + (k0 + kk) * A_hstep + i + ii); + if (input_scale_ptr) + { + _p = _mm512_mul_ps(_p, _mm512_set1_ps(input_scale_ptr[k0 + kk])); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+x"(_p)); +#else + volatile __m512 _p_ordered = _p; + _p = _p_ordered; +#endif + } + _mm_storeu_si128((__m128i*)pp, float2int8_avx512(_mm512_mul_ps(_p, _scale))); + } + } + } +#endif // __AVX512F__ + for (; ii + 7 < max_ii; ii += 8) + { + signed char* outptr0 = outptr + ii * out_hstep; + float* descale_ptr0 = descale_ptr + ii * block_count; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + + __m256 _absmax = _mm256_setzero_ps(); + for (int kk = 0; kk < max_kk; kk++) + { + __m256 _p = _mm256_loadu_ps((const float*)A + (k0 + kk) * A_hstep + i + ii); + if (input_scale_ptr) + _p = _mm256_mul_ps(_p, _mm256_set1_ps(input_scale_ptr[k0 + kk])); + _absmax = _mm256_max_ps(_absmax, abs256_ps(_p)); + } + + const __m256 _descale = _mm256_div_ps(_absmax, _mm256_set1_ps(127.f)); + const __m256 _nonzero = _mm256_cmp_ps(_absmax, _mm256_setzero_ps(), _CMP_NEQ_OQ); + const __m256 _absmax_nonzero = _mm256_blendv_ps(_mm256_set1_ps(1.f), _absmax, _nonzero); + const __m256d _absmax0_fp64 = _mm256_cvtps_pd(_mm256_castps256_ps128(_absmax_nonzero)); + const __m256d _absmax1_fp64 = _mm256_cvtps_pd(_mm256_extractf128_ps(_absmax_nonzero, 1)); + const __m128 _scale0 = _mm256_cvtpd_ps(_mm256_div_pd(_mm256_set1_pd(127.0), _absmax0_fp64)); + const __m128 _scale1 = _mm256_cvtpd_ps(_mm256_div_pd(_mm256_set1_pd(127.0), _absmax1_fp64)); + const __m256 _scale = _mm256_and_ps(combine4x2_ps(_scale0, _scale1), _nonzero); + _mm256_storeu_ps(descale_ptr0 + g * 8, _descale); + +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + __m256i _w_shift = _mm256_setzero_si256(); + signed char* pp = outptr0 + (k0 + g * 4) * 8; +#else + signed char* pp = outptr0 + k0 * 8; +#endif + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __m256 _p0 = _mm256_loadu_ps((const float*)A + (k0 + kk) * A_hstep + i + ii); + __m256 _p1 = _mm256_loadu_ps((const float*)A + (k0 + kk + 1) * A_hstep + i + ii); + __m256 _p2 = _mm256_loadu_ps((const float*)A + (k0 + kk + 2) * A_hstep + i + ii); + __m256 _p3 = _mm256_loadu_ps((const float*)A + (k0 + kk + 3) * A_hstep + i + ii); + if (input_scale_ptr) + { + _p0 = _mm256_mul_ps(_p0, _mm256_set1_ps(input_scale_ptr[k0 + kk])); + _p1 = _mm256_mul_ps(_p1, _mm256_set1_ps(input_scale_ptr[k0 + kk + 1])); + _p2 = _mm256_mul_ps(_p2, _mm256_set1_ps(input_scale_ptr[k0 + kk + 2])); + _p3 = _mm256_mul_ps(_p3, _mm256_set1_ps(input_scale_ptr[k0 + kk + 3])); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+x"(_p0), "+x"(_p1), "+x"(_p2), "+x"(_p3)); +#else + volatile __m256 _p0_ordered = _p0; + volatile __m256 _p1_ordered = _p1; + volatile __m256 _p2_ordered = _p2; + volatile __m256 _p3_ordered = _p3; + _p0 = _p0_ordered; + _p1 = _p1_ordered; + _p2 = _p2_ordered; + _p3 = _p3_ordered; +#endif + } + _p0 = _mm256_mul_ps(_p0, _scale); + _p1 = _mm256_mul_ps(_p1, _scale); + _p2 = _mm256_mul_ps(_p2, _scale); + _p3 = _mm256_mul_ps(_p3, _scale); + + __m128i _q0 = float2int8_avx(_p0, _p2); + __m128i _q1 = float2int8_avx(_p1, _p3); + __m128i _q01 = _mm_unpacklo_epi8(_q0, _q1); + __m128i _q23 = _mm_unpackhi_epi8(_q0, _q1); +#if __AVX512VNNI__ || __AVXVNNI__ + _q0 = _mm_unpacklo_epi16(_q01, _q23); + _q1 = _mm_unpackhi_epi16(_q01, _q23); + const __m256i _q = combine4x2_epi32(_q0, _q1); + _mm256_storeu_si256((__m256i*)pp, _q); +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + _w_shift = _mm256_comp_dpbusd_epi32(_w_shift, _mm256_set1_epi8(127), _q); +#endif +#else + _mm_storeu_si128((__m128i*)pp, _q01); + _mm_storeu_si128((__m128i*)(pp + 16), _q23); +#endif // __AVX512VNNI__ || __AVXVNNI__ + pp += 32; + } +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + if (max_kk >= 4) + { + _mm256_storeu_si256((__m256i*)pp, _w_shift); + pp += 32; + } +#endif + for (; kk + 1 < max_kk; kk += 2) + { + __m256 _p0 = _mm256_loadu_ps((const float*)A + (k0 + kk) * A_hstep + i + ii); + __m256 _p1 = _mm256_loadu_ps((const float*)A + (k0 + kk + 1) * A_hstep + i + ii); + if (input_scale_ptr) + { + _p0 = _mm256_mul_ps(_p0, _mm256_set1_ps(input_scale_ptr[k0 + kk])); + _p1 = _mm256_mul_ps(_p1, _mm256_set1_ps(input_scale_ptr[k0 + kk + 1])); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+x"(_p0), "+x"(_p1)); +#else + volatile __m256 _p0_ordered = _p0; + volatile __m256 _p1_ordered = _p1; + _p0 = _p0_ordered; + _p1 = _p1_ordered; +#endif + } + _p0 = _mm256_mul_ps(_p0, _scale); + _p1 = _mm256_mul_ps(_p1, _scale); + __m128i _q = float2int8_avx(_p0, _p1); + const __m128i _si = _mm_setr_epi8(0, 8, 1, 9, 2, 10, 3, 11, 4, 12, 5, 13, 6, 14, 7, 15); + _q = _mm_shuffle_epi8(_q, _si); + _mm_storeu_si128((__m128i*)pp, _q); + pp += 16; + } + for (; kk < max_kk; kk++) + { + __m256 _p = _mm256_loadu_ps((const float*)A + (k0 + kk) * A_hstep + i + ii); + if (input_scale_ptr) + { + _p = _mm256_mul_ps(_p, _mm256_set1_ps(input_scale_ptr[k0 + kk])); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+x"(_p)); +#else + volatile __m256 _p_ordered = _p; + _p = _p_ordered; +#endif + } + *(int64_t*)pp = float2int8_avx(_mm256_mul_ps(_p, _scale)); + pp += 8; + } + } + } +#endif // __AVX2__ + for (; ii + 3 < max_ii; ii += 4) + { + signed char* outptr0 = outptr + ii * out_hstep; + float* descale_ptr0 = descale_ptr + ii * block_count; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + + __m128 _absmax = _mm_setzero_ps(); + for (int kk = 0; kk < max_kk; kk++) + { + __m128 _p = _mm_loadu_ps((const float*)A + (k0 + kk) * A_hstep + i + ii); + if (input_scale_ptr) + _p = _mm_mul_ps(_p, _mm_set1_ps(input_scale_ptr[k0 + kk])); + _absmax = _mm_max_ps(_absmax, abs_ps(_p)); + } + + const __m128 _descale = _mm_div_ps(_absmax, _mm_set1_ps(127.f)); + const __m128 _nonzero = _mm_cmpneq_ps(_absmax, _mm_setzero_ps()); + const __m128 _absmax_nonzero = _mm_or_ps(_mm_and_ps(_absmax, _nonzero), _mm_andnot_ps(_nonzero, _mm_set1_ps(1.f))); + const __m128d _absmax01_fp64 = _mm_cvtps_pd(_absmax_nonzero); + const __m128d _absmax23_fp64 = _mm_cvtps_pd(_mm_movehl_ps(_absmax_nonzero, _absmax_nonzero)); + const __m128 _scale01 = _mm_cvtpd_ps(_mm_div_pd(_mm_set1_pd(127.0), _absmax01_fp64)); + const __m128 _scale23 = _mm_cvtpd_ps(_mm_div_pd(_mm_set1_pd(127.0), _absmax23_fp64)); + const __m128 _scale = _mm_and_ps(_mm_movelh_ps(_scale01, _scale23), _nonzero); + _mm_storeu_ps(descale_ptr0 + g * 4, _descale); + +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + __m128i _w_shift = _mm_setzero_si128(); + signed char* pp = outptr0 + (k0 + g * 4) * 4; +#else + signed char* pp = outptr0 + k0 * 4; +#endif + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __m128 _p0 = _mm_loadu_ps((const float*)A + (k0 + kk) * A_hstep + i + ii); + __m128 _p1 = _mm_loadu_ps((const float*)A + (k0 + kk + 1) * A_hstep + i + ii); + __m128 _p2 = _mm_loadu_ps((const float*)A + (k0 + kk + 2) * A_hstep + i + ii); + __m128 _p3 = _mm_loadu_ps((const float*)A + (k0 + kk + 3) * A_hstep + i + ii); + if (input_scale_ptr) + { + _p0 = _mm_mul_ps(_p0, _mm_set1_ps(input_scale_ptr[k0 + kk])); + _p1 = _mm_mul_ps(_p1, _mm_set1_ps(input_scale_ptr[k0 + kk + 1])); + _p2 = _mm_mul_ps(_p2, _mm_set1_ps(input_scale_ptr[k0 + kk + 2])); + _p3 = _mm_mul_ps(_p3, _mm_set1_ps(input_scale_ptr[k0 + kk + 3])); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+x"(_p0), "+x"(_p1), "+x"(_p2), "+x"(_p3)); +#else + volatile __m128 _p0_ordered = _p0; + volatile __m128 _p1_ordered = _p1; + volatile __m128 _p2_ordered = _p2; + volatile __m128 _p3_ordered = _p3; + _p0 = _p0_ordered; + _p1 = _p1_ordered; + _p2 = _p2_ordered; + _p3 = _p3_ordered; +#endif + } + const __m128i _q0 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p0, _scale))); + const __m128i _q1 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p1, _scale))); + const __m128i _q2 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p2, _scale))); + const __m128i _q3 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p3, _scale))); + const __m128i _q01 = _mm_unpacklo_epi8(_q0, _q1); + const __m128i _q23 = _mm_unpacklo_epi8(_q2, _q3); +#if __AVX512VNNI__ || __AVXVNNI__ + const __m128i _q = _mm_unpacklo_epi16(_q01, _q23); + _mm_storeu_si128((__m128i*)pp, _q); +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + _w_shift = _mm_comp_dpbusd_epi32(_w_shift, _mm_set1_epi8(127), _q); +#endif +#else + _mm_storeu_si128((__m128i*)pp, _mm_unpacklo_epi64(_q01, _q23)); +#endif // __AVX512VNNI__ || __AVXVNNI__ + pp += 16; + } +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + if (max_kk >= 4) + { + _mm_storeu_si128((__m128i*)pp, _w_shift); + pp += 16; + } +#endif + for (; kk + 1 < max_kk; kk += 2) + { + __m128 _p0 = _mm_loadu_ps((const float*)A + (k0 + kk) * A_hstep + i + ii); + __m128 _p1 = _mm_loadu_ps((const float*)A + (k0 + kk + 1) * A_hstep + i + ii); + if (input_scale_ptr) + { + _p0 = _mm_mul_ps(_p0, _mm_set1_ps(input_scale_ptr[k0 + kk])); + _p1 = _mm_mul_ps(_p1, _mm_set1_ps(input_scale_ptr[k0 + kk + 1])); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+x"(_p0), "+x"(_p1)); +#else + volatile __m128 _p0_ordered = _p0; + volatile __m128 _p1_ordered = _p1; + _p0 = _p0_ordered; + _p1 = _p1_ordered; +#endif + } + const __m128i _q0 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p0, _scale))); + const __m128i _q1 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p1, _scale))); + _mm_storel_epi64((__m128i*)pp, _mm_unpacklo_epi8(_q0, _q1)); + pp += 8; + } + for (; kk < max_kk; kk++) + { + __m128 _p = _mm_loadu_ps((const float*)A + (k0 + kk) * A_hstep + i + ii); + if (input_scale_ptr) + { + _p = _mm_mul_ps(_p, _mm_set1_ps(input_scale_ptr[k0 + kk])); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+x"(_p)); +#else + volatile __m128 _p_ordered = _p; + _p = _p_ordered; +#endif + } + *(int*)pp = float2int8_sse(_mm_mul_ps(_p, _scale)); + pp += 4; + } + } + } +#endif // __SSE2__ + for (; ii + 1 < max_ii; ii += 2) + { +#if __SSE2__ + signed char* outptr0 = outptr + ii * out_hstep; + float* descale_ptr0 = descale_ptr + ii * block_count; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + + __m128 _absmax = _mm_setzero_ps(); + for (int kk = 0; kk < max_kk; kk++) + { + __m128 _p = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)((const float*)A + (k0 + kk) * A_hstep + i + ii)); + if (input_scale_ptr) + _p = _mm_mul_ps(_p, _mm_set1_ps(input_scale_ptr[k0 + kk])); + _absmax = _mm_max_ps(_absmax, abs_ps(_p)); + } + + const __m128 _descale = _mm_div_ps(_absmax, _mm_set1_ps(127.f)); + const __m128 _nonzero = _mm_cmpneq_ps(_absmax, _mm_setzero_ps()); + const __m128 _absmax_nonzero = _mm_or_ps(_mm_and_ps(_absmax, _nonzero), _mm_andnot_ps(_nonzero, _mm_set1_ps(1.f))); + const __m128 _scale = _mm_and_ps(_mm_cvtpd_ps(_mm_div_pd(_mm_set1_pd(127.0), _mm_cvtps_pd(_absmax_nonzero))), _nonzero); + _mm_storel_pi((__m64*)(descale_ptr0 + g * 2), _descale); + +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + __m128i _w_shift = _mm_setzero_si128(); + signed char* pp = outptr0 + (k0 + g * 4) * 2; +#else + signed char* pp = outptr0 + k0 * 2; +#endif + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __m128 _p0 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)((const float*)A + (k0 + kk) * A_hstep + i + ii)); + __m128 _p1 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)((const float*)A + (k0 + kk + 1) * A_hstep + i + ii)); + __m128 _p2 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)((const float*)A + (k0 + kk + 2) * A_hstep + i + ii)); + __m128 _p3 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)((const float*)A + (k0 + kk + 3) * A_hstep + i + ii)); + if (input_scale_ptr) + { + _p0 = _mm_mul_ps(_p0, _mm_set1_ps(input_scale_ptr[k0 + kk])); + _p1 = _mm_mul_ps(_p1, _mm_set1_ps(input_scale_ptr[k0 + kk + 1])); + _p2 = _mm_mul_ps(_p2, _mm_set1_ps(input_scale_ptr[k0 + kk + 2])); + _p3 = _mm_mul_ps(_p3, _mm_set1_ps(input_scale_ptr[k0 + kk + 3])); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+x"(_p0), "+x"(_p1), "+x"(_p2), "+x"(_p3)); +#else + volatile __m128 _p0_ordered = _p0; + volatile __m128 _p1_ordered = _p1; + volatile __m128 _p2_ordered = _p2; + volatile __m128 _p3_ordered = _p3; + _p0 = _p0_ordered; + _p1 = _p1_ordered; + _p2 = _p2_ordered; + _p3 = _p3_ordered; +#endif + } + const __m128i _q0 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p0, _scale))); + const __m128i _q1 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p1, _scale))); + const __m128i _q2 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p2, _scale))); + const __m128i _q3 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p3, _scale))); + const __m128i _q01 = _mm_unpacklo_epi8(_q0, _q1); + const __m128i _q23 = _mm_unpacklo_epi8(_q2, _q3); +#if __AVX512VNNI__ || __AVXVNNI__ + const __m128i _q = _mm_unpacklo_epi16(_q01, _q23); + _mm_storel_epi64((__m128i*)pp, _q); +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + _w_shift = _mm_comp_dpbusd_epi32(_w_shift, _mm_set1_epi8(127), _q); +#endif +#else + _mm_storel_epi64((__m128i*)pp, _mm_unpacklo_epi32(_q01, _q23)); +#endif // __AVX512VNNI__ || __AVXVNNI__ + pp += 8; + } +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + if (max_kk >= 4) + { + _mm_storel_epi64((__m128i*)pp, _w_shift); + pp += 8; + } +#endif + for (; kk + 1 < max_kk; kk += 2) + { + __m128 _p0 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)((const float*)A + (k0 + kk) * A_hstep + i + ii)); + __m128 _p1 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)((const float*)A + (k0 + kk + 1) * A_hstep + i + ii)); + if (input_scale_ptr) + { + _p0 = _mm_mul_ps(_p0, _mm_set1_ps(input_scale_ptr[k0 + kk])); + _p1 = _mm_mul_ps(_p1, _mm_set1_ps(input_scale_ptr[k0 + kk + 1])); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+x"(_p0), "+x"(_p1)); +#else + volatile __m128 _p0_ordered = _p0; + volatile __m128 _p1_ordered = _p1; + _p0 = _p0_ordered; + _p1 = _p1_ordered; +#endif + } + const __m128i _q0 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p0, _scale))); + const __m128i _q1 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p1, _scale))); + *(int*)pp = _mm_cvtsi128_si32(_mm_unpacklo_epi8(_q0, _q1)); + pp += 4; + } + for (; kk < max_kk; kk++) + { + __m128 _p = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)((const float*)A + (k0 + kk) * A_hstep + i + ii)); + if (input_scale_ptr) + { + _p = _mm_mul_ps(_p, _mm_set1_ps(input_scale_ptr[k0 + kk])); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+x"(_p)); +#else + volatile __m128 _p_ordered = _p; + _p = _p_ordered; +#endif + } + *(unsigned short*)pp = (unsigned short)float2int8_sse(_mm_mul_ps(_p, _scale)); + pp += 2; + } + } +#else + signed char* outptr0 = outptr + ii * out_hstep; + float* descale_ptr0 = descale_ptr + ii * block_count; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + float absmax0 = 0.f; + float absmax1 = 0.f; + for (int kk = 0; kk < max_kk; kk++) + { + const float* ptrA = (const float*)A + (k0 + kk) * A_hstep + i + ii; + float v0 = ptrA[0]; + float v1 = ptrA[1]; + if (input_scale_ptr) + { + const float s = input_scale_ptr[k0 + kk]; + v0 *= s; + v1 *= s; + } + absmax0 = std::max(absmax0, (float)fabsf(v0)); + absmax1 = std::max(absmax1, (float)fabsf(v1)); + } + + float scale0 = 0.f; + float scale1 = 0.f; + if (absmax0 != 0.f) + { + volatile double scale_fp64 = 127.0 / (double)absmax0; + scale0 = (float)scale_fp64; + } + if (absmax1 != 0.f) + { + volatile double scale_fp64 = 127.0 / (double)absmax1; + scale1 = (float)scale_fp64; + } + descale_ptr0[g * 2] = absmax0 / 127.f; + descale_ptr0[g * 2 + 1] = absmax1 / 127.f; + + signed char* pp = outptr0 + k0 * 2; + for (int kk = 0; kk < max_kk; kk++) + { + const float* ptrA = (const float*)A + (k0 + kk) * A_hstep + i + ii; + float v0 = ptrA[0]; + float v1 = ptrA[1]; + if (input_scale_ptr) + { + const float s = input_scale_ptr[k0 + kk]; + v0 *= s; + v1 *= s; + volatile float v0_ordered = v0; + volatile float v1_ordered = v1; + v0 = v0_ordered; + v1 = v1_ordered; + } + pp[0] = float2int8(v0 * scale0); + pp[1] = float2int8(v1 * scale1); + pp += 2; + } + } +#endif // __SSE2__ + } + +#if __SSE2__ +#if __AVX2__ +#if __AVX512F__ + __m512i _vindex512 = _mm512_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15); + _vindex512 = _mm512_mullo_epi32(_vindex512, _mm512_set1_epi32((int)A_hstep)); +#endif // __AVX512F__ + __m256i _vindex256 = _mm256_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7); + _vindex256 = _mm256_mullo_epi32(_vindex256, _mm256_set1_epi32((int)A_hstep)); +#endif // __AVX2__ +#endif // __SSE2__ + + for (; ii < max_ii; ii++) + { + signed char* outptr0 = outptr + ii * out_hstep; + float* descale_ptr0 = descale_ptr + ii * block_count; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + signed char* pp = outptr0 + k0 + g * 4; +#else + signed char* pp = outptr0 + k0; +#endif + float absmax = 0.f; + + int kk = 0; +#if __SSE2__ +#if __AVX2__ +#if __AVX512F__ + __m512 _absmax512 = _mm512_setzero_ps(); + for (; kk + 15 < max_kk; kk += 16) + { + const float* ptrA = (const float*)A + (k0 + kk) * A_hstep + i + ii; + __m512 _p = _mm512_i32gather_ps(_vindex512, ptrA, sizeof(float)); + if (input_scale_ptr) + _p = _mm512_mul_ps(_p, _mm512_loadu_ps(input_scale_ptr + k0 + kk)); + _absmax512 = _mm512_max_ps(_absmax512, abs512_ps(_p)); + } + absmax = std::max(absmax, _mm512_comp_reduce_max_ps(_absmax512)); +#endif // __AVX512F__ + __m256 _absmax256 = _mm256_setzero_ps(); + for (; kk + 7 < max_kk; kk += 8) + { + const float* ptrA = (const float*)A + (k0 + kk) * A_hstep + i + ii; + __m256 _p = _mm256_i32gather_ps(ptrA, _vindex256, sizeof(float)); + if (input_scale_ptr) + _p = _mm256_mul_ps(_p, _mm256_loadu_ps(input_scale_ptr + k0 + kk)); + _absmax256 = _mm256_max_ps(_absmax256, abs256_ps(_p)); + } + absmax = std::max(absmax, _mm256_reduce_max_ps(_absmax256)); +#endif // __AVX2__ + __m128 _absmax128 = _mm_setzero_ps(); + for (; kk + 3 < max_kk; kk += 4) + { + const float* ptrA = (const float*)A + (k0 + kk) * A_hstep + i + ii; + __m128 _p = _mm_setr_ps(ptrA[0], ptrA[A_hstep], ptrA[A_hstep * 2], ptrA[A_hstep * 3]); + if (input_scale_ptr) + _p = _mm_mul_ps(_p, _mm_loadu_ps(input_scale_ptr + k0 + kk)); + _absmax128 = _mm_max_ps(_absmax128, abs_ps(_p)); + } + absmax = std::max(absmax, _mm_reduce_max_ps(_absmax128)); +#endif // __SSE2__ + for (; kk < max_kk; kk++) + { + const int k = k0 + kk; + float v = ((const float*)A)[k * A_hstep + i + ii]; + if (input_scale_ptr) + v *= input_scale_ptr[k]; + absmax = std::max(absmax, (float)fabsf(v)); + } + + if (absmax == 0.f) + { + descale_ptr0[g] = 0.f; +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + memset(pp, 0, max_kk >= 4 ? max_kk + 4 : max_kk); +#else + memset(pp, 0, max_kk); +#endif + continue; + } + +#if __SSE2__ + const float scale = (float)(127.0 / (double)absmax); +#else + volatile double scale_fp64 = 127.0 / (double)absmax; + const float scale = (float)scale_fp64; +#endif + descale_ptr0[g] = absmax / 127.f; +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + int w_shift = 0; +#endif + kk = 0; +#if __SSE2__ +#if __AVX2__ +#if __AVX512F__ + const __m512 _scale512 = _mm512_set1_ps(scale); + for (; kk + 15 < max_kk; kk += 16) + { + const float* ptrA = (const float*)A + (k0 + kk) * A_hstep + i + ii; + __m512 _p = _mm512_i32gather_ps(_vindex512, ptrA, sizeof(float)); + if (input_scale_ptr) + { + _p = _mm512_mul_ps(_p, _mm512_loadu_ps(input_scale_ptr + k0 + kk)); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+x"(_p)); +#else + volatile __m512 _p_ordered = _p; + _p = _p_ordered; +#endif + } + const __m128i _q = float2int8_avx512(_mm512_mul_ps(_p, _scale512)); + _mm_storeu_si128((__m128i*)pp, _q); + pp += 16; +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + const __m256i _q16 = _mm256_cvtepi8_epi16(_q); + const __m256i _q32 = _mm256_madd_epi16(_q16, _mm256_set1_epi16(1)); + w_shift += _mm_reduce_add_epi32(_mm256_castsi256_si128(_q32)); + w_shift += _mm_reduce_add_epi32(_mm256_extracti128_si256(_q32, 1)); +#endif + } +#endif // __AVX512F__ + const __m256 _scale256 = _mm256_set1_ps(scale); + for (; kk + 7 < max_kk; kk += 8) + { + const float* ptrA = (const float*)A + (k0 + kk) * A_hstep + i + ii; + __m256 _p = _mm256_i32gather_ps(ptrA, _vindex256, sizeof(float)); + if (input_scale_ptr) + { + _p = _mm256_mul_ps(_p, _mm256_loadu_ps(input_scale_ptr + k0 + kk)); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+x"(_p)); +#else + volatile __m256 _p_ordered = _p; + _p = _p_ordered; +#endif + } + const int64_t q = float2int8_avx(_mm256_mul_ps(_p, _scale256)); + *(int64_t*)pp = q; + pp += 8; +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) +#if defined(__x86_64__) || defined(_M_X64) + const __m128i _q8 = _mm_cvtsi64_si128(q); +#else + const __m128i _q8 = _mm_loadl_epi64((const __m128i*)(pp - 8)); +#endif + const __m128i _q16 = _mm_unpacklo_epi8(_q8, _mm_cmpgt_epi8(_mm_setzero_si128(), _q8)); + w_shift += _mm_reduce_add_epi32(_mm_madd_epi16(_q16, _mm_set1_epi16(1))); +#endif + } +#endif // __AVX2__ + const __m128 _scale128 = _mm_set1_ps(scale); + for (; kk + 3 < max_kk; kk += 4) + { + const float* ptrA = (const float*)A + (k0 + kk) * A_hstep + i + ii; + __m128 _p = _mm_setr_ps(ptrA[0], ptrA[A_hstep], ptrA[A_hstep * 2], ptrA[A_hstep * 3]); + if (input_scale_ptr) + { + _p = _mm_mul_ps(_p, _mm_loadu_ps(input_scale_ptr + k0 + kk)); +#if NCNN_GNU_INLINE_ASM + asm volatile("" : "+x"(_p)); +#else + volatile __m128 _p_ordered = _p; + _p = _p_ordered; +#endif + } + const int32_t q = float2int8_sse(_mm_mul_ps(_p, _scale128)); + *(int32_t*)pp = q; + pp += 4; +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + const __m128i _q8 = _mm_cvtsi32_si128(q); + const __m128i _q16 = _mm_unpacklo_epi8(_q8, _mm_cmpgt_epi8(_mm_setzero_si128(), _q8)); + w_shift += _mm_reduce_add_epi32(_mm_madd_epi16(_q16, _mm_set1_epi16(1))); +#endif + } +#endif // __SSE2__ +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + if (max_kk >= 4) + { + ((int*)pp)[0] = w_shift * 127; + pp += 4; + } +#endif + for (; kk < max_kk; kk++) + { + const int k = k0 + kk; + float v = ((const float*)A)[k * A_hstep + i + ii]; + if (input_scale_ptr) + { + v *= input_scale_ptr[k]; +#if NCNN_GNU_INLINE_ASM && __SSE2__ + asm volatile("" : "+x"(v)); +#else + volatile float v_ordered = v; + v = v_ordered; +#endif + } + *pp++ = float2int8(v * scale); + } + } + } +} + +static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int block_size) +{ +#if NCNN_RUNTIME_CPU && NCNN_AVX512VNNI && __AVX512F__ && !__AVX512VNNI__ + if (ncnn::cpu_support_x86_avx512_vnni()) + { + gemm_transB_packed_tile_wq_int8_avx512vnni(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); + return; + } +#endif +#if NCNN_RUNTIME_CPU && NCNN_AVXVNNIINT8 && __AVX__ && !__AVXVNNIINT8__ && !__AVX512VNNI__ + if (ncnn::cpu_support_x86_avx_vnni_int8()) + { + gemm_transB_packed_tile_wq_int8_avxvnniint8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); + return; + } +#endif +#if NCNN_RUNTIME_CPU && NCNN_AVXVNNI && __AVX__ && !__AVXVNNI__ && !__AVXVNNIINT8__ && !__AVX512VNNI__ + if (ncnn::cpu_support_x86_avx_vnni()) + { + gemm_transB_packed_tile_wq_int8_avxvnni(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); + return; + } +#endif +#if NCNN_RUNTIME_CPU && NCNN_AVX2 && __AVX__ && !__AVX2__ && !__AVXVNNI__ && !__AVXVNNIINT8__ && !__AVX512VNNI__ + if (ncnn::cpu_support_x86_avx2()) + { + gemm_transB_packed_tile_wq_int8_avx2(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); + return; + } +#endif +#if NCNN_RUNTIME_CPU && NCNN_XOP && __SSE2__ && !__XOP__ && !__AVX2__ && !__AVXVNNI__ && !__AVXVNNIINT8__ && !__AVX512VNNI__ + if (ncnn::cpu_support_x86_xop()) + { + gemm_transB_packed_tile_wq_int8_xop(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); + return; + } +#endif + + const signed char* pAT = AT_tile; + const int A_hstep = AT_tile.w; + const float* pAT_descales = AT_descales_tile; + const int A_descales_hstep = AT_descales_tile.w; + const signed char* pBT = BT_tile; + const float* pBT_descales = BT_descales_tile; + float* outptr = topT_tile; + + int ii = 0; +#if __SSE2__ +#if __AVX2__ +#if __AVX512F__ + for (; ii + 15 < max_ii; ii += 16) + { + const signed char* pB = pBT; + const float* pB_descales = pBT_descales; + + int jj = 0; +#if defined(__x86_64__) || defined(_M_X64) + for (; jj + 7 < max_jj; jj += 8) + { + __m512 _fsum0 = _mm512_setzero_ps(); + __m512 _fsum1 = _mm512_setzero_ps(); + __m512 _fsum2 = _mm512_setzero_ps(); + __m512 _fsum3 = _mm512_setzero_ps(); + __m512 _fsum4 = _mm512_setzero_ps(); + __m512 _fsum5 = _mm512_setzero_ps(); + __m512 _fsum6 = _mm512_setzero_ps(); + __m512 _fsum7 = _mm512_setzero_ps(); + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + __m512i _sum0 = _mm512_setzero_si512(); + __m512i _sum1 = _mm512_setzero_si512(); + __m512i _sum2 = _mm512_setzero_si512(); + __m512i _sum3 = _mm512_setzero_si512(); + __m512i _sum4 = _mm512_setzero_si512(); + __m512i _sum5 = _mm512_setzero_si512(); + __m512i _sum6 = _mm512_setzero_si512(); + __m512i _sum7 = _mm512_setzero_si512(); + const int max_kk = std::min(K - k, block_size); + int kk = 0; +#if __AVX512VNNI__ + for (; kk + 3 < max_kk; kk += 4) + { + const __m512i _pA0 = _mm512_loadu_si512((const __m512i*)pA); + const __m256i _pB = _mm256_loadu_si256((const __m256i*)pB); + const __m512i _pB0 = combine8x2_epi32(_pB, _pB); + const __m512i _pA1 = _mm512_alignr_epi8(_pA0, _pA0, 8); + const __m512i _pB1 = _mm512_alignr_epi8(_pB0, _pB0, 4); + const __m512i _pB2 = _mm512_permutex_epi64(_pB0, _MM_SHUFFLE(1, 0, 3, 2)); + const __m512i _pB3 = _mm512_alignr_epi8(_pB2, _pB2, 4); + _sum0 = _mm512_dpbusd_epi32(_sum0, _pB0, _pA0); + _sum1 = _mm512_dpbusd_epi32(_sum1, _pB1, _pA0); + _sum2 = _mm512_dpbusd_epi32(_sum2, _pB0, _pA1); + _sum3 = _mm512_dpbusd_epi32(_sum3, _pB1, _pA1); + _sum4 = _mm512_dpbusd_epi32(_sum4, _pB2, _pA0); + _sum5 = _mm512_dpbusd_epi32(_sum5, _pB3, _pA0); + _sum6 = _mm512_dpbusd_epi32(_sum6, _pB2, _pA1); + _sum7 = _mm512_dpbusd_epi32(_sum7, _pB3, _pA1); + pB += 32; + pA += 64; + } + if (max_kk >= 4) + { + const __m512i _shift0 = _mm512_loadu_si512((const __m512i*)pA); + const __m512i _shift1 = _mm512_alignr_epi8(_shift0, _shift0, 8); + _sum0 = _mm512_sub_epi32(_sum0, _shift0); + _sum1 = _mm512_sub_epi32(_sum1, _shift0); + _sum2 = _mm512_sub_epi32(_sum2, _shift1); + _sum3 = _mm512_sub_epi32(_sum3, _shift1); + _sum4 = _mm512_sub_epi32(_sum4, _shift0); + _sum5 = _mm512_sub_epi32(_sum5, _shift0); + _sum6 = _mm512_sub_epi32(_sum6, _shift1); + _sum7 = _mm512_sub_epi32(_sum7, _shift1); + pA += 64; + } +#endif // __AVX512VNNI__ + for (; kk + 1 < max_kk; kk += 2) + { + const __m256i _pA = _mm256_loadu_si256((const __m256i*)pA); + const __m512i _pA0 = _mm512_cvtepi8_epi16(_pA); + const __m512i _pA1 = _mm512_alignr_epi8(_pA0, _pA0, 8); + const __m128i _pB = _mm_loadu_si128((const __m128i*)pB); + const __m256i _pBB = _mm256_cvtepi8_epi16(_pB); + const __m512i _pB0 = combine8x2_epi32(_pBB, _pBB); + const __m512i _pB1 = _mm512_alignr_epi8(_pB0, _pB0, 4); + const __m512i _pB2 = _mm512_permutex_epi64(_pB0, _MM_SHUFFLE(1, 0, 3, 2)); + const __m512i _pB3 = _mm512_alignr_epi8(_pB2, _pB2, 4); + _sum0 = _mm512_comp_dpwssd_epi32(_sum0, _pA0, _pB0); + _sum1 = _mm512_comp_dpwssd_epi32(_sum1, _pA0, _pB1); + _sum2 = _mm512_comp_dpwssd_epi32(_sum2, _pA1, _pB0); + _sum3 = _mm512_comp_dpwssd_epi32(_sum3, _pA1, _pB1); + _sum4 = _mm512_comp_dpwssd_epi32(_sum4, _pA0, _pB2); + _sum5 = _mm512_comp_dpwssd_epi32(_sum5, _pA0, _pB3); + _sum6 = _mm512_comp_dpwssd_epi32(_sum6, _pA1, _pB2); + _sum7 = _mm512_comp_dpwssd_epi32(_sum7, _pA1, _pB3); + pB += 16; + pA += 32; + } + for (; kk < max_kk; kk++) + { + const __m128i _pA = _mm_loadu_si128((const __m128i*)pA); + const __m256i _pA0 = _mm256_cvtepi8_epi16(_pA); + const __m256i _pA1 = _mm256_shuffle_epi32(_pA0, _MM_SHUFFLE(2, 3, 0, 1)); + __m128i _pB = _mm_loadl_epi64((const __m128i*)pB); + _pB = _mm_cvtepi8_epi16(_pB); + const __m256i _pB0 = combine4x2_epi32(_pB, _pB); + const __m256i _pB1 = _mm256_shufflehi_epi16(_mm256_shufflelo_epi16(_pB0, _MM_SHUFFLE(0, 3, 2, 1)), _MM_SHUFFLE(0, 3, 2, 1)); + const __m256i _pB2 = _mm256_alignr_epi8(_pB0, _pB0, 8); + const __m256i _pB3 = _mm256_shufflehi_epi16(_mm256_shufflelo_epi16(_pB2, _MM_SHUFFLE(0, 3, 2, 1)), _MM_SHUFFLE(0, 3, 2, 1)); + _sum0 = _mm512_add_epi32(_sum0, _mm512_cvtepi16_epi32(_mm256_mullo_epi16(_pA0, _pB0))); + _sum1 = _mm512_add_epi32(_sum1, _mm512_cvtepi16_epi32(_mm256_mullo_epi16(_pA0, _pB1))); + _sum2 = _mm512_add_epi32(_sum2, _mm512_cvtepi16_epi32(_mm256_mullo_epi16(_pA1, _pB0))); + _sum3 = _mm512_add_epi32(_sum3, _mm512_cvtepi16_epi32(_mm256_mullo_epi16(_pA1, _pB1))); + _sum4 = _mm512_add_epi32(_sum4, _mm512_cvtepi16_epi32(_mm256_mullo_epi16(_pA0, _pB2))); + _sum5 = _mm512_add_epi32(_sum5, _mm512_cvtepi16_epi32(_mm256_mullo_epi16(_pA0, _pB3))); + _sum6 = _mm512_add_epi32(_sum6, _mm512_cvtepi16_epi32(_mm256_mullo_epi16(_pA1, _pB2))); + _sum7 = _mm512_add_epi32(_sum7, _mm512_cvtepi16_epi32(_mm256_mullo_epi16(_pA1, _pB3))); + pB += 8; + pA += 16; + } + + const __m512 _A0 = _mm512_loadu_ps(pA_descales); + const __m512 _A1 = _mm512_castsi512_ps(_mm512_alignr_epi8(_mm512_castps_si512(_A0), _mm512_castps_si512(_A0), 8)); + const __m256 _b = _mm256_loadu_ps(pB_descales); + const __m512 _B0 = combine8x2_ps(_b, _b); + const __m512 _B1 = _mm512_castsi512_ps(_mm512_alignr_epi8(_mm512_castps_si512(_B0), _mm512_castps_si512(_B0), 4)); + const __m512 _B2 = _mm512_castsi512_ps(_mm512_permutex_epi64(_mm512_castps_si512(_B0), _MM_SHUFFLE(1, 0, 3, 2))); + const __m512 _B3 = _mm512_castsi512_ps(_mm512_alignr_epi8(_mm512_castps_si512(_B2), _mm512_castps_si512(_B2), 4)); + _fsum0 = _mm512_add_ps(_fsum0, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum0), _mm512_mul_ps(_A0, _B0))); + _fsum1 = _mm512_add_ps(_fsum1, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum1), _mm512_mul_ps(_A0, _B1))); + _fsum2 = _mm512_add_ps(_fsum2, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum2), _mm512_mul_ps(_A1, _B0))); + _fsum3 = _mm512_add_ps(_fsum3, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum3), _mm512_mul_ps(_A1, _B1))); + _fsum4 = _mm512_add_ps(_fsum4, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum4), _mm512_mul_ps(_A0, _B2))); + _fsum5 = _mm512_add_ps(_fsum5, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum5), _mm512_mul_ps(_A0, _B3))); + _fsum6 = _mm512_add_ps(_fsum6, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum6), _mm512_mul_ps(_A1, _B2))); + _fsum7 = _mm512_add_ps(_fsum7, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum7), _mm512_mul_ps(_A1, _B3))); + pA_descales += 16; + pB_descales += 8; + } + + _mm512_storeu_ps(outptr + 0, _fsum0); + _mm512_storeu_ps(outptr + 16, _fsum1); + _mm512_storeu_ps(outptr + 32, _fsum2); + _mm512_storeu_ps(outptr + 48, _fsum3); + _mm512_storeu_ps(outptr + 64, _fsum4); + _mm512_storeu_ps(outptr + 80, _fsum5); + _mm512_storeu_ps(outptr + 96, _fsum6); + _mm512_storeu_ps(outptr + 112, _fsum7); + outptr += 128; + } + for (; jj + 3 < max_jj; jj += 4) + { + __m512 _fsum0 = _mm512_setzero_ps(); + __m512 _fsum1 = _mm512_setzero_ps(); + __m512 _fsum2 = _mm512_setzero_ps(); + __m512 _fsum3 = _mm512_setzero_ps(); + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + __m512i _sum0 = _mm512_setzero_si512(); + __m512i _sum1 = _mm512_setzero_si512(); + __m512i _sum2 = _mm512_setzero_si512(); + __m512i _sum3 = _mm512_setzero_si512(); + const int max_kk = std::min(K - k, block_size); + int kk = 0; +#if __AVX512VNNI__ + for (; kk + 3 < max_kk; kk += 4) + { + const __m512i _pA0 = _mm512_loadu_si512((const __m512i*)pA); + const __m512i _pB0 = _mm512_broadcast_i32x4(_mm_loadu_si128((const __m128i*)pB)); + const __m512i _pA1 = _mm512_alignr_epi8(_pA0, _pA0, 8); + const __m512i _pB1 = _mm512_alignr_epi8(_pB0, _pB0, 4); + _sum0 = _mm512_dpbusd_epi32(_sum0, _pB0, _pA0); + _sum1 = _mm512_dpbusd_epi32(_sum1, _pB1, _pA0); + _sum2 = _mm512_dpbusd_epi32(_sum2, _pB0, _pA1); + _sum3 = _mm512_dpbusd_epi32(_sum3, _pB1, _pA1); + pB += 16; + pA += 64; + } + if (max_kk >= 4) + { + const __m512i _shift0 = _mm512_loadu_si512((const __m512i*)pA); + const __m512i _shift1 = _mm512_alignr_epi8(_shift0, _shift0, 8); + _sum0 = _mm512_sub_epi32(_sum0, _shift0); + _sum1 = _mm512_sub_epi32(_sum1, _shift0); + _sum2 = _mm512_sub_epi32(_sum2, _shift1); + _sum3 = _mm512_sub_epi32(_sum3, _shift1); + pA += 64; + } +#endif // __AVX512VNNI__ + for (; kk + 1 < max_kk; kk += 2) + { + const __m256i _pA = _mm256_loadu_si256((const __m256i*)pA); + const __m512i _pA0 = _mm512_cvtepi8_epi16(_pA); + const __m512i _pA1 = _mm512_alignr_epi8(_pA0, _pA0, 8); + const __m256i _pB = _mm256_castpd_si256(_mm256_broadcast_sd((const double*)pB)); + const __m512i _pB0 = _mm512_cvtepi8_epi16(_pB); + const __m512i _pB1 = _mm512_alignr_epi8(_pB0, _pB0, 4); + _sum0 = _mm512_comp_dpwssd_epi32(_sum0, _pA0, _pB0); + _sum1 = _mm512_comp_dpwssd_epi32(_sum1, _pA0, _pB1); + _sum2 = _mm512_comp_dpwssd_epi32(_sum2, _pA1, _pB0); + _sum3 = _mm512_comp_dpwssd_epi32(_sum3, _pA1, _pB1); + pB += 8; + pA += 32; + } + for (; kk < max_kk; kk++) + { + const __m128i _pA = _mm_loadu_si128((const __m128i*)pA); + const __m256i _pA0 = _mm256_cvtepi8_epi16(_pA); + const __m256i _pA1 = _mm256_shuffle_epi32(_pA0, _MM_SHUFFLE(2, 3, 0, 1)); + const __m256i _pB0 = _mm256_cvtepi8_epi16(_mm_castps_si128(_mm_load1_ps((const float*)pB))); + const __m256i _pB1 = _mm256_shufflehi_epi16(_mm256_shufflelo_epi16(_pB0, _MM_SHUFFLE(0, 3, 2, 1)), _MM_SHUFFLE(0, 3, 2, 1)); + _sum0 = _mm512_add_epi32(_sum0, _mm512_cvtepi16_epi32(_mm256_mullo_epi16(_pA0, _pB0))); + _sum1 = _mm512_add_epi32(_sum1, _mm512_cvtepi16_epi32(_mm256_mullo_epi16(_pA0, _pB1))); + _sum2 = _mm512_add_epi32(_sum2, _mm512_cvtepi16_epi32(_mm256_mullo_epi16(_pA1, _pB0))); + _sum3 = _mm512_add_epi32(_sum3, _mm512_cvtepi16_epi32(_mm256_mullo_epi16(_pA1, _pB1))); + pB += 4; + pA += 16; + } + + const __m512 _A0 = _mm512_loadu_ps(pA_descales); + const __m512 _A1 = _mm512_castsi512_ps(_mm512_alignr_epi8(_mm512_castps_si512(_A0), _mm512_castps_si512(_A0), 8)); + const __m512 _B0 = _mm512_broadcast_f32x4(_mm_loadu_ps(pB_descales)); + const __m512 _B1 = _mm512_castsi512_ps(_mm512_alignr_epi8(_mm512_castps_si512(_B0), _mm512_castps_si512(_B0), 4)); + _fsum0 = _mm512_add_ps(_fsum0, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum0), _mm512_mul_ps(_A0, _B0))); + _fsum1 = _mm512_add_ps(_fsum1, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum1), _mm512_mul_ps(_A0, _B1))); + _fsum2 = _mm512_add_ps(_fsum2, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum2), _mm512_mul_ps(_A1, _B0))); + _fsum3 = _mm512_add_ps(_fsum3, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum3), _mm512_mul_ps(_A1, _B1))); + pA_descales += 16; + pB_descales += 4; + } + + _mm512_storeu_ps(outptr + 0, _fsum0); + _mm512_storeu_ps(outptr + 16, _fsum1); + _mm512_storeu_ps(outptr + 32, _fsum2); + _mm512_storeu_ps(outptr + 48, _fsum3); + outptr += 64; + } +#endif // defined(__x86_64__) || defined(_M_X64) + for (; jj + 1 < max_jj; jj += 2) + { + __m512 _fsum0 = _mm512_setzero_ps(); + __m512 _fsum1 = _mm512_setzero_ps(); + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + __m512i _sum0 = _mm512_setzero_si512(); + __m512i _sum1 = _mm512_setzero_si512(); + const int max_kk = std::min(K - k, block_size); + int kk = 0; +#if __AVX512VNNI__ + for (; kk + 3 < max_kk; kk += 4) + { + const __m512i _pA0 = _mm512_loadu_si512((const __m512i*)pA); + const __m512i _pB0 = _mm512_castpd_si512(_mm512_set1_pd(*(const double*)pB)); + const __m512i _pB1 = _mm512_alignr_epi8(_pB0, _pB0, 4); + _sum0 = _mm512_dpbusd_epi32(_sum0, _pB0, _pA0); + _sum1 = _mm512_dpbusd_epi32(_sum1, _pB1, _pA0); + pB += 8; + pA += 64; + } + if (max_kk >= 4) + { + const __m512i _shift0 = _mm512_loadu_si512((const __m512i*)pA); + _sum0 = _mm512_sub_epi32(_sum0, _shift0); + _sum1 = _mm512_sub_epi32(_sum1, _shift0); + pA += 64; + } +#endif // __AVX512VNNI__ + for (; kk + 1 < max_kk; kk += 2) + { + const __m256i _pA = _mm256_loadu_si256((const __m256i*)pA); + const __m512i _pA0 = _mm512_cvtepi8_epi16(_pA); + const __m256i _pB = _mm256_castps_si256(_mm256_broadcast_ss((const float*)pB)); + const __m512i _pB0 = _mm512_cvtepi8_epi16(_pB); + const __m512i _pB1 = _mm512_alignr_epi8(_pB0, _pB0, 4); + _sum0 = _mm512_comp_dpwssd_epi32(_sum0, _pA0, _pB0); + _sum1 = _mm512_comp_dpwssd_epi32(_sum1, _pA0, _pB1); + pB += 4; + pA += 32; + } + for (; kk < max_kk; kk++) + { + const __m128i _pA = _mm_loadu_si128((const __m128i*)pA); + const __m256i _pA0 = _mm256_cvtepi8_epi16(_pA); + const __m128i _pB = _mm_set1_epi16(*(const short*)pB); + const __m256i _pB0 = _mm256_cvtepi8_epi16(_pB); + const __m256i _pB1 = _mm256_shufflehi_epi16(_mm256_shufflelo_epi16(_pB0, _MM_SHUFFLE(0, 1, 0, 1)), _MM_SHUFFLE(0, 1, 0, 1)); + _sum0 = _mm512_add_epi32(_sum0, _mm512_cvtepi16_epi32(_mm256_mullo_epi16(_pA0, _pB0))); + _sum1 = _mm512_add_epi32(_sum1, _mm512_cvtepi16_epi32(_mm256_mullo_epi16(_pA0, _pB1))); + pB += 2; + pA += 16; + } + + const __m512 _A0 = _mm512_loadu_ps(pA_descales); + const __m512 _B0 = _mm512_castpd_ps(_mm512_set1_pd(*(const double*)pB_descales)); + const __m512 _B1 = _mm512_castsi512_ps(_mm512_alignr_epi8(_mm512_castps_si512(_B0), _mm512_castps_si512(_B0), 4)); + _fsum0 = _mm512_add_ps(_fsum0, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum0), _mm512_mul_ps(_A0, _B0))); + _fsum1 = _mm512_add_ps(_fsum1, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum1), _mm512_mul_ps(_A0, _B1))); + pA_descales += 16; + pB_descales += 2; + } + + _mm512_storeu_ps(outptr + 0, _fsum0); + _mm512_storeu_ps(outptr + 16, _fsum1); + outptr += 32; + } + for (; jj < max_jj; jj++) + { + __m512 _fsum0 = _mm512_setzero_ps(); + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + __m512i _sum0 = _mm512_setzero_si512(); + const int max_kk = std::min(K - k, block_size); + int kk = 0; +#if __AVX512VNNI__ + for (; kk + 3 < max_kk; kk += 4) + { + const __m512i _pA0 = _mm512_loadu_si512((const __m512i*)pA); + const __m512i _pB0 = _mm512_set1_epi32(*(const int*)pB); + _sum0 = _mm512_dpbusd_epi32(_sum0, _pB0, _pA0); + pB += 4; + pA += 64; + } + if (max_kk >= 4) + { + const __m512i _shift0 = _mm512_loadu_si512((const __m512i*)pA); + _sum0 = _mm512_sub_epi32(_sum0, _shift0); + pA += 64; + } +#endif // __AVX512VNNI__ + for (; kk + 1 < max_kk; kk += 2) + { + const __m256i _pA = _mm256_loadu_si256((const __m256i*)pA); + const __m512i _pA0 = _mm512_cvtepi8_epi16(_pA); + const __m512i _pB0 = _mm512_cvtepi8_epi16(_mm256_set1_epi16(*(const short*)pB)); + _sum0 = _mm512_comp_dpwssd_epi32(_sum0, _pA0, _pB0); + pB += 2; + pA += 32; + } + for (; kk < max_kk; kk++) + { + const __m512i _pA0 = _mm512_cvtepi8_epi32(_mm_loadu_si128((const __m128i*)pA)); + _sum0 = _mm512_add_epi32(_sum0, _mm512_mullo_epi32(_pA0, _mm512_set1_epi32((signed char)pB[0]))); + pB += 1; + pA += 16; + } + + const __m512 _A0 = _mm512_loadu_ps(pA_descales); + const __m512 _B0 = _mm512_set1_ps(pB_descales[0]); + _fsum0 = _mm512_add_ps(_fsum0, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum0), _mm512_mul_ps(_A0, _B0))); + pA_descales += 16; + pB_descales += 1; + } + + _mm512_storeu_ps(outptr + 0, _fsum0); + outptr += 16; + } + + pAT += A_hstep * 16; + pAT_descales += A_descales_hstep * 16; + } +#endif // __AVX512F__ + for (; ii + 7 < max_ii; ii += 8) + { + const signed char* pB = pBT; + const float* pB_descales = pBT_descales; + + int jj = 0; +#if defined(__x86_64__) || defined(_M_X64) +#if __AVX512F__ + for (; jj + 7 < max_jj; jj += 8) + { + __m256 _fsum0 = _mm256_setzero_ps(); + __m256 _fsum1 = _mm256_setzero_ps(); + __m256 _fsum2 = _mm256_setzero_ps(); + __m256 _fsum3 = _mm256_setzero_ps(); + __m256 _fsum4 = _mm256_setzero_ps(); + __m256 _fsum5 = _mm256_setzero_ps(); + __m256 _fsum6 = _mm256_setzero_ps(); + __m256 _fsum7 = _mm256_setzero_ps(); + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + __m256i _sum0 = _mm256_setzero_si256(); + __m256i _sum1 = _mm256_setzero_si256(); + __m256i _sum2 = _mm256_setzero_si256(); + __m256i _sum3 = _mm256_setzero_si256(); + __m256i _sum4 = _mm256_setzero_si256(); + __m256i _sum5 = _mm256_setzero_si256(); + __m256i _sum6 = _mm256_setzero_si256(); + __m256i _sum7 = _mm256_setzero_si256(); + const int max_kk = std::min(K - k, block_size); + int kk = 0; +#if __AVX512VNNI__ + for (; kk + 3 < max_kk; kk += 4) + { + const __m256i _pA0 = _mm256_loadu_si256((const __m256i*)pA); + const __m256i _pA1 = _mm256_alignr_epi8(_pA0, _pA0, 8); + const __m256i _pB0 = _mm256_loadu_si256((const __m256i*)pB); + const __m256i _pB1 = _mm256_alignr_epi8(_pB0, _pB0, 4); + const __m256i _pB2 = _mm256_permute4x64_epi64(_pB0, _MM_SHUFFLE(1, 0, 3, 2)); + const __m256i _pB3 = _mm256_alignr_epi8(_pB2, _pB2, 4); + _sum0 = _mm256_comp_dpbusd_epi32(_sum0, _pB0, _pA0); + _sum1 = _mm256_comp_dpbusd_epi32(_sum1, _pB1, _pA0); + _sum2 = _mm256_comp_dpbusd_epi32(_sum2, _pB0, _pA1); + _sum3 = _mm256_comp_dpbusd_epi32(_sum3, _pB1, _pA1); + _sum4 = _mm256_comp_dpbusd_epi32(_sum4, _pB2, _pA0); + _sum5 = _mm256_comp_dpbusd_epi32(_sum5, _pB3, _pA0); + _sum6 = _mm256_comp_dpbusd_epi32(_sum6, _pB2, _pA1); + _sum7 = _mm256_comp_dpbusd_epi32(_sum7, _pB3, _pA1); + pB += 32; + pA += 32; + } + if (max_kk >= 4) + { + const __m256i _shift0 = _mm256_loadu_si256((const __m256i*)pA); + const __m256i _shift1 = _mm256_alignr_epi8(_shift0, _shift0, 8); + _sum0 = _mm256_sub_epi32(_sum0, _shift0); + _sum1 = _mm256_sub_epi32(_sum1, _shift0); + _sum2 = _mm256_sub_epi32(_sum2, _shift1); + _sum3 = _mm256_sub_epi32(_sum3, _shift1); + _sum4 = _mm256_sub_epi32(_sum4, _shift0); + _sum5 = _mm256_sub_epi32(_sum5, _shift0); + _sum6 = _mm256_sub_epi32(_sum6, _shift1); + _sum7 = _mm256_sub_epi32(_sum7, _shift1); + pA += 32; + } +#endif // __AVX512VNNI__ + for (; kk + 1 < max_kk; kk += 2) + { + const __m128i _pA8 = _mm_loadu_si128((const __m128i*)pA); + const __m128i _pB8 = _mm_loadu_si128((const __m128i*)pB); + const __m256i _pA0 = _mm256_cvtepi8_epi16(_pA8); + const __m256i _pB0 = _mm256_cvtepi8_epi16(_pB8); + const __m256i _pA1 = _mm256_alignr_epi8(_pA0, _pA0, 8); + const __m256i _pB1 = _mm256_alignr_epi8(_pB0, _pB0, 4); + const __m256i _pB2 = _mm256_permute4x64_epi64(_pB0, _MM_SHUFFLE(1, 0, 3, 2)); + const __m256i _pB3 = _mm256_alignr_epi8(_pB2, _pB2, 4); + _sum0 = _mm256_comp_dpwssd_epi32(_sum0, _pA0, _pB0); + _sum1 = _mm256_comp_dpwssd_epi32(_sum1, _pA0, _pB1); + _sum2 = _mm256_comp_dpwssd_epi32(_sum2, _pA1, _pB0); + _sum3 = _mm256_comp_dpwssd_epi32(_sum3, _pA1, _pB1); + _sum4 = _mm256_comp_dpwssd_epi32(_sum4, _pA0, _pB2); + _sum5 = _mm256_comp_dpwssd_epi32(_sum5, _pA0, _pB3); + _sum6 = _mm256_comp_dpwssd_epi32(_sum6, _pA1, _pB2); + _sum7 = _mm256_comp_dpwssd_epi32(_sum7, _pA1, _pB3); + pB += 16; + pA += 16; + } + if (kk < max_kk) + { + __m128i _pA0 = _mm_loadl_epi64((const __m128i*)pA); + __m128i _pB0 = _mm_loadl_epi64((const __m128i*)pB); + _pA0 = _mm_cvtepi8_epi16(_pA0); + _pB0 = _mm_cvtepi8_epi16(_pB0); + const __m128i _pA1 = _mm_shufflehi_epi16(_mm_shufflelo_epi16(_pA0, _MM_SHUFFLE(1, 0, 3, 2)), _MM_SHUFFLE(1, 0, 3, 2)); + const __m128i _pB1 = _mm_shufflehi_epi16(_mm_shufflelo_epi16(_pB0, _MM_SHUFFLE(0, 3, 2, 1)), _MM_SHUFFLE(0, 3, 2, 1)); + const __m128i _pB2 = _mm_alignr_epi8(_pB0, _pB0, 8); + const __m128i _pB3 = _mm_shufflehi_epi16(_mm_shufflelo_epi16(_pB2, _MM_SHUFFLE(0, 3, 2, 1)), _MM_SHUFFLE(0, 3, 2, 1)); + _sum0 = _mm256_add_epi32(_sum0, _mm256_cvtepi16_epi32(_mm_mullo_epi16(_pA0, _pB0))); + _sum1 = _mm256_add_epi32(_sum1, _mm256_cvtepi16_epi32(_mm_mullo_epi16(_pA0, _pB1))); + _sum2 = _mm256_add_epi32(_sum2, _mm256_cvtepi16_epi32(_mm_mullo_epi16(_pA1, _pB0))); + _sum3 = _mm256_add_epi32(_sum3, _mm256_cvtepi16_epi32(_mm_mullo_epi16(_pA1, _pB1))); + _sum4 = _mm256_add_epi32(_sum4, _mm256_cvtepi16_epi32(_mm_mullo_epi16(_pA0, _pB2))); + _sum5 = _mm256_add_epi32(_sum5, _mm256_cvtepi16_epi32(_mm_mullo_epi16(_pA0, _pB3))); + _sum6 = _mm256_add_epi32(_sum6, _mm256_cvtepi16_epi32(_mm_mullo_epi16(_pA1, _pB2))); + _sum7 = _mm256_add_epi32(_sum7, _mm256_cvtepi16_epi32(_mm_mullo_epi16(_pA1, _pB3))); + pB += 8; + } + + const __m256 _A0 = _mm256_loadu_ps(pA_descales); + const __m256 _A1 = _mm256_castsi256_ps(_mm256_alignr_epi8(_mm256_castps_si256(_A0), _mm256_castps_si256(_A0), 8)); + const __m256 _B0 = _mm256_loadu_ps(pB_descales); + const __m256 _B1 = _mm256_castsi256_ps(_mm256_alignr_epi8(_mm256_castps_si256(_B0), _mm256_castps_si256(_B0), 4)); + const __m256 _B2 = _mm256_castsi256_ps(_mm256_permute4x64_epi64(_mm256_castps_si256(_B0), _MM_SHUFFLE(1, 0, 3, 2))); + const __m256 _B3 = _mm256_castsi256_ps(_mm256_alignr_epi8(_mm256_castps_si256(_B2), _mm256_castps_si256(_B2), 4)); + _fsum0 = _mm256_add_ps(_fsum0, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum0), _mm256_mul_ps(_A0, _B0))); + _fsum1 = _mm256_add_ps(_fsum1, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum1), _mm256_mul_ps(_A0, _B1))); + _fsum2 = _mm256_add_ps(_fsum2, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum2), _mm256_mul_ps(_A1, _B0))); + _fsum3 = _mm256_add_ps(_fsum3, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum3), _mm256_mul_ps(_A1, _B1))); + _fsum4 = _mm256_add_ps(_fsum4, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum4), _mm256_mul_ps(_A0, _B2))); + _fsum5 = _mm256_add_ps(_fsum5, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum5), _mm256_mul_ps(_A0, _B3))); + _fsum6 = _mm256_add_ps(_fsum6, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum6), _mm256_mul_ps(_A1, _B2))); + _fsum7 = _mm256_add_ps(_fsum7, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum7), _mm256_mul_ps(_A1, _B3))); + pA_descales += 8; + pB_descales += 8; + } + + _mm256_storeu_ps(outptr + 0, _fsum0); + _mm256_storeu_ps(outptr + 8, _fsum1); + _mm256_storeu_ps(outptr + 16, _fsum2); + _mm256_storeu_ps(outptr + 24, _fsum3); + _mm256_storeu_ps(outptr + 32, _fsum4); + _mm256_storeu_ps(outptr + 40, _fsum5); + _mm256_storeu_ps(outptr + 48, _fsum6); + _mm256_storeu_ps(outptr + 56, _fsum7); + outptr += 64; + } +#endif // __AVX512F__ + for (; jj + 3 < max_jj; jj += 4) + { + __m256 _fsum0 = _mm256_setzero_ps(); + __m256 _fsum1 = _mm256_setzero_ps(); + __m256 _fsum2 = _mm256_setzero_ps(); + __m256 _fsum3 = _mm256_setzero_ps(); + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + __m256i _sum0 = _mm256_setzero_si256(); + __m256i _sum1 = _mm256_setzero_si256(); + __m256i _sum2 = _mm256_setzero_si256(); + __m256i _sum3 = _mm256_setzero_si256(); + const int max_kk = std::min(K - k, block_size); + int kk = 0; +#if __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 3 < max_kk; kk += 4) + { + const __m256i _pA0 = _mm256_loadu_si256((const __m256i*)pA); + const __m128i _pB = _mm_loadu_si128((const __m128i*)pB); + const __m256i _pB0 = combine4x2_epi32(_pB, _pB); + const __m256i _pA1 = _mm256_alignr_epi8(_pA0, _pA0, 8); + const __m256i _pB1 = _mm256_alignr_epi8(_pB0, _pB0, 4); +#if __AVXVNNIINT8__ + _sum0 = _mm256_dpbssd_epi32(_sum0, _pB0, _pA0); + _sum1 = _mm256_dpbssd_epi32(_sum1, _pB1, _pA0); + _sum2 = _mm256_dpbssd_epi32(_sum2, _pB0, _pA1); + _sum3 = _mm256_dpbssd_epi32(_sum3, _pB1, _pA1); +#else // __AVXVNNIINT8__ + _sum0 = _mm256_comp_dpbusd_epi32(_sum0, _pB0, _pA0); + _sum1 = _mm256_comp_dpbusd_epi32(_sum1, _pB1, _pA0); + _sum2 = _mm256_comp_dpbusd_epi32(_sum2, _pB0, _pA1); + _sum3 = _mm256_comp_dpbusd_epi32(_sum3, _pB1, _pA1); +#endif // __AVXVNNIINT8__ + pB += 16; + pA += 32; + } +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + if (max_kk >= 4) + { + const __m256i _shift0 = _mm256_loadu_si256((const __m256i*)pA); + const __m256i _shift1 = _mm256_alignr_epi8(_shift0, _shift0, 8); + _sum0 = _mm256_sub_epi32(_sum0, _shift0); + _sum1 = _mm256_sub_epi32(_sum1, _shift0); + _sum2 = _mm256_sub_epi32(_sum2, _shift1); + _sum3 = _mm256_sub_epi32(_sum3, _shift1); + pA += 32; + } +#endif +#endif // __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 1 < max_kk; kk += 2) + { + const __m128i _pA = _mm_loadu_si128((const __m128i*)pA); + const __m128i _pB = _mm_castpd_si128(_mm_load1_pd((const double*)pB)); + const __m256i _pA0 = _mm256_cvtepi8_epi16(_pA); + const __m256i _pA1 = _mm256_alignr_epi8(_pA0, _pA0, 8); + const __m256i _pB0 = _mm256_cvtepi8_epi16(_pB); + const __m256i _pB1 = _mm256_alignr_epi8(_pB0, _pB0, 4); + _sum0 = _mm256_comp_dpwssd_epi32(_sum0, _pA0, _pB0); + _sum1 = _mm256_comp_dpwssd_epi32(_sum1, _pA0, _pB1); + _sum2 = _mm256_comp_dpwssd_epi32(_sum2, _pA1, _pB0); + _sum3 = _mm256_comp_dpwssd_epi32(_sum3, _pA1, _pB1); + pB += 8; + pA += 16; + } + for (; kk < max_kk; kk++) + { + __m128i _pA0 = _mm_loadl_epi64((const __m128i*)pA); + __m128i _pB0 = _mm_castps_si128(_mm_load1_ps((const float*)pB)); + _pA0 = _mm_cvtepi8_epi16(_pA0); + _pB0 = _mm_cvtepi8_epi16(_pB0); + const __m128i _pA1 = _mm_shuffle_epi32(_pA0, _MM_SHUFFLE(2, 3, 0, 1)); + const __m128i _pB1 = _mm_shufflehi_epi16(_mm_shufflelo_epi16(_pB0, _MM_SHUFFLE(0, 3, 2, 1)), _MM_SHUFFLE(0, 3, 2, 1)); + _sum0 = _mm256_add_epi32(_sum0, _mm256_cvtepi16_epi32(_mm_mullo_epi16(_pA0, _pB0))); + _sum1 = _mm256_add_epi32(_sum1, _mm256_cvtepi16_epi32(_mm_mullo_epi16(_pA0, _pB1))); + _sum2 = _mm256_add_epi32(_sum2, _mm256_cvtepi16_epi32(_mm_mullo_epi16(_pA1, _pB0))); + _sum3 = _mm256_add_epi32(_sum3, _mm256_cvtepi16_epi32(_mm_mullo_epi16(_pA1, _pB1))); + pB += 4; + pA += 8; + } + + const __m256 _A0 = _mm256_loadu_ps(pA_descales); + const __m256 _A1 = _mm256_castsi256_ps(_mm256_alignr_epi8(_mm256_castps_si256(_A0), _mm256_castps_si256(_A0), 8)); + const __m128 _b = _mm_loadu_ps(pB_descales); + const __m256 _B0 = combine4x2_ps(_b, _b); + const __m256 _B1 = _mm256_castsi256_ps(_mm256_alignr_epi8(_mm256_castps_si256(_B0), _mm256_castps_si256(_B0), 4)); + _fsum0 = _mm256_add_ps(_fsum0, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum0), _mm256_mul_ps(_A0, _B0))); + _fsum1 = _mm256_add_ps(_fsum1, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum1), _mm256_mul_ps(_A0, _B1))); + _fsum2 = _mm256_add_ps(_fsum2, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum2), _mm256_mul_ps(_A1, _B0))); + _fsum3 = _mm256_add_ps(_fsum3, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum3), _mm256_mul_ps(_A1, _B1))); + pA_descales += 8; + pB_descales += 4; + } + + _mm256_storeu_ps(outptr + 0, _fsum0); + _mm256_storeu_ps(outptr + 8, _fsum1); + _mm256_storeu_ps(outptr + 16, _fsum2); + _mm256_storeu_ps(outptr + 24, _fsum3); + outptr += 32; + } +#endif // defined(__x86_64__) || defined(_M_X64) + for (; jj + 1 < max_jj; jj += 2) + { + __m256 _fsum0 = _mm256_setzero_ps(); + __m256 _fsum1 = _mm256_setzero_ps(); + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + __m256i _sum0 = _mm256_setzero_si256(); + __m256i _sum1 = _mm256_setzero_si256(); + const int max_kk = std::min(K - k, block_size); + int kk = 0; +#if __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 3 < max_kk; kk += 4) + { + const __m256i _pA0 = _mm256_loadu_si256((const __m256i*)pA); + const __m256i _pB0 = _mm256_castpd_si256(_mm256_broadcast_sd((const double*)pB)); + const __m256i _pB1 = _mm256_alignr_epi8(_pB0, _pB0, 4); +#if __AVXVNNIINT8__ + _sum0 = _mm256_dpbssd_epi32(_sum0, _pB0, _pA0); + _sum1 = _mm256_dpbssd_epi32(_sum1, _pB1, _pA0); +#else // __AVXVNNIINT8__ + _sum0 = _mm256_comp_dpbusd_epi32(_sum0, _pB0, _pA0); + _sum1 = _mm256_comp_dpbusd_epi32(_sum1, _pB1, _pA0); +#endif // __AVXVNNIINT8__ + pB += 8; + pA += 32; + } +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + if (max_kk >= 4) + { + const __m256i _shift0 = _mm256_loadu_si256((const __m256i*)pA); + _sum0 = _mm256_sub_epi32(_sum0, _shift0); + _sum1 = _mm256_sub_epi32(_sum1, _shift0); + pA += 32; + } +#endif +#else + for (; kk + 3 < max_kk; kk += 4) + { + const __m256i _pA = _mm256_loadu_si256((const __m256i*)pA); + const __m128i _pB = _mm_loadl_epi64((const __m128i*)pB); + const __m256i _pA01 = _mm256_cvtepi8_epi16(_mm256_castsi256_si128(_pA)); + const __m256i _pA23 = _mm256_cvtepi8_epi16(_mm256_extracti128_si256(_pA, 1)); + const __m256i _pB01 = _mm256_cvtepi8_epi16(_mm_shuffle_epi32(_pB, _MM_SHUFFLE(0, 0, 0, 0))); + const __m256i _pB23 = _mm256_cvtepi8_epi16(_mm_shuffle_epi32(_pB, _MM_SHUFFLE(1, 1, 1, 1))); + const __m256i _pB01_1 = _mm256_alignr_epi8(_pB01, _pB01, 4); + const __m256i _pB23_1 = _mm256_alignr_epi8(_pB23, _pB23, 4); + _sum0 = _mm256_comp_dpwssd_epi32(_sum0, _pA01, _pB01); + _sum1 = _mm256_comp_dpwssd_epi32(_sum1, _pA01, _pB01_1); + _sum0 = _mm256_comp_dpwssd_epi32(_sum0, _pA23, _pB23); + _sum1 = _mm256_comp_dpwssd_epi32(_sum1, _pA23, _pB23_1); + pB += 8; + pA += 32; + } +#endif // __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 1 < max_kk; kk += 2) + { + const __m128i _pA = _mm_loadu_si128((const __m128i*)pA); + const __m128i _pB = _mm_castps_si128(_mm_load1_ps((const float*)pB)); + const __m256i _pA0 = _mm256_cvtepi8_epi16(_pA); + const __m256i _pB0 = _mm256_cvtepi8_epi16(_pB); + const __m256i _pB1 = _mm256_alignr_epi8(_pB0, _pB0, 4); + _sum0 = _mm256_comp_dpwssd_epi32(_sum0, _pA0, _pB0); + _sum1 = _mm256_comp_dpwssd_epi32(_sum1, _pA0, _pB1); + pB += 4; + pA += 16; + } + for (; kk < max_kk; kk++) + { + __m128i _pA = _mm_loadl_epi64((const __m128i*)pA); + __m128i _pB0 = _mm_set1_epi16(*(const short*)pB); + _pA = _mm_cvtepi8_epi16(_pA); + _pB0 = _mm_cvtepi8_epi16(_pB0); + const __m128i _pB1 = _mm_shufflehi_epi16(_mm_shufflelo_epi16(_pB0, _MM_SHUFFLE(0, 1, 0, 1)), _MM_SHUFFLE(0, 1, 0, 1)); + _sum0 = _mm256_add_epi32(_sum0, _mm256_cvtepi16_epi32(_mm_mullo_epi16(_pA, _pB0))); + _sum1 = _mm256_add_epi32(_sum1, _mm256_cvtepi16_epi32(_mm_mullo_epi16(_pA, _pB1))); + pB += 2; + pA += 8; + } + + const __m256 _A0 = _mm256_loadu_ps(pA_descales); + const __m256 _B0 = _mm256_castpd_ps(_mm256_broadcast_sd((const double*)pB_descales)); + const __m256 _B1 = _mm256_castsi256_ps(_mm256_alignr_epi8(_mm256_castps_si256(_B0), _mm256_castps_si256(_B0), 4)); + _fsum0 = _mm256_add_ps(_fsum0, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum0), _mm256_mul_ps(_A0, _B0))); + _fsum1 = _mm256_add_ps(_fsum1, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum1), _mm256_mul_ps(_A0, _B1))); + pA_descales += 8; + pB_descales += 2; + } + + _mm256_storeu_ps(outptr + 0, _fsum0); + _mm256_storeu_ps(outptr + 8, _fsum1); + outptr += 16; + } + for (; jj < max_jj; jj++) + { + __m256 _fsum0 = _mm256_setzero_ps(); + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + __m256i _sum0 = _mm256_setzero_si256(); + const int max_kk = std::min(K - k, block_size); + int kk = 0; +#if __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 3 < max_kk; kk += 4) + { + const __m256i _pA0 = _mm256_loadu_si256((const __m256i*)pA); + const __m256i _pB0 = _mm256_castps_si256(_mm256_broadcast_ss((const float*)pB)); +#if __AVXVNNIINT8__ + _sum0 = _mm256_dpbssd_epi32(_sum0, _pB0, _pA0); +#else // __AVXVNNIINT8__ + _sum0 = _mm256_comp_dpbusd_epi32(_sum0, _pB0, _pA0); +#endif // __AVXVNNIINT8__ + pB += 4; + pA += 32; + } +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + if (max_kk >= 4) + { + const __m256i _shift0 = _mm256_loadu_si256((const __m256i*)pA); + _sum0 = _mm256_sub_epi32(_sum0, _shift0); + pA += 32; + } +#endif +#else + for (; kk + 3 < max_kk; kk += 4) + { + const __m256i _pA = _mm256_loadu_si256((const __m256i*)pA); + const __m256i _pA01 = _mm256_cvtepi8_epi16(_mm256_castsi256_si128(_pA)); + const __m256i _pA23 = _mm256_cvtepi8_epi16(_mm256_extracti128_si256(_pA, 1)); + const __m128i _pB16 = _mm_cvtepi8_epi16(_mm_cvtsi32_si128(*(const int*)pB)); + const __m256i _pB01 = _mm256_broadcastsi128_si256(_mm_shuffle_epi32(_pB16, _MM_SHUFFLE(0, 0, 0, 0))); + const __m256i _pB23 = _mm256_broadcastsi128_si256(_mm_shuffle_epi32(_pB16, _MM_SHUFFLE(1, 1, 1, 1))); + _sum0 = _mm256_comp_dpwssd_epi32(_sum0, _pA01, _pB01); + _sum0 = _mm256_comp_dpwssd_epi32(_sum0, _pA23, _pB23); + pB += 4; + pA += 32; + } +#endif // __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 1 < max_kk; kk += 2) + { + const __m128i _pA = _mm_loadu_si128((const __m128i*)pA); + const __m256i _pA0 = _mm256_cvtepi8_epi16(_pA); + const __m256i _pB0 = _mm256_cvtepi8_epi16(_mm_set1_epi16(*(const short*)pB)); + _sum0 = _mm256_comp_dpwssd_epi32(_sum0, _pA0, _pB0); + pB += 2; + pA += 16; + } + for (; kk < max_kk; kk++) + { + __m128i _pA = _mm_loadl_epi64((const __m128i*)pA); + _pA = _mm_cvtepi8_epi16(_pA); + _sum0 = _mm256_add_epi32(_sum0, _mm256_cvtepi16_epi32(_mm_mullo_epi16(_pA, _mm_set1_epi16((signed char)pB[0])))); + pB += 1; + pA += 8; + } + + const __m256 _A0 = _mm256_loadu_ps(pA_descales); + _fsum0 = _mm256_add_ps(_fsum0, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum0), _mm256_mul_ps(_A0, _mm256_set1_ps(pB_descales[0])))); + pA_descales += 8; + pB_descales++; + } + + _mm256_storeu_ps(outptr, _fsum0); + outptr += 8; + } + + pAT += A_hstep * 8; + pAT_descales += A_descales_hstep * 8; + } +#endif // __AVX2__ + for (; ii + 3 < max_ii; ii += 4) + { + const signed char* pB = pBT; + const float* pB_descales = pBT_descales; + + int jj = 0; +#if defined(__x86_64__) || defined(_M_X64) +#if __AVX512F__ + for (; jj + 7 < max_jj; jj += 8) + { + __m256 _fsum0 = _mm256_setzero_ps(); + __m256 _fsum1 = _mm256_setzero_ps(); + __m256 _fsum2 = _mm256_setzero_ps(); + __m256 _fsum3 = _mm256_setzero_ps(); + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + __m256i _sum0 = _mm256_setzero_si256(); + __m256i _sum1 = _mm256_setzero_si256(); + __m256i _sum2 = _mm256_setzero_si256(); + __m256i _sum3 = _mm256_setzero_si256(); + const int max_kk = std::min(K - k, block_size); + int kk = 0; +#if __AVX512VNNI__ + for (; kk + 3 < max_kk; kk += 4) + { + const __m256i _pA0 = _mm256_broadcastsi128_si256(_mm_loadu_si128((const __m128i*)pA)); + const __m256i _pA1 = _mm256_alignr_epi8(_pA0, _pA0, 8); + const __m256i _pB0 = _mm256_loadu_si256((const __m256i*)pB); + const __m256i _pB1 = _mm256_alignr_epi8(_pB0, _pB0, 4); + _sum0 = _mm256_comp_dpbusd_epi32(_sum0, _pB0, _pA0); + _sum1 = _mm256_comp_dpbusd_epi32(_sum1, _pB1, _pA0); + _sum2 = _mm256_comp_dpbusd_epi32(_sum2, _pB0, _pA1); + _sum3 = _mm256_comp_dpbusd_epi32(_sum3, _pB1, _pA1); + pA += 16; + pB += 32; + } + if (max_kk >= 4) + { + const __m256i _shift0 = _mm256_broadcastsi128_si256(_mm_loadu_si128((const __m128i*)pA)); + const __m256i _shift1 = _mm256_alignr_epi8(_shift0, _shift0, 8); + _sum0 = _mm256_sub_epi32(_sum0, _shift0); + _sum1 = _mm256_sub_epi32(_sum1, _shift0); + _sum2 = _mm256_sub_epi32(_sum2, _shift1); + _sum3 = _mm256_sub_epi32(_sum3, _shift1); + pA += 16; + } +#endif // __AVX512VNNI__ + for (; kk + 1 < max_kk; kk += 2) + { + const __m128i _pA8x1 = _mm_loadl_epi64((const __m128i*)pA); + const __m128i _pA8 = _mm_unpacklo_epi64(_pA8x1, _pA8x1); + const __m128i _pB8 = _mm_loadu_si128((const __m128i*)pB); + const __m256i _pA0 = _mm256_cvtepi8_epi16(_pA8); + const __m256i _pA1 = _mm256_alignr_epi8(_pA0, _pA0, 8); + const __m256i _pB0 = _mm256_cvtepi8_epi16(_pB8); + const __m256i _pB1 = _mm256_alignr_epi8(_pB0, _pB0, 4); + _sum0 = _mm256_comp_dpwssd_epi32(_sum0, _pA0, _pB0); + _sum1 = _mm256_comp_dpwssd_epi32(_sum1, _pA0, _pB1); + _sum2 = _mm256_comp_dpwssd_epi32(_sum2, _pA1, _pB0); + _sum3 = _mm256_comp_dpwssd_epi32(_sum3, _pA1, _pB1); + pA += 8; + pB += 16; + } + for (; kk < max_kk; kk++) + { + const __m128i _pA8 = _mm_cvtsi32_si128(*(const int*)pA); + const __m128i _pB8 = _mm_loadl_epi64((const __m128i*)pB); + const __m128i _pA32 = _mm_cvtepi8_epi32(_pA8); + const __m256i _pA0 = combine4x2_epi32(_pA32, _pA32); + const __m256i _pA1 = _mm256_shuffle_epi32(_pA0, _MM_SHUFFLE(1, 0, 3, 2)); + const __m256i _pB0 = combine4x2_epi32(_mm_cvtepi8_epi32(_pB8), _mm_cvtepi8_epi32(_mm_srli_si128(_pB8, 4))); + const __m256i _pB1 = _mm256_shuffle_epi32(_pB0, _MM_SHUFFLE(0, 3, 2, 1)); + _sum0 = _mm256_add_epi32(_sum0, _mm256_mullo_epi32(_pA0, _pB0)); + _sum1 = _mm256_add_epi32(_sum1, _mm256_mullo_epi32(_pA0, _pB1)); + _sum2 = _mm256_add_epi32(_sum2, _mm256_mullo_epi32(_pA1, _pB0)); + _sum3 = _mm256_add_epi32(_sum3, _mm256_mullo_epi32(_pA1, _pB1)); + pA += 4; + pB += 8; + } + + const __m128 _ad128 = _mm_loadu_ps(pA_descales); + const __m256 _ad0 = combine4x2_ps(_ad128, _ad128); + const __m256 _ad1 = _mm256_shuffle_ps(_ad0, _ad0, _MM_SHUFFLE(1, 0, 3, 2)); + const __m256 _bd0 = _mm256_loadu_ps(pB_descales); + const __m256 _bd1 = _mm256_shuffle_ps(_bd0, _bd0, _MM_SHUFFLE(0, 3, 2, 1)); + _fsum0 = _mm256_add_ps(_fsum0, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum0), _mm256_mul_ps(_ad0, _bd0))); + _fsum1 = _mm256_add_ps(_fsum1, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum1), _mm256_mul_ps(_ad0, _bd1))); + _fsum2 = _mm256_add_ps(_fsum2, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum2), _mm256_mul_ps(_ad1, _bd0))); + _fsum3 = _mm256_add_ps(_fsum3, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum3), _mm256_mul_ps(_ad1, _bd1))); + pA_descales += 4; + pB_descales += 8; + } + + _mm256_storeu_ps(outptr, _fsum0); + _mm256_storeu_ps(outptr + 8, _fsum1); + _mm256_storeu_ps(outptr + 16, _fsum2); + _mm256_storeu_ps(outptr + 24, _fsum3); + outptr += 32; + } +#endif // __AVX512F__ + for (; jj + 3 < max_jj; jj += 4) + { + __m128 _fsum0 = _mm_setzero_ps(); + __m128 _fsum1 = _mm_setzero_ps(); + __m128 _fsum2 = _mm_setzero_ps(); + __m128 _fsum3 = _mm_setzero_ps(); + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + __m128i _sum0 = _mm_setzero_si128(); + __m128i _sum1 = _mm_setzero_si128(); + __m128i _sum2 = _mm_setzero_si128(); + __m128i _sum3 = _mm_setzero_si128(); + const int max_kk = std::min(K - k, block_size); + int kk = 0; +#if __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 3 < max_kk; kk += 4) + { + const __m128i _pA0 = _mm_loadu_si128((const __m128i*)pA); + const __m128i _pA1 = _mm_alignr_epi8(_pA0, _pA0, 8); + const __m128i _pB0 = _mm_loadu_si128((const __m128i*)pB); + const __m128i _pB1 = _mm_alignr_epi8(_pB0, _pB0, 4); +#if __AVXVNNIINT8__ + _sum0 = _mm_dpbssd_epi32(_sum0, _pB0, _pA0); + _sum1 = _mm_dpbssd_epi32(_sum1, _pB1, _pA0); + _sum2 = _mm_dpbssd_epi32(_sum2, _pB0, _pA1); + _sum3 = _mm_dpbssd_epi32(_sum3, _pB1, _pA1); +#else // __AVXVNNIINT8__ + _sum0 = _mm_comp_dpbusd_epi32(_sum0, _pB0, _pA0); + _sum1 = _mm_comp_dpbusd_epi32(_sum1, _pB1, _pA0); + _sum2 = _mm_comp_dpbusd_epi32(_sum2, _pB0, _pA1); + _sum3 = _mm_comp_dpbusd_epi32(_sum3, _pB1, _pA1); +#endif // __AVXVNNIINT8__ + pA += 16; + pB += 16; + } +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + if (max_kk >= 4) + { + const __m128i _shift0 = _mm_loadu_si128((const __m128i*)pA); + const __m128i _shift1 = _mm_alignr_epi8(_shift0, _shift0, 8); + _sum0 = _mm_sub_epi32(_sum0, _shift0); + _sum1 = _mm_sub_epi32(_sum1, _shift0); + _sum2 = _mm_sub_epi32(_sum2, _shift1); + _sum3 = _mm_sub_epi32(_sum3, _shift1); + pA += 16; + } +#endif +#endif // __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 1 < max_kk; kk += 2) + { + const __m128i _pA8 = _mm_loadl_epi64((const __m128i*)pA); + const __m128i _pB8 = _mm_loadl_epi64((const __m128i*)pB); + const __m128i _pA0 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + const __m128i _pA1 = _mm_shuffle_epi32(_pA0, _MM_SHUFFLE(1, 0, 3, 2)); + const __m128i _pB0 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); + const __m128i _pB1 = _mm_shuffle_epi32(_pB0, _MM_SHUFFLE(0, 3, 2, 1)); + _sum0 = _mm_comp_dpwssd_epi32(_sum0, _pA0, _pB0); + _sum1 = _mm_comp_dpwssd_epi32(_sum1, _pA0, _pB1); + _sum2 = _mm_comp_dpwssd_epi32(_sum2, _pA1, _pB0); + _sum3 = _mm_comp_dpwssd_epi32(_sum3, _pA1, _pB1); + pA += 8; + pB += 8; + } + for (; kk < max_kk; kk++) + { + const __m128i _pA8 = _mm_cvtsi32_si128(*(const int*)pA); + const __m128i _pB8 = _mm_cvtsi32_si128(*(const int*)pB); + const __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + const __m128i _pB16 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); + const __m128i _pA0 = _mm_unpacklo_epi16(_pA16, _pA16); + const __m128i _pA1 = _mm_shuffle_epi32(_pA0, _MM_SHUFFLE(1, 0, 3, 2)); + const __m128i _pB0 = _mm_unpacklo_epi16(_pB16, _mm_setzero_si128()); + const __m128i _pB1 = _mm_shuffle_epi32(_pB0, _MM_SHUFFLE(0, 3, 2, 1)); + _sum0 = _mm_comp_dpwssd_epi32(_sum0, _pA0, _pB0); + _sum1 = _mm_comp_dpwssd_epi32(_sum1, _pA0, _pB1); + _sum2 = _mm_comp_dpwssd_epi32(_sum2, _pA1, _pB0); + _sum3 = _mm_comp_dpwssd_epi32(_sum3, _pA1, _pB1); + pA += 4; + pB += 4; + } + + const __m128 _ad0 = _mm_loadu_ps(pA_descales); + const __m128 _ad1 = _mm_shuffle_ps(_ad0, _ad0, _MM_SHUFFLE(1, 0, 3, 2)); + const __m128 _bd0 = _mm_loadu_ps(pB_descales); + const __m128 _bd1 = _mm_shuffle_ps(_bd0, _bd0, _MM_SHUFFLE(0, 3, 2, 1)); + _fsum0 = _mm_add_ps(_fsum0, _mm_mul_ps(_mm_cvtepi32_ps(_sum0), _mm_mul_ps(_ad0, _bd0))); + _fsum1 = _mm_add_ps(_fsum1, _mm_mul_ps(_mm_cvtepi32_ps(_sum1), _mm_mul_ps(_ad0, _bd1))); + _fsum2 = _mm_add_ps(_fsum2, _mm_mul_ps(_mm_cvtepi32_ps(_sum2), _mm_mul_ps(_ad1, _bd0))); + _fsum3 = _mm_add_ps(_fsum3, _mm_mul_ps(_mm_cvtepi32_ps(_sum3), _mm_mul_ps(_ad1, _bd1))); + pA_descales += 4; + pB_descales += 4; + } + + _mm_storeu_ps(outptr, _fsum0); + _mm_storeu_ps(outptr + 4, _fsum1); + _mm_storeu_ps(outptr + 8, _fsum2); + _mm_storeu_ps(outptr + 12, _fsum3); + outptr += 16; + } +#endif // defined(__x86_64__) || defined(_M_X64) + for (; jj + 1 < max_jj; jj += 2) + { + __m128 _fsum0 = _mm_setzero_ps(); + __m128 _fsum1 = _mm_setzero_ps(); + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + __m128i _sum0 = _mm_setzero_si128(); + __m128i _sum1 = _mm_setzero_si128(); + const int max_kk = std::min(K - k, block_size); + int kk = 0; +#if __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 3 < max_kk; kk += 4) + { + const __m128i _pA0 = _mm_loadu_si128((const __m128i*)pA); + const __m128i _pB8 = _mm_loadl_epi64((const __m128i*)pB); + const __m128i _pB0 = _mm_unpacklo_epi64(_pB8, _pB8); + const __m128i _pB1 = _mm_alignr_epi8(_pB0, _pB0, 4); +#if __AVXVNNIINT8__ + _sum0 = _mm_dpbssd_epi32(_sum0, _pB0, _pA0); + _sum1 = _mm_dpbssd_epi32(_sum1, _pB1, _pA0); +#else // __AVXVNNIINT8__ + _sum0 = _mm_comp_dpbusd_epi32(_sum0, _pB0, _pA0); + _sum1 = _mm_comp_dpbusd_epi32(_sum1, _pB1, _pA0); +#endif // __AVXVNNIINT8__ + pA += 16; + pB += 8; + } +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + if (max_kk >= 4) + { + const __m128i _shift0 = _mm_loadu_si128((const __m128i*)pA); + _sum0 = _mm_sub_epi32(_sum0, _shift0); + _sum1 = _mm_sub_epi32(_sum1, _shift0); + pA += 16; + } +#endif +#endif // __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 1 < max_kk; kk += 2) + { + const __m128i _pA8 = _mm_loadl_epi64((const __m128i*)pA); + const __m128i _pB8 = _mm_castps_si128(_mm_load1_ps((const float*)pB)); + const __m128i _pA0 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + const __m128i _pB0 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); + const __m128i _pB1 = _mm_shuffle_epi32(_pB0, _MM_SHUFFLE(0, 3, 2, 1)); + _sum0 = _mm_comp_dpwssd_epi32(_sum0, _pA0, _pB0); + _sum1 = _mm_comp_dpwssd_epi32(_sum1, _pA0, _pB1); + pA += 8; + pB += 4; + } + for (; kk < max_kk; kk++) + { + const __m128i _pA8 = _mm_cvtsi32_si128(*(const int*)pA); + const __m128i _pB8 = _mm_set1_epi16(*(const short*)pB); + const __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + const __m128i _pB16 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); + const __m128i _pA0 = _mm_unpacklo_epi16(_pA16, _pA16); + const __m128i _pB0 = _mm_unpacklo_epi16(_pB16, _mm_setzero_si128()); + const __m128i _pB1 = _mm_shuffle_epi32(_pB0, _MM_SHUFFLE(0, 3, 2, 1)); + _sum0 = _mm_comp_dpwssd_epi32(_sum0, _pA0, _pB0); + _sum1 = _mm_comp_dpwssd_epi32(_sum1, _pA0, _pB1); + pA += 4; + pB += 2; + } + + const __m128 _ad = _mm_loadu_ps(pA_descales); + const __m128 _bd0 = _mm_setr_ps(pB_descales[0], pB_descales[1], pB_descales[0], pB_descales[1]); + const __m128 _bd1 = _mm_shuffle_ps(_bd0, _bd0, _MM_SHUFFLE(0, 3, 2, 1)); + _fsum0 = _mm_add_ps(_fsum0, _mm_mul_ps(_mm_cvtepi32_ps(_sum0), _mm_mul_ps(_ad, _bd0))); + _fsum1 = _mm_add_ps(_fsum1, _mm_mul_ps(_mm_cvtepi32_ps(_sum1), _mm_mul_ps(_ad, _bd1))); + pA_descales += 4; + pB_descales += 2; + } + + _mm_storeu_ps(outptr, _fsum0); + _mm_storeu_ps(outptr + 4, _fsum1); + outptr += 8; + } + for (; jj < max_jj; jj++) + { + __m128 _fsum = _mm_setzero_ps(); + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + __m128i _sum = _mm_setzero_si128(); + const int max_kk = std::min(K - k, block_size); + int kk = 0; +#if __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 3 < max_kk; kk += 4) + { + const __m128i _pA = _mm_loadu_si128((const __m128i*)pA); + const __m128i _pB = _mm_set1_epi32(*(const int*)pB); +#if __AVXVNNIINT8__ + _sum = _mm_dpbssd_epi32(_sum, _pB, _pA); +#else // __AVXVNNIINT8__ + _sum = _mm_comp_dpbusd_epi32(_sum, _pB, _pA); +#endif // __AVXVNNIINT8__ + pA += 16; + pB += 4; + } +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + if (max_kk >= 4) + { + _sum = _mm_sub_epi32(_sum, _mm_loadu_si128((const __m128i*)pA)); + pA += 16; + } +#endif +#endif // __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 1 < max_kk; kk += 2) + { + const __m128i _pA8 = _mm_loadl_epi64((const __m128i*)pA); + const __m128i _pB8 = _mm_set1_epi16(*(const short*)pB); + const __m128i _pA = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + const __m128i _pB = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); + _sum = _mm_comp_dpwssd_epi32(_sum, _pA, _pB); + pA += 8; + pB += 2; + } + for (; kk < max_kk; kk++) + { + const __m128i _pA8 = _mm_cvtsi32_si128(*(const int*)pA); + const __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + const __m128i _pA = _mm_unpacklo_epi16(_pA16, _mm_setzero_si128()); + _sum = _mm_comp_dpwssd_epi32(_sum, _pA, _mm_set1_epi16((signed char)pB[0])); + pA += 4; + pB++; + } + + const __m128 _ad = _mm_loadu_ps(pA_descales); + _fsum = _mm_add_ps(_fsum, _mm_mul_ps(_mm_cvtepi32_ps(_sum), _mm_mul_ps(_ad, _mm_set1_ps(pB_descales[0])))); + pA_descales += 4; + pB_descales++; + } + + _mm_storeu_ps(outptr, _fsum); + outptr += 4; + } + + pAT += A_hstep * 4; + pAT_descales += A_descales_hstep * 4; + } +#endif // __SSE2__ + for (; ii + 1 < max_ii; ii += 2) + { + const signed char* pB = pBT; + const float* pB_descales = pBT_descales; + + int jj = 0; +#if __SSE2__ +#if defined(__x86_64__) || defined(_M_X64) +#if __AVX512F__ + for (; jj + 7 < max_jj; jj += 8) + { + __m256 _fsum0 = _mm256_setzero_ps(); + __m256 _fsum1 = _mm256_setzero_ps(); + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + __m256i _sum0 = _mm256_setzero_si256(); + __m256i _sum1 = _mm256_setzero_si256(); + const int max_kk = std::min(K - k, block_size); + int kk = 0; +#if __AVX512VNNI__ + for (; kk + 3 < max_kk; kk += 4) + { + const __m128i _pA8 = _mm_loadl_epi64((const __m128i*)pA); + const __m128i _pA128 = _mm_unpacklo_epi64(_pA8, _pA8); + const __m256i _pA0 = _mm256_broadcastsi128_si256(_pA128); + const __m256i _pA1 = _mm256_shuffle_epi32(_pA0, _MM_SHUFFLE(2, 3, 0, 1)); + const __m256i _pB = _mm256_loadu_si256((const __m256i*)pB); + _sum0 = _mm256_comp_dpbusd_epi32(_sum0, _pB, _pA0); + _sum1 = _mm256_comp_dpbusd_epi32(_sum1, _pB, _pA1); + pA += 8; + pB += 32; + } + if (max_kk >= 4) + { + const __m128i _shift64 = _mm_loadl_epi64((const __m128i*)pA); + const __m128i _shift128 = _mm_unpacklo_epi64(_shift64, _shift64); + const __m256i _shift0 = _mm256_broadcastsi128_si256(_shift128); + const __m256i _shift1 = _mm256_shuffle_epi32(_shift0, _MM_SHUFFLE(2, 3, 0, 1)); + _sum0 = _mm256_sub_epi32(_sum0, _shift0); + _sum1 = _mm256_sub_epi32(_sum1, _shift1); + pA += 8; + } +#endif // __AVX512VNNI__ + for (; kk + 1 < max_kk; kk += 2) + { + const __m128i _pA8 = _mm_cvtsi32_si128(*(const int*)pA); + const __m128i _pA16x1 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + const __m128i _pA16 = _mm_unpacklo_epi64(_pA16x1, _pA16x1); + const __m256i _pA0 = _mm256_broadcastsi128_si256(_pA16); + const __m256i _pA1 = _mm256_shuffle_epi32(_pA0, _MM_SHUFFLE(2, 3, 0, 1)); + const __m256i _pB = _mm256_cvtepi8_epi16(_mm_loadu_si128((const __m128i*)pB)); + _sum0 = _mm256_comp_dpwssd_epi32(_sum0, _pA0, _pB); + _sum1 = _mm256_comp_dpwssd_epi32(_sum1, _pA1, _pB); + pA += 4; + pB += 16; + } + for (; kk < max_kk; kk++) + { + const __m128i _pA8 = _mm_cvtsi32_si128(*(const unsigned short*)pA); + const __m128i _pA32x1 = _mm_cvtepi8_epi32(_pA8); + const __m128i _pA128 = _mm_shuffle_epi32(_pA32x1, _MM_SHUFFLE(1, 0, 1, 0)); + const __m256i _pA0 = _mm256_broadcastsi128_si256(_pA128); + const __m256i _pA1 = _mm256_shuffle_epi32(_pA0, _MM_SHUFFLE(2, 3, 0, 1)); + const __m256i _pB = _mm256_cvtepi8_epi32(_mm_loadl_epi64((const __m128i*)pB)); + _sum0 = _mm256_add_epi32(_sum0, _mm256_mullo_epi32(_pA0, _pB)); + _sum1 = _mm256_add_epi32(_sum1, _mm256_mullo_epi32(_pA1, _pB)); + pA += 2; + pB += 8; + } + + const __m128 _ad2 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pA_descales); + const __m128 _ad128 = _mm_movelh_ps(_ad2, _ad2); + const __m256 _ad0 = combine4x2_ps(_ad128, _ad128); + const __m256 _ad1 = _mm256_shuffle_ps(_ad0, _ad0, _MM_SHUFFLE(2, 3, 0, 1)); + const __m256 _bd = _mm256_loadu_ps(pB_descales); + _fsum0 = _mm256_add_ps(_fsum0, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum0), _mm256_mul_ps(_ad0, _bd))); + _fsum1 = _mm256_add_ps(_fsum1, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum1), _mm256_mul_ps(_ad1, _bd))); + pA_descales += 2; + pB_descales += 8; + } + + _mm256_storeu_ps(outptr, _fsum0); + _mm256_storeu_ps(outptr + 8, _fsum1); + outptr += 16; + } +#endif // __AVX512F__ + for (; jj + 3 < max_jj; jj += 4) + { + __m128 _fsum0 = _mm_setzero_ps(); + __m128 _fsum1 = _mm_setzero_ps(); + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + __m128i _sum0 = _mm_setzero_si128(); + __m128i _sum1 = _mm_setzero_si128(); + const int max_kk = std::min(K - k, block_size); + int kk = 0; +#if __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 3 < max_kk; kk += 4) + { + const __m128i _pA8 = _mm_loadl_epi64((const __m128i*)pA); + const __m128i _pA = _mm_unpacklo_epi64(_pA8, _pA8); + const __m128i _pB0 = _mm_loadu_si128((const __m128i*)pB); + const __m128i _pB1 = _mm_shuffle_epi32(_pB0, _MM_SHUFFLE(0, 3, 2, 1)); +#if __AVXVNNIINT8__ + _sum0 = _mm_dpbssd_epi32(_sum0, _pB0, _pA); + _sum1 = _mm_dpbssd_epi32(_sum1, _pB1, _pA); +#else // __AVXVNNIINT8__ + _sum0 = _mm_comp_dpbusd_epi32(_sum0, _pB0, _pA); + _sum1 = _mm_comp_dpbusd_epi32(_sum1, _pB1, _pA); +#endif // __AVXVNNIINT8__ + pA += 8; + pB += 16; + } +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + if (max_kk >= 4) + { + const __m128i _shift64 = _mm_loadl_epi64((const __m128i*)pA); + const __m128i _shift = _mm_unpacklo_epi64(_shift64, _shift64); + _sum0 = _mm_sub_epi32(_sum0, _shift); + _sum1 = _mm_sub_epi32(_sum1, _shift); + pA += 8; + } +#endif +#endif // __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 1 < max_kk; kk += 2) + { + const __m128i _pA8 = _mm_cvtsi32_si128(*(const int*)pA); + const __m128i _pA16x1 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + const __m128i _pA = _mm_unpacklo_epi64(_pA16x1, _pA16x1); + const __m128i _pB8 = _mm_loadl_epi64((const __m128i*)pB); + const __m128i _pB0 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); + const __m128i _pB1 = _mm_shuffle_epi32(_pB0, _MM_SHUFFLE(0, 3, 2, 1)); + _sum0 = _mm_comp_dpwssd_epi32(_sum0, _pA, _pB0); + _sum1 = _mm_comp_dpwssd_epi32(_sum1, _pA, _pB1); + pA += 4; + pB += 8; + } + for (; kk < max_kk; kk++) + { + const __m128i _pA8 = _mm_cvtsi32_si128(*(const unsigned short*)pA); + const __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + const __m128i _pA32x1 = _mm_unpacklo_epi16(_pA16, _pA16); + const __m128i _pA = _mm_shuffle_epi32(_pA32x1, _MM_SHUFFLE(1, 0, 1, 0)); + const __m128i _pB8 = _mm_cvtsi32_si128(*(const int*)pB); + const __m128i _pB16 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); + const __m128i _pB0 = _mm_unpacklo_epi16(_pB16, _mm_setzero_si128()); + const __m128i _pB1 = _mm_shuffle_epi32(_pB0, _MM_SHUFFLE(0, 3, 2, 1)); + _sum0 = _mm_comp_dpwssd_epi32(_sum0, _pA, _pB0); + _sum1 = _mm_comp_dpwssd_epi32(_sum1, _pA, _pB1); + pA += 2; + pB += 4; + } + + const __m128 _ad2 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pA_descales); + const __m128 _ad = _mm_movelh_ps(_ad2, _ad2); + const __m128 _bd0 = _mm_loadu_ps(pB_descales); + const __m128 _bd1 = _mm_shuffle_ps(_bd0, _bd0, _MM_SHUFFLE(0, 3, 2, 1)); + _fsum0 = _mm_add_ps(_fsum0, _mm_mul_ps(_mm_cvtepi32_ps(_sum0), _mm_mul_ps(_ad, _bd0))); + _fsum1 = _mm_add_ps(_fsum1, _mm_mul_ps(_mm_cvtepi32_ps(_sum1), _mm_mul_ps(_ad, _bd1))); + pA_descales += 2; + pB_descales += 4; + } + + _mm_storeu_ps(outptr, _fsum0); + _mm_storeu_ps(outptr + 4, _fsum1); + outptr += 8; + } +#endif // defined(__x86_64__) || defined(_M_X64) +#endif // __SSE2__ + for (; jj + 1 < max_jj; jj += 2) + { +#if __SSE2__ + __m128 _fsum = _mm_setzero_ps(); + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + __m128i _sum = _mm_setzero_si128(); + const int max_kk = std::min(K - k, block_size); + int kk = 0; +#if __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 3 < max_kk; kk += 4) + { + const __m128i _pA8 = _mm_loadl_epi64((const __m128i*)pA); + const __m128i _pA = _mm_unpacklo_epi32(_pA8, _pA8); + const __m128i _pB8 = _mm_loadl_epi64((const __m128i*)pB); + const __m128i _pB = _mm_unpacklo_epi64(_pB8, _pB8); +#if __AVXVNNIINT8__ + _sum = _mm_dpbssd_epi32(_sum, _pB, _pA); +#else // __AVXVNNIINT8__ + _sum = _mm_comp_dpbusd_epi32(_sum, _pB, _pA); +#endif // __AVXVNNIINT8__ + pA += 8; + pB += 8; + } +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + if (max_kk >= 4) + { + const __m128i _shift64 = _mm_loadl_epi64((const __m128i*)pA); + const __m128i _shift = _mm_unpacklo_epi32(_shift64, _shift64); + _sum = _mm_sub_epi32(_sum, _shift); + pA += 8; + } +#endif +#endif // __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 1 < max_kk; kk += 2) + { + const __m128i _pA8 = _mm_cvtsi32_si128(*(const int*)pA); + const __m128i _pA16x1 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + const __m128i _pA = _mm_unpacklo_epi32(_pA16x1, _pA16x1); + const __m128i _pB8 = _mm_cvtsi32_si128(*(const int*)pB); + const __m128i _pB16x1 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); + const __m128i _pB = _mm_unpacklo_epi64(_pB16x1, _pB16x1); + _sum = _mm_comp_dpwssd_epi32(_sum, _pA, _pB); + pA += 4; + pB += 4; + } + for (; kk < max_kk; kk++) + { + const __m128i _pA8 = _mm_cvtsi32_si128(*(const unsigned short*)pA); + const __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + const __m128i _pA32x1 = _mm_unpacklo_epi16(_pA16, _pA16); + const __m128i _pA = _mm_unpacklo_epi32(_pA32x1, _pA32x1); + const __m128i _pB8 = _mm_cvtsi32_si128(*(const unsigned short*)pB); + const __m128i _pB16 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); + const __m128i _pB32x1 = _mm_unpacklo_epi16(_pB16, _mm_setzero_si128()); + const __m128i _pB = _mm_unpacklo_epi64(_pB32x1, _pB32x1); + _sum = _mm_comp_dpwssd_epi32(_sum, _pA, _pB); + pA += 2; + pB += 2; + } + + const __m128 _ad2 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pA_descales); + const __m128 _ad = _mm_unpacklo_ps(_ad2, _ad2); + const __m128 _bd2 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pB_descales); + const __m128 _bd = _mm_movelh_ps(_bd2, _bd2); + _fsum = _mm_add_ps(_fsum, _mm_mul_ps(_mm_cvtepi32_ps(_sum), _mm_mul_ps(_ad, _bd))); + pA_descales += 2; + pB_descales += 2; + } + + _mm_storeu_ps(outptr, _fsum); + outptr += 4; + +#else + float fsum00 = 0.f; + float fsum01 = 0.f; + float fsum10 = 0.f; + float fsum11 = 0.f; + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + int sum00 = 0; + int sum01 = 0; + int sum10 = 0; + int sum11 = 0; + const int max_kk = std::min(K - k, block_size); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + int b0 = (signed char)pB[0]; + int b1 = (signed char)pB[1]; + sum00 += pA[0] * b0; + sum01 += pA[0] * b1; + sum10 += pA[1] * b0; + sum11 += pA[1] * b1; + b0 = (signed char)pB[2]; + b1 = (signed char)pB[3]; + sum00 += pA[2] * b0; + sum01 += pA[2] * b1; + sum10 += pA[3] * b0; + sum11 += pA[3] * b1; + b0 = (signed char)pB[4]; + b1 = (signed char)pB[5]; + sum00 += pA[4] * b0; + sum01 += pA[4] * b1; + sum10 += pA[5] * b0; + sum11 += pA[5] * b1; + b0 = (signed char)pB[6]; + b1 = (signed char)pB[7]; + sum00 += pA[6] * b0; + sum01 += pA[6] * b1; + sum10 += pA[7] * b0; + sum11 += pA[7] * b1; + pA += 8; + pB += 8; + } + for (; kk < max_kk; kk++) + { + const int b0 = (signed char)pB[0]; + const int b1 = (signed char)pB[1]; + sum00 += pA[0] * b0; + sum01 += pA[0] * b1; + sum10 += pA[1] * b0; + sum11 += pA[1] * b1; + pA += 2; + pB += 2; + } + + const float ad0 = pA_descales[0]; + const float ad1 = pA_descales[1]; + fsum00 += sum00 * ad0 * pB_descales[0]; + fsum01 += sum01 * ad0 * pB_descales[1]; + fsum10 += sum10 * ad1 * pB_descales[0]; + fsum11 += sum11 * ad1 * pB_descales[1]; + pA_descales += 2; + pB_descales += 2; + } + + outptr[0] = fsum00; + outptr[1] = fsum01; + outptr[2] = fsum10; + outptr[3] = fsum11; + outptr += 4; +#endif // __SSE2__ + } + for (; jj < max_jj; jj++) + { +#if __SSE2__ + __m128 _fsum = _mm_setzero_ps(); + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + __m128i _sum = _mm_setzero_si128(); + const int max_kk = std::min(K - k, block_size); + int kk = 0; +#if __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 3 < max_kk; kk += 4) + { + const __m128i _pA = _mm_loadl_epi64((const __m128i*)pA); + const __m128i _pB8 = _mm_cvtsi32_si128(*(const int*)pB); + const __m128i _pB = _mm_shuffle_epi32(_pB8, _MM_SHUFFLE(0, 0, 0, 0)); +#if __AVXVNNIINT8__ + _sum = _mm_dpbssd_epi32(_sum, _pB, _pA); +#else // __AVXVNNIINT8__ + _sum = _mm_comp_dpbusd_epi32(_sum, _pB, _pA); +#endif // __AVXVNNIINT8__ + pA += 8; + pB += 4; + } +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + if (max_kk >= 4) + { + _sum = _mm_sub_epi32(_sum, _mm_loadl_epi64((const __m128i*)pA)); + pA += 8; + } +#endif +#endif // __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 1 < max_kk; kk += 2) + { + const __m128i _pA8 = _mm_cvtsi32_si128(*(const int*)pA); + const __m128i _pA = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + const __m128i _pB8 = _mm_cvtsi32_si128(*(const unsigned short*)pB); + const __m128i _pB16 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); + const __m128i _pB = _mm_shuffle_epi32(_pB16, _MM_SHUFFLE(0, 0, 0, 0)); + _sum = _mm_comp_dpwssd_epi32(_sum, _pA, _pB); + pA += 4; + pB += 2; + } + for (; kk < max_kk; kk++) + { + const __m128i _pA8 = _mm_cvtsi32_si128(*(const unsigned short*)pA); + const __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + const __m128i _pA = _mm_unpacklo_epi16(_pA16, _pA16); + const __m128i _pB8 = _mm_cvtsi32_si128((signed char)pB[0]); + const __m128i _pB16 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); + const __m128i _pB32 = _mm_unpacklo_epi16(_pB16, _mm_setzero_si128()); + const __m128i _pB = _mm_shuffle_epi32(_pB32, _MM_SHUFFLE(0, 0, 0, 0)); + _sum = _mm_comp_dpwssd_epi32(_sum, _pA, _pB); + pA += 2; + pB++; + } + + const __m128 _ad = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pA_descales); + _fsum = _mm_add_ps(_fsum, _mm_mul_ps(_mm_cvtepi32_ps(_sum), _mm_mul_ps(_ad, _mm_set1_ps(pB_descales[0])))); + pA_descales += 2; + pB_descales++; + } + + _mm_storel_pi((__m64*)outptr, _fsum); + outptr += 2; + +#else + float fsum0 = 0.f; + float fsum1 = 0.f; + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + int sum0 = 0; + int sum1 = 0; + const int max_kk = std::min(K - k, block_size); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + int b0 = (signed char)pB[0]; + sum0 += pA[0] * b0; + sum1 += pA[1] * b0; + b0 = (signed char)pB[1]; + sum0 += pA[2] * b0; + sum1 += pA[3] * b0; + b0 = (signed char)pB[2]; + sum0 += pA[4] * b0; + sum1 += pA[5] * b0; + b0 = (signed char)pB[3]; + sum0 += pA[6] * b0; + sum1 += pA[7] * b0; + pA += 8; + pB += 4; + } + for (; kk < max_kk; kk++) + { + const int b0 = (signed char)pB[0]; + sum0 += pA[0] * b0; + sum1 += pA[1] * b0; + pA += 2; + pB++; + } + + fsum0 += sum0 * pA_descales[0] * pB_descales[0]; + fsum1 += sum1 * pA_descales[1] * pB_descales[0]; + pA_descales += 2; + pB_descales++; + } + + outptr[0] = fsum0; + outptr[1] = fsum1; + outptr += 2; +#endif // __SSE2__ + } + + pAT += A_hstep * 2; + pAT_descales += A_descales_hstep * 2; + } + for (; ii < max_ii; ii++) + { + const signed char* pB = pBT; + const float* pB_descales = pBT_descales; + + int jj = 0; +#if __SSE2__ +#if defined(__x86_64__) || defined(_M_X64) +#if __AVX512F__ + for (; jj + 7 < max_jj; jj += 8) + { + __m256 _fsum = _mm256_setzero_ps(); + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + __m256i _sum = _mm256_setzero_si256(); + const int max_kk = std::min(K - k, block_size); + int kk = 0; +#if __AVX512VNNI__ + for (; kk + 3 < max_kk; kk += 4) + { + const __m128i _pA32 = _mm_cvtsi32_si128(*(const int*)pA); + const __m256i _pA = _mm256_broadcastd_epi32(_pA32); + const __m256i _pB = _mm256_loadu_si256((const __m256i*)pB); + _sum = _mm256_comp_dpbusd_epi32(_sum, _pB, _pA); + pA += 4; + pB += 32; + } + if (max_kk >= 4) + { + _sum = _mm256_sub_epi32(_sum, _mm256_set1_epi32(*(const int*)pA)); + pA += 4; + } +#else + for (; kk + 3 < max_kk; kk += 4) + { + const __m128i _pA8 = _mm_cvtsi32_si128(*(const int*)pA); + const __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + const __m256i _pA01 = _mm256_broadcastsi128_si256(_mm_shuffle_epi32(_pA16, _MM_SHUFFLE(0, 0, 0, 0))); + const __m256i _pA23 = _mm256_broadcastsi128_si256(_mm_shuffle_epi32(_pA16, _MM_SHUFFLE(1, 1, 1, 1))); + const __m256i _pB01 = _mm256_cvtepi8_epi16(_mm_loadu_si128((const __m128i*)pB)); + const __m256i _pB23 = _mm256_cvtepi8_epi16(_mm_loadu_si128((const __m128i*)(pB + 16))); + _sum = _mm256_comp_dpwssd_epi32(_sum, _pA01, _pB01); + _sum = _mm256_comp_dpwssd_epi32(_sum, _pA23, _pB23); + pA += 4; + pB += 32; + } +#endif // __AVX512VNNI__ + for (; kk + 1 < max_kk; kk += 2) + { + const __m128i _pA8 = _mm_cvtsi32_si128(*(const unsigned short*)pA); + const __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + const __m256i _pA = _mm256_broadcastsi128_si256(_mm_shuffle_epi32(_pA16, _MM_SHUFFLE(0, 0, 0, 0))); + const __m256i _pB = _mm256_cvtepi8_epi16(_mm_loadu_si128((const __m128i*)pB)); + _sum = _mm256_comp_dpwssd_epi32(_sum, _pA, _pB); + pA += 2; + pB += 16; + } + for (; kk < max_kk; kk++) + { + const __m128i _pA8 = _mm_cvtsi32_si128((signed char)pA[0]); + const __m256i _pA = _mm256_broadcastd_epi32(_pA8); + const __m256i _pB = _mm256_cvtepi8_epi32(_mm_loadl_epi64((const __m128i*)pB)); + _sum = _mm256_add_epi32(_sum, _mm256_mullo_epi32(_pA, _pB)); + pA++; + pB += 8; + } + const __m128 _ad1 = _mm_load_ss(pA_descales); + const __m256 _ad = _mm256_broadcastss_ps(_ad1); + const __m256 _descale = _mm256_mul_ps(_ad, _mm256_loadu_ps(pB_descales)); + _fsum = _mm256_add_ps(_fsum, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum), _descale)); + pA_descales += 1; + pB_descales += 8; + } + _mm256_storeu_ps(outptr, _fsum); + outptr += 8; + } +#endif // __AVX512F__ + for (; jj + 3 < max_jj; jj += 4) + { + __m128 _fsum = _mm_setzero_ps(); + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + __m128i _sum = _mm_setzero_si128(); + const int max_kk = std::min(K - k, block_size); + int kk = 0; +#if __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 3 < max_kk; kk += 4) + { + const __m128i _pA32 = _mm_cvtsi32_si128(*(const int*)pA); + const __m128i _pA = _mm_shuffle_epi32(_pA32, _MM_SHUFFLE(0, 0, 0, 0)); + const __m128i _pB = _mm_loadu_si128((const __m128i*)pB); +#if __AVXVNNIINT8__ + _sum = _mm_dpbssd_epi32(_sum, _pB, _pA); +#else // __AVXVNNIINT8__ + _sum = _mm_comp_dpbusd_epi32(_sum, _pB, _pA); +#endif // __AVXVNNIINT8__ + pA += 4; + pB += 16; + } +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + if (max_kk >= 4) + { + _sum = _mm_sub_epi32(_sum, _mm_set1_epi32(*(const int*)pA)); + pA += 4; + } +#endif +#else + for (; kk + 3 < max_kk; kk += 4) + { + const __m128i _pA8 = _mm_cvtsi32_si128(*(const int*)pA); + const __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + const __m128i _pA01 = _mm_shuffle_epi32(_pA16, _MM_SHUFFLE(0, 0, 0, 0)); + const __m128i _pA23 = _mm_shuffle_epi32(_pA16, _MM_SHUFFLE(1, 1, 1, 1)); + const __m128i _pB01x1 = _mm_loadl_epi64((const __m128i*)pB); + const __m128i _pB23x1 = _mm_loadl_epi64((const __m128i*)(pB + 8)); + const __m128i _pB01 = _mm_unpacklo_epi8(_pB01x1, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB01x1)); + const __m128i _pB23 = _mm_unpacklo_epi8(_pB23x1, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB23x1)); + _sum = _mm_comp_dpwssd_epi32(_sum, _pA01, _pB01); + _sum = _mm_comp_dpwssd_epi32(_sum, _pA23, _pB23); + pA += 4; + pB += 16; + } +#endif // __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 1 < max_kk; kk += 2) + { + const __m128i _pA8 = _mm_cvtsi32_si128(*(const unsigned short*)pA); + const __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + const __m128i _pA = _mm_shuffle_epi32(_pA16, _MM_SHUFFLE(0, 0, 0, 0)); + const __m128i _pB8 = _mm_loadl_epi64((const __m128i*)pB); + const __m128i _pB = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); + _sum = _mm_comp_dpwssd_epi32(_sum, _pA, _pB); + pA += 2; + pB += 8; + } + for (; kk < max_kk; kk++) + { + const __m128i _pA8 = _mm_cvtsi32_si128((signed char)pA[0]); + const __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + const __m128i _pA = _mm_shuffle_epi32(_mm_unpacklo_epi16(_pA16, _pA16), _MM_SHUFFLE(0, 0, 0, 0)); + const __m128i _pB8 = _mm_cvtsi32_si128(*(const int*)pB); + const __m128i _pB16 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); + const __m128i _pB = _mm_unpacklo_epi16(_pB16, _mm_setzero_si128()); + _sum = _mm_comp_dpwssd_epi32(_sum, _pA, _pB); + pA++; + pB += 4; + } + const __m128 _ad1 = _mm_load_ss(pA_descales); + const __m128 _ad = _mm_shuffle_ps(_ad1, _ad1, _MM_SHUFFLE(0, 0, 0, 0)); + const __m128 _descale = _mm_mul_ps(_ad, _mm_loadu_ps(pB_descales)); + _fsum = _mm_add_ps(_fsum, _mm_mul_ps(_mm_cvtepi32_ps(_sum), _descale)); + pA_descales += 1; + pB_descales += 4; + } + _mm_storeu_ps(outptr, _fsum); + outptr += 4; + } +#endif // defined(__x86_64__) || defined(_M_X64) +#endif // __SSE2__ + for (; jj + 1 < max_jj; jj += 2) + { +#if __SSE2__ + __m128 _fsum = _mm_setzero_ps(); + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + __m128i _sum = _mm_setzero_si128(); + const int max_kk = std::min(K - k, block_size); + int kk = 0; +#if __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 3 < max_kk; kk += 4) + { + const __m128i _pA32 = _mm_cvtsi32_si128(*(const int*)pA); + const __m128i _pA = _mm_shuffle_epi32(_pA32, _MM_SHUFFLE(0, 0, 0, 0)); + const __m128i _pB = _mm_loadl_epi64((const __m128i*)pB); +#if __AVXVNNIINT8__ + _sum = _mm_dpbssd_epi32(_sum, _pB, _pA); +#else // __AVXVNNIINT8__ + _sum = _mm_comp_dpbusd_epi32(_sum, _pB, _pA); +#endif // __AVXVNNIINT8__ + pA += 4; + pB += 8; + } +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + if (max_kk >= 4) + { + _sum = _mm_sub_epi32(_sum, _mm_set1_epi32(*(const int*)pA)); + pA += 4; + } +#endif +#else + for (; kk + 3 < max_kk; kk += 4) + { + const __m128i _pA8 = _mm_cvtsi32_si128(*(const int*)pA); + const __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + const __m128i _pA01 = _mm_shuffle_epi32(_pA16, _MM_SHUFFLE(0, 0, 0, 0)); + const __m128i _pA23 = _mm_shuffle_epi32(_pA16, _MM_SHUFFLE(1, 1, 1, 1)); + const __m128i _pB8 = _mm_loadl_epi64((const __m128i*)pB); + const __m128i _pB16 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); + const __m128i _pB23 = _mm_shuffle_epi32(_pB16, _MM_SHUFFLE(3, 2, 3, 2)); + _sum = _mm_comp_dpwssd_epi32(_sum, _pA01, _pB16); + _sum = _mm_comp_dpwssd_epi32(_sum, _pA23, _pB23); + pA += 4; + pB += 8; + } +#endif // __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 1 < max_kk; kk += 2) + { + const __m128i _pA8 = _mm_cvtsi32_si128(*(const unsigned short*)pA); + const __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + const __m128i _pA = _mm_shuffle_epi32(_pA16, _MM_SHUFFLE(0, 0, 0, 0)); + const __m128i _pB8 = _mm_cvtsi32_si128(*(const int*)pB); + const __m128i _pB = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); + _sum = _mm_comp_dpwssd_epi32(_sum, _pA, _pB); + pA += 2; + pB += 4; + } + for (; kk < max_kk; kk++) + { + const __m128i _pA8 = _mm_cvtsi32_si128((signed char)pA[0]); + const __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + const __m128i _pA = _mm_shuffle_epi32(_mm_unpacklo_epi16(_pA16, _pA16), _MM_SHUFFLE(0, 0, 0, 0)); + const __m128i _pB8 = _mm_cvtsi32_si128(*(const unsigned short*)pB); + const __m128i _pB16 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); + const __m128i _pB = _mm_unpacklo_epi16(_pB16, _mm_setzero_si128()); + _sum = _mm_comp_dpwssd_epi32(_sum, _pA, _pB); + pA++; + pB += 2; + } + const __m128 _ad1 = _mm_load_ss(pA_descales); + const __m128 _ad = _mm_shuffle_ps(_ad1, _ad1, _MM_SHUFFLE(0, 0, 0, 0)); + const __m128 _bd = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pB_descales); + _fsum = _mm_add_ps(_fsum, _mm_mul_ps(_mm_cvtepi32_ps(_sum), _mm_mul_ps(_ad, _bd))); + pA_descales += 1; + pB_descales += 2; + } + _mm_storel_pi((__m64*)outptr, _fsum); + outptr += 2; +#else + float fsum0 = 0.f; + float fsum1 = 0.f; + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + int sum0 = 0; + int sum1 = 0; + const int max_kk = std::min(K - k, block_size); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + sum0 += pA[0] * (signed char)pB[0]; + sum0 += pA[1] * (signed char)pB[2]; + sum0 += pA[2] * (signed char)pB[4]; + sum0 += pA[3] * (signed char)pB[6]; + sum1 += pA[0] * (signed char)pB[1]; + sum1 += pA[1] * (signed char)pB[3]; + sum1 += pA[2] * (signed char)pB[5]; + sum1 += pA[3] * (signed char)pB[7]; + pA += 4; + pB += 8; + } + for (; kk < max_kk; kk++) + { + sum0 += pA[0] * (signed char)pB[0]; + sum1 += pA[0] * (signed char)pB[1]; + pA++; + pB += 2; + } + + const float ad = pA_descales[0]; + fsum0 += sum0 * ad * pB_descales[0]; + fsum1 += sum1 * ad * pB_descales[1]; + pA_descales += 1; + pB_descales += 2; + } + + outptr[0] = fsum0; + outptr[1] = fsum1; + outptr += 2; +#endif // __SSE2__ + } + for (; jj < max_jj; jj++) + { +#if __SSE2__ + float fsum = 0.f; + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + __m128i _sum = _mm_setzero_si128(); + const int max_kk = std::min(K - k, block_size); + int kk = 0; +#if __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 3 < max_kk; kk += 4) + { + const __m128i _pA = _mm_cvtsi32_si128(*(const int*)pA); + const __m128i _pB = _mm_cvtsi32_si128(*(const int*)pB); +#if __AVXVNNIINT8__ + _sum = _mm_dpbssd_epi32(_sum, _pB, _pA); +#else // __AVXVNNIINT8__ + _sum = _mm_comp_dpbusd_epi32(_sum, _pB, _pA); +#endif // __AVXVNNIINT8__ + pA += 4; + pB += 4; + } +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + if (max_kk >= 4) + { + _sum = _mm_sub_epi32(_sum, _mm_cvtsi32_si128(*(const int*)pA)); + pA += 4; + } +#endif +#else + for (; kk + 3 < max_kk; kk += 4) + { + const __m128i _pA8 = _mm_cvtsi32_si128(*(const int*)pA); + const __m128i _pB8 = _mm_cvtsi32_si128(*(const int*)pB); + const __m128i _pA = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + const __m128i _pB = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); + _sum = _mm_comp_dpwssd_epi32(_sum, _pA, _pB); + pA += 4; + pB += 4; + } +#endif // __AVX512VNNI__ || __AVXVNNI__ + for (; kk + 1 < max_kk; kk += 2) + { + const __m128i _pA8 = _mm_cvtsi32_si128(*(const unsigned short*)pA); + const __m128i _pB8 = _mm_cvtsi32_si128(*(const unsigned short*)pB); + const __m128i _pA = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + const __m128i _pB = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); + _sum = _mm_comp_dpwssd_epi32(_sum, _pA, _pB); + pA += 2; + pB += 2; + } + for (; kk < max_kk; kk++) + { + const __m128i _pA8 = _mm_cvtsi32_si128((unsigned char)pA[0]); + const __m128i _pB8 = _mm_cvtsi32_si128((unsigned char)pB[0]); + const __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + const __m128i _pB16 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); + _sum = _mm_comp_dpwssd_epi32(_sum, _mm_unpacklo_epi16(_pA16, _pA16), _mm_unpacklo_epi16(_pB16, _mm_setzero_si128())); + pA++; + pB++; + } + fsum += _mm_reduce_add_epi32(_sum) * pA_descales[0] * pB_descales[0]; + pA_descales += 1; + pB_descales++; + } + outptr[0] = fsum; + outptr++; +#else + float fsum = 0.f; + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + int sum = 0; + const int max_kk = std::min(K - k, block_size); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + sum += pA[0] * (signed char)pB[0]; + sum += pA[1] * (signed char)pB[1]; + sum += pA[2] * (signed char)pB[2]; + sum += pA[3] * (signed char)pB[3]; + pA += 4; + pB += 4; + } + for (; kk < max_kk; kk++) + { + sum += pA[0] * (signed char)pB[0]; + pA++; + pB++; + } + + fsum += sum * pA_descales[0] * pB_descales[0]; + pA_descales += 1; + pB_descales++; + } + + outptr[0] = fsum; + outptr++; +#endif // __SSE2__ + } + + pAT += A_hstep; + pAT_descales += A_descales_hstep; + } +} + +static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta) +{ +#if NCNN_RUNTIME_CPU && NCNN_AVX512VNNI && __AVX512F__ && !__AVX512VNNI__ + if (ncnn::cpu_support_x86_avx512_vnni()) + { + unpack_output_tile_wq_int8_avx512vnni(topT, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, N, alpha, beta); + return; + } +#endif +#if NCNN_RUNTIME_CPU && NCNN_AVXVNNIINT8 && __AVX__ && !__AVXVNNIINT8__ && !__AVX512VNNI__ + if (ncnn::cpu_support_x86_avx_vnni_int8()) + { + unpack_output_tile_wq_int8_avxvnniint8(topT, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, N, alpha, beta); + return; + } +#endif +#if NCNN_RUNTIME_CPU && NCNN_AVXVNNI && __AVX__ && !__AVXVNNI__ && !__AVXVNNIINT8__ && !__AVX512VNNI__ + if (ncnn::cpu_support_x86_avx_vnni()) + { + unpack_output_tile_wq_int8_avxvnni(topT, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, N, alpha, beta); + return; + } +#endif +#if NCNN_RUNTIME_CPU && NCNN_AVX2 && __AVX__ && !__AVX2__ && !__AVXVNNI__ && !__AVXVNNIINT8__ && !__AVX512VNNI__ + if (ncnn::cpu_support_x86_avx2()) + { + unpack_output_tile_wq_int8_avx2(topT, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, N, alpha, beta); + return; + } +#endif + + const float* pC = C; + const float* pp = topT; + const size_t out_hstep = top_blob.dims == 3 ? top_blob.cstep : (size_t)top_blob.w; + + int ii = 0; +#if __SSE2__ +#if __AVX__ +#if __AVX512F__ + for (; ii + 15 < max_ii; ii += 16) + { + float* p0 = (float*)top_blob + (i + ii) * out_hstep + j; + + __m512 _c = _mm512_setzero_ps(); + __m512i _c_vindex = _mm512_setzero_si512(); + if (pC) + { + if (broadcast_type_C == 0) + _c = _mm512_set1_ps(pC[0]); + if (broadcast_type_C == 1 || broadcast_type_C == 2) + _c = _mm512_loadu_ps(pC + i + ii); + if (broadcast_type_C == 3) + { + pC = (const float*)C + (i + ii) * N + j; + _c_vindex = _mm512_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15); + _c_vindex = _mm512_mullo_epi32(_c_vindex, _mm512_set1_epi32(N)); + } + if (broadcast_type_C == 4) + pC = (const float*)C + j; + if ((broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) && beta != 1.f) + _c = _mm512_mul_ps(_c, _mm512_set1_ps(beta)); + } + + int jj = 0; +#if defined(__x86_64__) || defined(_M_X64) + for (; jj + 7 < max_jj; jj += 8) + { + __m512 _f0 = _mm512_loadu_ps(pp + 0); + __m512 _f1 = _mm512_loadu_ps(pp + 16); + __m512 _f2 = _mm512_loadu_ps(pp + 32); + __m512 _f3 = _mm512_loadu_ps(pp + 48); + __m512 _f4 = _mm512_loadu_ps(pp + 64); + __m512 _f5 = _mm512_loadu_ps(pp + 80); + __m512 _f6 = _mm512_loadu_ps(pp + 96); + __m512 _f7 = _mm512_loadu_ps(pp + 112); + pp += 128; + _f1 = _mm512_permute_ps(_f1, _MM_SHUFFLE(2, 1, 0, 3)); + _f3 = _mm512_permute_ps(_f3, _MM_SHUFFLE(2, 1, 0, 3)); + _f5 = _mm512_permute_ps(_f5, _MM_SHUFFLE(2, 1, 0, 3)); + _f7 = _mm512_permute_ps(_f7, _MM_SHUFFLE(2, 1, 0, 3)); + __m512 _tmp0 = _mm512_unpacklo_ps(_f0, _f3); + __m512 _tmp1 = _mm512_unpackhi_ps(_f0, _f3); + __m512 _tmp2 = _mm512_unpacklo_ps(_f2, _f1); + __m512 _tmp3 = _mm512_unpackhi_ps(_f2, _f1); + __m512 _tmp4 = _mm512_unpacklo_ps(_f4, _f7); + __m512 _tmp5 = _mm512_unpackhi_ps(_f4, _f7); + __m512 _tmp6 = _mm512_unpacklo_ps(_f6, _f5); + __m512 _tmp7 = _mm512_unpackhi_ps(_f6, _f5); + _f0 = _mm512_castpd_ps(_mm512_unpacklo_pd(_mm512_castps_pd(_tmp0), _mm512_castps_pd(_tmp2))); + _f1 = _mm512_castpd_ps(_mm512_unpackhi_pd(_mm512_castps_pd(_tmp0), _mm512_castps_pd(_tmp2))); + _f2 = _mm512_castpd_ps(_mm512_unpacklo_pd(_mm512_castps_pd(_tmp3), _mm512_castps_pd(_tmp1))); + _f3 = _mm512_castpd_ps(_mm512_unpackhi_pd(_mm512_castps_pd(_tmp3), _mm512_castps_pd(_tmp1))); + _f4 = _mm512_castpd_ps(_mm512_unpacklo_pd(_mm512_castps_pd(_tmp4), _mm512_castps_pd(_tmp6))); + _f5 = _mm512_castpd_ps(_mm512_unpackhi_pd(_mm512_castps_pd(_tmp4), _mm512_castps_pd(_tmp6))); + _f6 = _mm512_castpd_ps(_mm512_unpacklo_pd(_mm512_castps_pd(_tmp7), _mm512_castps_pd(_tmp5))); + _f7 = _mm512_castpd_ps(_mm512_unpackhi_pd(_mm512_castps_pd(_tmp7), _mm512_castps_pd(_tmp5))); + _f1 = _mm512_permute_ps(_f1, _MM_SHUFFLE(2, 1, 0, 3)); + _f3 = _mm512_permute_ps(_f3, _MM_SHUFFLE(2, 1, 0, 3)); + _f5 = _mm512_permute_ps(_f5, _MM_SHUFFLE(2, 1, 0, 3)); + _f7 = _mm512_permute_ps(_f7, _MM_SHUFFLE(2, 1, 0, 3)); + _tmp0 = _mm512_shuffle_f32x4(_f0, _f4, _MM_SHUFFLE(0, 1, 1, 0)); + _tmp1 = _mm512_shuffle_f32x4(_f1, _f5, _MM_SHUFFLE(0, 1, 1, 0)); + _tmp2 = _mm512_shuffle_f32x4(_f2, _f6, _MM_SHUFFLE(0, 1, 1, 0)); + _tmp3 = _mm512_shuffle_f32x4(_f3, _f7, _MM_SHUFFLE(0, 1, 1, 0)); + _tmp4 = _mm512_shuffle_f32x4(_f0, _f4, _MM_SHUFFLE(2, 3, 3, 2)); + _tmp5 = _mm512_shuffle_f32x4(_f1, _f5, _MM_SHUFFLE(2, 3, 3, 2)); + _tmp6 = _mm512_shuffle_f32x4(_f2, _f6, _MM_SHUFFLE(2, 3, 3, 2)); + _tmp7 = _mm512_shuffle_f32x4(_f3, _f7, _MM_SHUFFLE(2, 3, 3, 2)); + _f0 = _mm512_shuffle_f32x4(_tmp0, _tmp4, _MM_SHUFFLE(2, 0, 2, 0)); + _f1 = _mm512_shuffle_f32x4(_tmp1, _tmp5, _MM_SHUFFLE(2, 0, 2, 0)); + _f2 = _mm512_shuffle_f32x4(_tmp2, _tmp6, _MM_SHUFFLE(2, 0, 2, 0)); + _f3 = _mm512_shuffle_f32x4(_tmp3, _tmp7, _MM_SHUFFLE(2, 0, 2, 0)); + _f4 = _mm512_shuffle_f32x4(_tmp0, _tmp4, _MM_SHUFFLE(1, 3, 1, 3)); + _f5 = _mm512_shuffle_f32x4(_tmp1, _tmp5, _MM_SHUFFLE(1, 3, 1, 3)); + _f6 = _mm512_shuffle_f32x4(_tmp2, _tmp6, _MM_SHUFFLE(1, 3, 1, 3)); + _f7 = _mm512_shuffle_f32x4(_tmp3, _tmp7, _MM_SHUFFLE(1, 3, 1, 3)); + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = _mm512_add_ps(_f0, _c); + _f1 = _mm512_add_ps(_f1, _c); + _f2 = _mm512_add_ps(_f2, _c); + _f3 = _mm512_add_ps(_f3, _c); + _f4 = _mm512_add_ps(_f4, _c); + _f5 = _mm512_add_ps(_f5, _c); + _f6 = _mm512_add_ps(_f6, _c); + _f7 = _mm512_add_ps(_f7, _c); + } + if (broadcast_type_C == 3) + { + __m512 _c0 = _mm512_i32gather_ps(_c_vindex, pC, sizeof(float)); + if (beta != 1.f) _c0 = _mm512_mul_ps(_c0, _mm512_set1_ps(beta)); + _f0 = _mm512_add_ps(_f0, _c0); + __m512 _c1 = _mm512_i32gather_ps(_c_vindex, pC + 1, sizeof(float)); + if (beta != 1.f) _c1 = _mm512_mul_ps(_c1, _mm512_set1_ps(beta)); + _f1 = _mm512_add_ps(_f1, _c1); + __m512 _c2 = _mm512_i32gather_ps(_c_vindex, pC + 2, sizeof(float)); + if (beta != 1.f) _c2 = _mm512_mul_ps(_c2, _mm512_set1_ps(beta)); + _f2 = _mm512_add_ps(_f2, _c2); + __m512 _c3 = _mm512_i32gather_ps(_c_vindex, pC + 3, sizeof(float)); + if (beta != 1.f) _c3 = _mm512_mul_ps(_c3, _mm512_set1_ps(beta)); + _f3 = _mm512_add_ps(_f3, _c3); + __m512 _c4 = _mm512_i32gather_ps(_c_vindex, pC + 4, sizeof(float)); + if (beta != 1.f) _c4 = _mm512_mul_ps(_c4, _mm512_set1_ps(beta)); + _f4 = _mm512_add_ps(_f4, _c4); + __m512 _c5 = _mm512_i32gather_ps(_c_vindex, pC + 5, sizeof(float)); + if (beta != 1.f) _c5 = _mm512_mul_ps(_c5, _mm512_set1_ps(beta)); + _f5 = _mm512_add_ps(_f5, _c5); + __m512 _c6 = _mm512_i32gather_ps(_c_vindex, pC + 6, sizeof(float)); + if (beta != 1.f) _c6 = _mm512_mul_ps(_c6, _mm512_set1_ps(beta)); + _f6 = _mm512_add_ps(_f6, _c6); + __m512 _c7 = _mm512_i32gather_ps(_c_vindex, pC + 7, sizeof(float)); + if (beta != 1.f) _c7 = _mm512_mul_ps(_c7, _mm512_set1_ps(beta)); + _f7 = _mm512_add_ps(_f7, _c7); + pC += 8; + } + if (broadcast_type_C == 4) + { + __m512 _c0 = _mm512_set1_ps(pC[0]); + if (beta != 1.f) _c0 = _mm512_mul_ps(_c0, _mm512_set1_ps(beta)); + _f0 = _mm512_add_ps(_f0, _c0); + __m512 _c1 = _mm512_set1_ps(pC[1]); + if (beta != 1.f) _c1 = _mm512_mul_ps(_c1, _mm512_set1_ps(beta)); + _f1 = _mm512_add_ps(_f1, _c1); + __m512 _c2 = _mm512_set1_ps(pC[2]); + if (beta != 1.f) _c2 = _mm512_mul_ps(_c2, _mm512_set1_ps(beta)); + _f2 = _mm512_add_ps(_f2, _c2); + __m512 _c3 = _mm512_set1_ps(pC[3]); + if (beta != 1.f) _c3 = _mm512_mul_ps(_c3, _mm512_set1_ps(beta)); + _f3 = _mm512_add_ps(_f3, _c3); + __m512 _c4 = _mm512_set1_ps(pC[4]); + if (beta != 1.f) _c4 = _mm512_mul_ps(_c4, _mm512_set1_ps(beta)); + _f4 = _mm512_add_ps(_f4, _c4); + __m512 _c5 = _mm512_set1_ps(pC[5]); + if (beta != 1.f) _c5 = _mm512_mul_ps(_c5, _mm512_set1_ps(beta)); + _f5 = _mm512_add_ps(_f5, _c5); + __m512 _c6 = _mm512_set1_ps(pC[6]); + if (beta != 1.f) _c6 = _mm512_mul_ps(_c6, _mm512_set1_ps(beta)); + _f6 = _mm512_add_ps(_f6, _c6); + __m512 _c7 = _mm512_set1_ps(pC[7]); + if (beta != 1.f) _c7 = _mm512_mul_ps(_c7, _mm512_set1_ps(beta)); + _f7 = _mm512_add_ps(_f7, _c7); + pC += 8; + } + } + if (alpha != 1.f) + { + const __m512 _alpha = _mm512_set1_ps(alpha); + _f0 = _mm512_mul_ps(_f0, _alpha); + _f1 = _mm512_mul_ps(_f1, _alpha); + _f2 = _mm512_mul_ps(_f2, _alpha); + _f3 = _mm512_mul_ps(_f3, _alpha); + _f4 = _mm512_mul_ps(_f4, _alpha); + _f5 = _mm512_mul_ps(_f5, _alpha); + _f6 = _mm512_mul_ps(_f6, _alpha); + _f7 = _mm512_mul_ps(_f7, _alpha); + } + transpose16x8_ps(_f0, _f1, _f2, _f3, _f4, _f5, _f6, _f7); + _mm256_storeu_ps(p0, _mm512_castps512_ps256(_f0)); + _mm256_storeu_ps(p0 + out_hstep, _mm512_extractf32x8_ps(_f0, 1)); + _mm256_storeu_ps(p0 + out_hstep * 2, _mm512_castps512_ps256(_f1)); + _mm256_storeu_ps(p0 + out_hstep * 3, _mm512_extractf32x8_ps(_f1, 1)); + _mm256_storeu_ps(p0 + out_hstep * 4, _mm512_castps512_ps256(_f2)); + _mm256_storeu_ps(p0 + out_hstep * 5, _mm512_extractf32x8_ps(_f2, 1)); + _mm256_storeu_ps(p0 + out_hstep * 6, _mm512_castps512_ps256(_f3)); + _mm256_storeu_ps(p0 + out_hstep * 7, _mm512_extractf32x8_ps(_f3, 1)); + _mm256_storeu_ps(p0 + out_hstep * 8, _mm512_castps512_ps256(_f4)); + _mm256_storeu_ps(p0 + out_hstep * 9, _mm512_extractf32x8_ps(_f4, 1)); + _mm256_storeu_ps(p0 + out_hstep * 10, _mm512_castps512_ps256(_f5)); + _mm256_storeu_ps(p0 + out_hstep * 11, _mm512_extractf32x8_ps(_f5, 1)); + _mm256_storeu_ps(p0 + out_hstep * 12, _mm512_castps512_ps256(_f6)); + _mm256_storeu_ps(p0 + out_hstep * 13, _mm512_extractf32x8_ps(_f6, 1)); + _mm256_storeu_ps(p0 + out_hstep * 14, _mm512_castps512_ps256(_f7)); + _mm256_storeu_ps(p0 + out_hstep * 15, _mm512_extractf32x8_ps(_f7, 1)); + p0 += 8; + + } + for (; jj + 3 < max_jj; jj += 4) + { + __m512 _f0 = _mm512_loadu_ps(pp + 0); + __m512 _f1 = _mm512_loadu_ps(pp + 16); + __m512 _f2 = _mm512_loadu_ps(pp + 32); + __m512 _f3 = _mm512_loadu_ps(pp + 48); + pp += 64; + _f1 = _mm512_permute_ps(_f1, _MM_SHUFFLE(2, 1, 0, 3)); + _f3 = _mm512_permute_ps(_f3, _MM_SHUFFLE(2, 1, 0, 3)); + __m512 _tmp0 = _mm512_unpacklo_ps(_f0, _f3); + __m512 _tmp1 = _mm512_unpackhi_ps(_f0, _f3); + __m512 _tmp2 = _mm512_unpacklo_ps(_f2, _f1); + __m512 _tmp3 = _mm512_unpackhi_ps(_f2, _f1); + _f0 = _mm512_castpd_ps(_mm512_unpacklo_pd(_mm512_castps_pd(_tmp0), _mm512_castps_pd(_tmp2))); + _f1 = _mm512_castpd_ps(_mm512_unpackhi_pd(_mm512_castps_pd(_tmp0), _mm512_castps_pd(_tmp2))); + _f2 = _mm512_castpd_ps(_mm512_unpacklo_pd(_mm512_castps_pd(_tmp3), _mm512_castps_pd(_tmp1))); + _f3 = _mm512_castpd_ps(_mm512_unpackhi_pd(_mm512_castps_pd(_tmp3), _mm512_castps_pd(_tmp1))); + _f1 = _mm512_permute_ps(_f1, _MM_SHUFFLE(2, 1, 0, 3)); + _f3 = _mm512_permute_ps(_f3, _MM_SHUFFLE(2, 1, 0, 3)); + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = _mm512_add_ps(_f0, _c); + _f1 = _mm512_add_ps(_f1, _c); + _f2 = _mm512_add_ps(_f2, _c); + _f3 = _mm512_add_ps(_f3, _c); + } + if (broadcast_type_C == 3) + { + __m512 _c0 = _mm512_i32gather_ps(_c_vindex, pC, sizeof(float)); + if (beta != 1.f) _c0 = _mm512_mul_ps(_c0, _mm512_set1_ps(beta)); + _f0 = _mm512_add_ps(_f0, _c0); + __m512 _c1 = _mm512_i32gather_ps(_c_vindex, pC + 1, sizeof(float)); + if (beta != 1.f) _c1 = _mm512_mul_ps(_c1, _mm512_set1_ps(beta)); + _f1 = _mm512_add_ps(_f1, _c1); + __m512 _c2 = _mm512_i32gather_ps(_c_vindex, pC + 2, sizeof(float)); + if (beta != 1.f) _c2 = _mm512_mul_ps(_c2, _mm512_set1_ps(beta)); + _f2 = _mm512_add_ps(_f2, _c2); + __m512 _c3 = _mm512_i32gather_ps(_c_vindex, pC + 3, sizeof(float)); + if (beta != 1.f) _c3 = _mm512_mul_ps(_c3, _mm512_set1_ps(beta)); + _f3 = _mm512_add_ps(_f3, _c3); + pC += 4; + } + if (broadcast_type_C == 4) + { + __m512 _c0 = _mm512_set1_ps(pC[0]); + if (beta != 1.f) _c0 = _mm512_mul_ps(_c0, _mm512_set1_ps(beta)); + _f0 = _mm512_add_ps(_f0, _c0); + __m512 _c1 = _mm512_set1_ps(pC[1]); + if (beta != 1.f) _c1 = _mm512_mul_ps(_c1, _mm512_set1_ps(beta)); + _f1 = _mm512_add_ps(_f1, _c1); + __m512 _c2 = _mm512_set1_ps(pC[2]); + if (beta != 1.f) _c2 = _mm512_mul_ps(_c2, _mm512_set1_ps(beta)); + _f2 = _mm512_add_ps(_f2, _c2); + __m512 _c3 = _mm512_set1_ps(pC[3]); + if (beta != 1.f) _c3 = _mm512_mul_ps(_c3, _mm512_set1_ps(beta)); + _f3 = _mm512_add_ps(_f3, _c3); + pC += 4; + } + } + if (alpha != 1.f) + { + const __m512 _alpha = _mm512_set1_ps(alpha); + _f0 = _mm512_mul_ps(_f0, _alpha); + _f1 = _mm512_mul_ps(_f1, _alpha); + _f2 = _mm512_mul_ps(_f2, _alpha); + _f3 = _mm512_mul_ps(_f3, _alpha); + } + transpose16x4_ps(_f0, _f1, _f2, _f3); + _mm_storeu_ps(p0, _mm512_extractf32x4_ps(_f0, 0)); + _mm_storeu_ps(p0 + out_hstep, _mm512_extractf32x4_ps(_f0, 1)); + _mm_storeu_ps(p0 + out_hstep * 2, _mm512_extractf32x4_ps(_f0, 2)); + _mm_storeu_ps(p0 + out_hstep * 3, _mm512_extractf32x4_ps(_f0, 3)); + _mm_storeu_ps(p0 + out_hstep * 4, _mm512_extractf32x4_ps(_f1, 0)); + _mm_storeu_ps(p0 + out_hstep * 5, _mm512_extractf32x4_ps(_f1, 1)); + _mm_storeu_ps(p0 + out_hstep * 6, _mm512_extractf32x4_ps(_f1, 2)); + _mm_storeu_ps(p0 + out_hstep * 7, _mm512_extractf32x4_ps(_f1, 3)); + _mm_storeu_ps(p0 + out_hstep * 8, _mm512_extractf32x4_ps(_f2, 0)); + _mm_storeu_ps(p0 + out_hstep * 9, _mm512_extractf32x4_ps(_f2, 1)); + _mm_storeu_ps(p0 + out_hstep * 10, _mm512_extractf32x4_ps(_f2, 2)); + _mm_storeu_ps(p0 + out_hstep * 11, _mm512_extractf32x4_ps(_f2, 3)); + _mm_storeu_ps(p0 + out_hstep * 12, _mm512_extractf32x4_ps(_f3, 0)); + _mm_storeu_ps(p0 + out_hstep * 13, _mm512_extractf32x4_ps(_f3, 1)); + _mm_storeu_ps(p0 + out_hstep * 14, _mm512_extractf32x4_ps(_f3, 2)); + _mm_storeu_ps(p0 + out_hstep * 15, _mm512_extractf32x4_ps(_f3, 3)); + p0 += 4; + + } +#endif // defined(__x86_64__) || defined(_M_X64) + for (; jj + 1 < max_jj; jj += 2) + { + __m512 _f0 = _mm512_loadu_ps(pp + 0); + __m512 _f1 = _mm512_loadu_ps(pp + 16); + pp += 32; + __m512 _tmp0 = _mm512_permute_ps(_f0, _MM_SHUFFLE(3, 1, 2, 0)); + __m512 _tmp1 = _mm512_permute_ps(_f1, _MM_SHUFFLE(0, 2, 3, 1)); + _f0 = _mm512_unpacklo_ps(_tmp0, _tmp1); + _f1 = _mm512_unpackhi_ps(_tmp0, _tmp1); + _f1 = _mm512_permute_ps(_f1, _MM_SHUFFLE(2, 1, 0, 3)); + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = _mm512_add_ps(_f0, _c); + _f1 = _mm512_add_ps(_f1, _c); + } + if (broadcast_type_C == 3) + { + __m512 _c0 = _mm512_i32gather_ps(_c_vindex, pC, sizeof(float)); + if (beta != 1.f) _c0 = _mm512_mul_ps(_c0, _mm512_set1_ps(beta)); + _f0 = _mm512_add_ps(_f0, _c0); + __m512 _c1 = _mm512_i32gather_ps(_c_vindex, pC + 1, sizeof(float)); + if (beta != 1.f) _c1 = _mm512_mul_ps(_c1, _mm512_set1_ps(beta)); + _f1 = _mm512_add_ps(_f1, _c1); + pC += 2; + } + if (broadcast_type_C == 4) + { + __m512 _c0 = _mm512_set1_ps(pC[0]); + if (beta != 1.f) _c0 = _mm512_mul_ps(_c0, _mm512_set1_ps(beta)); + _f0 = _mm512_add_ps(_f0, _c0); + __m512 _c1 = _mm512_set1_ps(pC[1]); + if (beta != 1.f) _c1 = _mm512_mul_ps(_c1, _mm512_set1_ps(beta)); + _f1 = _mm512_add_ps(_f1, _c1); + pC += 2; + } + } + if (alpha != 1.f) + { + const __m512 _alpha = _mm512_set1_ps(alpha); + _f0 = _mm512_mul_ps(_f0, _alpha); + _f1 = _mm512_mul_ps(_f1, _alpha); + } + transpose16x2_ps(_f0, _f1); + { + const __m128 _r = _mm512_extractf32x4_ps(_f0, 0); + _mm_storel_pi((__m64*)(p0), _r); + _mm_storeh_pi((__m64*)(p0 + out_hstep), _r); + } + { + const __m128 _r = _mm512_extractf32x4_ps(_f0, 1); + _mm_storel_pi((__m64*)(p0 + out_hstep * 2), _r); + _mm_storeh_pi((__m64*)(p0 + out_hstep * 3), _r); + } + { + const __m128 _r = _mm512_extractf32x4_ps(_f0, 2); + _mm_storel_pi((__m64*)(p0 + out_hstep * 4), _r); + _mm_storeh_pi((__m64*)(p0 + out_hstep * 5), _r); + } + { + const __m128 _r = _mm512_extractf32x4_ps(_f0, 3); + _mm_storel_pi((__m64*)(p0 + out_hstep * 6), _r); + _mm_storeh_pi((__m64*)(p0 + out_hstep * 7), _r); + } + { + const __m128 _r = _mm512_extractf32x4_ps(_f1, 0); + _mm_storel_pi((__m64*)(p0 + out_hstep * 8), _r); + _mm_storeh_pi((__m64*)(p0 + out_hstep * 9), _r); + } + { + const __m128 _r = _mm512_extractf32x4_ps(_f1, 1); + _mm_storel_pi((__m64*)(p0 + out_hstep * 10), _r); + _mm_storeh_pi((__m64*)(p0 + out_hstep * 11), _r); + } + { + const __m128 _r = _mm512_extractf32x4_ps(_f1, 2); + _mm_storel_pi((__m64*)(p0 + out_hstep * 12), _r); + _mm_storeh_pi((__m64*)(p0 + out_hstep * 13), _r); + } + { + const __m128 _r = _mm512_extractf32x4_ps(_f1, 3); + _mm_storel_pi((__m64*)(p0 + out_hstep * 14), _r); + _mm_storeh_pi((__m64*)(p0 + out_hstep * 15), _r); + } + p0 += 2; + + } + for (; jj < max_jj; jj++) + { + __m512 _f0 = _mm512_loadu_ps(pp + 0); + pp += 16; + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = _mm512_add_ps(_f0, _c); + } + if (broadcast_type_C == 3) + { + __m512 _c0 = _mm512_i32gather_ps(_c_vindex, pC, sizeof(float)); + if (beta != 1.f) _c0 = _mm512_mul_ps(_c0, _mm512_set1_ps(beta)); + _f0 = _mm512_add_ps(_f0, _c0); + pC++; + } + if (broadcast_type_C == 4) + { + __m512 _c0 = _mm512_set1_ps(pC[0]); + if (beta != 1.f) _c0 = _mm512_mul_ps(_c0, _mm512_set1_ps(beta)); + _f0 = _mm512_add_ps(_f0, _c0); + pC++; + } + } + if (alpha != 1.f) + { + const __m512 _alpha = _mm512_set1_ps(alpha); + _f0 = _mm512_mul_ps(_f0, _alpha); + } + __m512i _vindex = _mm512_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15); + _vindex = _mm512_mullo_epi32(_vindex, _mm512_set1_epi32((int)out_hstep)); + _mm512_i32scatter_ps(p0, _vindex, _f0, sizeof(float)); + p0++; + + } + } +#endif // __AVX512F__ +#if !__AVX2__ + const float* pp1 = pp + max_jj * 4; +#endif + for (; ii + 7 < max_ii; ii += 8) + { + float* p0 = (float*)top_blob + (i + ii) * out_hstep + j; + + float c0 = 0.f; + float c1 = c0; + float c2 = c0; + float c3 = c0; + float c4 = c0; + float c5 = c0; + float c6 = c0; + float c7 = c0; + if (pC) + { + if (broadcast_type_C == 0) + { + c0 = pC[0]; + c1 = c0; + c2 = c0; + c3 = c0; + c4 = c0; + c5 = c0; + c6 = c0; + c7 = c0; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + c0 = pC[i + ii]; + c1 = pC[i + ii + 1]; + c2 = pC[i + ii + 2]; + c3 = pC[i + ii + 3]; + c4 = pC[i + ii + 4]; + c5 = pC[i + ii + 5]; + c6 = pC[i + ii + 6]; + c7 = pC[i + ii + 7]; + } + if (broadcast_type_C == 3) + pC = (const float*)C + (i + ii) * N + j; + if (broadcast_type_C == 4) + pC = (const float*)C + j; + if (broadcast_type_C == 0 && beta != 1.f) + c0 *= beta; + if ((broadcast_type_C == 1 || broadcast_type_C == 2) && beta != 1.f) + { + c0 *= beta; + c1 *= beta; + c2 *= beta; + c3 *= beta; + c4 *= beta; + c5 *= beta; + c6 *= beta; + c7 *= beta; + } + } + + int jj = 0; +#if defined(__x86_64__) || defined(_M_X64) +#if __AVX512F__ + for (; jj + 7 < max_jj; jj += 8) + { + __m256 _f0 = _mm256_loadu_ps(pp + 0); + __m256 _f1 = _mm256_loadu_ps(pp + 8); + __m256 _f2 = _mm256_loadu_ps(pp + 16); + __m256 _f3 = _mm256_loadu_ps(pp + 24); + __m256 _f4 = _mm256_loadu_ps(pp + 32); + __m256 _f5 = _mm256_loadu_ps(pp + 40); + __m256 _f6 = _mm256_loadu_ps(pp + 48); + __m256 _f7 = _mm256_loadu_ps(pp + 56); + { + __m256 _tmp0 = _f0; + __m256 _tmp1 = _mm256_shuffle_ps(_f1, _f1, _MM_SHUFFLE(2, 1, 0, 3)); + __m256 _tmp2 = _f2; + __m256 _tmp3 = _mm256_shuffle_ps(_f3, _f3, _MM_SHUFFLE(2, 1, 0, 3)); + __m256 _tmp4 = _f4; + __m256 _tmp5 = _mm256_shuffle_ps(_f5, _f5, _MM_SHUFFLE(2, 1, 0, 3)); + __m256 _tmp6 = _f6; + __m256 _tmp7 = _mm256_shuffle_ps(_f7, _f7, _MM_SHUFFLE(2, 1, 0, 3)); + _f0 = _mm256_unpacklo_ps(_tmp0, _tmp3); + _f1 = _mm256_unpackhi_ps(_tmp0, _tmp3); + _f2 = _mm256_unpacklo_ps(_tmp2, _tmp1); + _f3 = _mm256_unpackhi_ps(_tmp2, _tmp1); + _f4 = _mm256_unpacklo_ps(_tmp4, _tmp7); + _f5 = _mm256_unpackhi_ps(_tmp4, _tmp7); + _f6 = _mm256_unpacklo_ps(_tmp6, _tmp5); + _f7 = _mm256_unpackhi_ps(_tmp6, _tmp5); + _tmp0 = _mm256_castpd_ps(_mm256_unpacklo_pd(_mm256_castps_pd(_f0), _mm256_castps_pd(_f2))); + _tmp1 = _mm256_castpd_ps(_mm256_unpackhi_pd(_mm256_castps_pd(_f0), _mm256_castps_pd(_f2))); + _tmp2 = _mm256_castpd_ps(_mm256_unpacklo_pd(_mm256_castps_pd(_f3), _mm256_castps_pd(_f1))); + _tmp3 = _mm256_castpd_ps(_mm256_unpackhi_pd(_mm256_castps_pd(_f3), _mm256_castps_pd(_f1))); + _tmp4 = _mm256_castpd_ps(_mm256_unpacklo_pd(_mm256_castps_pd(_f4), _mm256_castps_pd(_f6))); + _tmp5 = _mm256_castpd_ps(_mm256_unpackhi_pd(_mm256_castps_pd(_f4), _mm256_castps_pd(_f6))); + _tmp6 = _mm256_castpd_ps(_mm256_unpacklo_pd(_mm256_castps_pd(_f7), _mm256_castps_pd(_f5))); + _tmp7 = _mm256_castpd_ps(_mm256_unpackhi_pd(_mm256_castps_pd(_f7), _mm256_castps_pd(_f5))); + _tmp1 = _mm256_shuffle_ps(_tmp1, _tmp1, _MM_SHUFFLE(2, 1, 0, 3)); + _tmp3 = _mm256_shuffle_ps(_tmp3, _tmp3, _MM_SHUFFLE(2, 1, 0, 3)); + _tmp5 = _mm256_shuffle_ps(_tmp5, _tmp5, _MM_SHUFFLE(2, 1, 0, 3)); + _tmp7 = _mm256_shuffle_ps(_tmp7, _tmp7, _MM_SHUFFLE(2, 1, 0, 3)); + _f0 = _mm256_permute2f128_ps(_tmp0, _tmp4, _MM_SHUFFLE(0, 3, 0, 0)); + _f1 = _mm256_permute2f128_ps(_tmp1, _tmp5, _MM_SHUFFLE(0, 3, 0, 0)); + _f2 = _mm256_permute2f128_ps(_tmp2, _tmp6, _MM_SHUFFLE(0, 3, 0, 0)); + _f3 = _mm256_permute2f128_ps(_tmp3, _tmp7, _MM_SHUFFLE(0, 3, 0, 0)); + _f4 = _mm256_permute2f128_ps(_tmp4, _tmp0, _MM_SHUFFLE(0, 3, 0, 0)); + _f5 = _mm256_permute2f128_ps(_tmp5, _tmp1, _MM_SHUFFLE(0, 3, 0, 0)); + _f6 = _mm256_permute2f128_ps(_tmp6, _tmp2, _MM_SHUFFLE(0, 3, 0, 0)); + _f7 = _mm256_permute2f128_ps(_tmp7, _tmp3, _MM_SHUFFLE(0, 3, 0, 0)); + } + transpose8x8_ps(_f0, _f1, _f2, _f3, _f4, _f5, _f6, _f7); + if (pC) + { + if (broadcast_type_C == 0) + { + __m256 _c = _mm256_set1_ps(c0); + _f0 = _mm256_add_ps(_f0, _c); + _f1 = _mm256_add_ps(_f1, _c); + _f2 = _mm256_add_ps(_f2, _c); + _f3 = _mm256_add_ps(_f3, _c); + _f4 = _mm256_add_ps(_f4, _c); + _f5 = _mm256_add_ps(_f5, _c); + _f6 = _mm256_add_ps(_f6, _c); + _f7 = _mm256_add_ps(_f7, _c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = _mm256_add_ps(_f0, _mm256_set1_ps(c0)); + _f1 = _mm256_add_ps(_f1, _mm256_set1_ps(c1)); + _f2 = _mm256_add_ps(_f2, _mm256_set1_ps(c2)); + _f3 = _mm256_add_ps(_f3, _mm256_set1_ps(c3)); + _f4 = _mm256_add_ps(_f4, _mm256_set1_ps(c4)); + _f5 = _mm256_add_ps(_f5, _mm256_set1_ps(c5)); + _f6 = _mm256_add_ps(_f6, _mm256_set1_ps(c6)); + _f7 = _mm256_add_ps(_f7, _mm256_set1_ps(c7)); + } + if (broadcast_type_C == 3) + { + __m256 _c0 = _mm256_loadu_ps(pC); + __m256 _c1 = _mm256_loadu_ps(pC + N); + __m256 _c2 = _mm256_loadu_ps(pC + N * 2); + __m256 _c3 = _mm256_loadu_ps(pC + N * 3); + __m256 _c4 = _mm256_loadu_ps(pC + N * 4); + __m256 _c5 = _mm256_loadu_ps(pC + N * 5); + __m256 _c6 = _mm256_loadu_ps(pC + N * 6); + __m256 _c7 = _mm256_loadu_ps(pC + N * 7); + if (beta == 1.f) + { + _f0 = _mm256_add_ps(_f0, _c0); + _f1 = _mm256_add_ps(_f1, _c1); + _f2 = _mm256_add_ps(_f2, _c2); + _f3 = _mm256_add_ps(_f3, _c3); + _f4 = _mm256_add_ps(_f4, _c4); + _f5 = _mm256_add_ps(_f5, _c5); + _f6 = _mm256_add_ps(_f6, _c6); + _f7 = _mm256_add_ps(_f7, _c7); + } + else + { + __m256 _beta = _mm256_set1_ps(beta); +#if __FMA__ + _f0 = _mm256_fmadd_ps(_c0, _beta, _f0); + _f1 = _mm256_fmadd_ps(_c1, _beta, _f1); + _f2 = _mm256_fmadd_ps(_c2, _beta, _f2); + _f3 = _mm256_fmadd_ps(_c3, _beta, _f3); + _f4 = _mm256_fmadd_ps(_c4, _beta, _f4); + _f5 = _mm256_fmadd_ps(_c5, _beta, _f5); + _f6 = _mm256_fmadd_ps(_c6, _beta, _f6); + _f7 = _mm256_fmadd_ps(_c7, _beta, _f7); +#else + _f0 = _mm256_add_ps(_f0, _mm256_mul_ps(_c0, _beta)); + _f1 = _mm256_add_ps(_f1, _mm256_mul_ps(_c1, _beta)); + _f2 = _mm256_add_ps(_f2, _mm256_mul_ps(_c2, _beta)); + _f3 = _mm256_add_ps(_f3, _mm256_mul_ps(_c3, _beta)); + _f4 = _mm256_add_ps(_f4, _mm256_mul_ps(_c4, _beta)); + _f5 = _mm256_add_ps(_f5, _mm256_mul_ps(_c5, _beta)); + _f6 = _mm256_add_ps(_f6, _mm256_mul_ps(_c6, _beta)); + _f7 = _mm256_add_ps(_f7, _mm256_mul_ps(_c7, _beta)); +#endif + } + pC += 8; + } + if (broadcast_type_C == 4) + { + __m256 _c = _mm256_loadu_ps(pC); + if (beta != 1.f) + _c = _mm256_mul_ps(_c, _mm256_set1_ps(beta)); + _f0 = _mm256_add_ps(_f0, _c); + _f1 = _mm256_add_ps(_f1, _c); + _f2 = _mm256_add_ps(_f2, _c); + _f3 = _mm256_add_ps(_f3, _c); + _f4 = _mm256_add_ps(_f4, _c); + _f5 = _mm256_add_ps(_f5, _c); + _f6 = _mm256_add_ps(_f6, _c); + _f7 = _mm256_add_ps(_f7, _c); + pC += 8; + } + } + if (alpha != 1.f) + { + __m256 _alpha = _mm256_set1_ps(alpha); + _f0 = _mm256_mul_ps(_f0, _alpha); + _f1 = _mm256_mul_ps(_f1, _alpha); + _f2 = _mm256_mul_ps(_f2, _alpha); + _f3 = _mm256_mul_ps(_f3, _alpha); + _f4 = _mm256_mul_ps(_f4, _alpha); + _f5 = _mm256_mul_ps(_f5, _alpha); + _f6 = _mm256_mul_ps(_f6, _alpha); + _f7 = _mm256_mul_ps(_f7, _alpha); + } + _mm256_storeu_ps(p0, _f0); + _mm256_storeu_ps(p0 + out_hstep, _f1); + _mm256_storeu_ps(p0 + out_hstep * 2, _f2); + _mm256_storeu_ps(p0 + out_hstep * 3, _f3); + _mm256_storeu_ps(p0 + out_hstep * 4, _f4); + _mm256_storeu_ps(p0 + out_hstep * 5, _f5); + _mm256_storeu_ps(p0 + out_hstep * 6, _f6); + _mm256_storeu_ps(p0 + out_hstep * 7, _f7); + p0 += 8; + pp += 64; + } +#endif // __AVX512F__ + for (; jj + 3 < max_jj; jj += 4) + { +#if __AVX2__ + __m256 _t0 = _mm256_loadu_ps(pp + 0); + __m256 _t1 = _mm256_loadu_ps(pp + 8); + __m256 _t2 = _mm256_loadu_ps(pp + 16); + __m256 _t3 = _mm256_loadu_ps(pp + 24); +#else + __m256 _t0 = combine4x2_ps(_mm_loadu_ps(pp + 0), _mm_loadu_ps(pp1 + 0)); + __m256 _t1 = combine4x2_ps(_mm_loadu_ps(pp + 4), _mm_loadu_ps(pp1 + 4)); + __m256 _t2 = combine4x2_ps(_mm_loadu_ps(pp + 8), _mm_loadu_ps(pp1 + 8)); + __m256 _t3 = combine4x2_ps(_mm_loadu_ps(pp + 12), _mm_loadu_ps(pp1 + 12)); +#endif + _t1 = _mm256_shuffle_ps(_t1, _t1, _MM_SHUFFLE(2, 1, 0, 3)); + _t3 = _mm256_shuffle_ps(_t3, _t3, _MM_SHUFFLE(2, 1, 0, 3)); + __m256 _tmp0 = _mm256_unpacklo_ps(_t0, _t3); + __m256 _tmp1 = _mm256_unpackhi_ps(_t0, _t3); + __m256 _tmp2 = _mm256_unpacklo_ps(_t2, _t1); + __m256 _tmp3 = _mm256_unpackhi_ps(_t2, _t1); + _t0 = _mm256_castpd_ps(_mm256_unpacklo_pd(_mm256_castps_pd(_tmp0), _mm256_castps_pd(_tmp2))); + _t1 = _mm256_castpd_ps(_mm256_unpackhi_pd(_mm256_castps_pd(_tmp0), _mm256_castps_pd(_tmp2))); + _t2 = _mm256_castpd_ps(_mm256_unpacklo_pd(_mm256_castps_pd(_tmp3), _mm256_castps_pd(_tmp1))); + _t3 = _mm256_castpd_ps(_mm256_unpackhi_pd(_mm256_castps_pd(_tmp3), _mm256_castps_pd(_tmp1))); + _t1 = _mm256_shuffle_ps(_t1, _t1, _MM_SHUFFLE(2, 1, 0, 3)); + _t3 = _mm256_shuffle_ps(_t3, _t3, _MM_SHUFFLE(2, 1, 0, 3)); + transpose8x4_ps(_t0, _t1, _t2, _t3); + __m128 _f0 = _mm256_extractf128_ps(_t0, 0); + __m128 _f1 = _mm256_extractf128_ps(_t0, 1); + __m128 _f2 = _mm256_extractf128_ps(_t1, 0); + __m128 _f3 = _mm256_extractf128_ps(_t1, 1); + __m128 _f4 = _mm256_extractf128_ps(_t2, 0); + __m128 _f5 = _mm256_extractf128_ps(_t2, 1); + __m128 _f6 = _mm256_extractf128_ps(_t3, 0); + __m128 _f7 = _mm256_extractf128_ps(_t3, 1); + if (pC) + { + if (broadcast_type_C == 0) + { + __m128 _c = _mm_set1_ps(c0); + _f0 = _mm_add_ps(_f0, _c); + _f1 = _mm_add_ps(_f1, _c); + _f2 = _mm_add_ps(_f2, _c); + _f3 = _mm_add_ps(_f3, _c); + _f4 = _mm_add_ps(_f4, _c); + _f5 = _mm_add_ps(_f5, _c); + _f6 = _mm_add_ps(_f6, _c); + _f7 = _mm_add_ps(_f7, _c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = _mm_add_ps(_f0, _mm_set1_ps(c0)); + _f1 = _mm_add_ps(_f1, _mm_set1_ps(c1)); + _f2 = _mm_add_ps(_f2, _mm_set1_ps(c2)); + _f3 = _mm_add_ps(_f3, _mm_set1_ps(c3)); + _f4 = _mm_add_ps(_f4, _mm_set1_ps(c4)); + _f5 = _mm_add_ps(_f5, _mm_set1_ps(c5)); + _f6 = _mm_add_ps(_f6, _mm_set1_ps(c6)); + _f7 = _mm_add_ps(_f7, _mm_set1_ps(c7)); + } + if (broadcast_type_C == 3) + { + __m128 _c0 = _mm_loadu_ps(pC); + __m128 _c1 = _mm_loadu_ps(pC + N); + __m128 _c2 = _mm_loadu_ps(pC + N * 2); + __m128 _c3 = _mm_loadu_ps(pC + N * 3); + __m128 _c4 = _mm_loadu_ps(pC + N * 4); + __m128 _c5 = _mm_loadu_ps(pC + N * 5); + __m128 _c6 = _mm_loadu_ps(pC + N * 6); + __m128 _c7 = _mm_loadu_ps(pC + N * 7); + if (beta == 1.f) + { + _f0 = _mm_add_ps(_f0, _c0); + _f1 = _mm_add_ps(_f1, _c1); + _f2 = _mm_add_ps(_f2, _c2); + _f3 = _mm_add_ps(_f3, _c3); + _f4 = _mm_add_ps(_f4, _c4); + _f5 = _mm_add_ps(_f5, _c5); + _f6 = _mm_add_ps(_f6, _c6); + _f7 = _mm_add_ps(_f7, _c7); + } + else + { + __m128 _beta = _mm_set1_ps(beta); +#if __FMA__ + _f0 = _mm_fmadd_ps(_c0, _beta, _f0); + _f1 = _mm_fmadd_ps(_c1, _beta, _f1); + _f2 = _mm_fmadd_ps(_c2, _beta, _f2); + _f3 = _mm_fmadd_ps(_c3, _beta, _f3); + _f4 = _mm_fmadd_ps(_c4, _beta, _f4); + _f5 = _mm_fmadd_ps(_c5, _beta, _f5); + _f6 = _mm_fmadd_ps(_c6, _beta, _f6); + _f7 = _mm_fmadd_ps(_c7, _beta, _f7); +#else + _f0 = _mm_add_ps(_f0, _mm_mul_ps(_c0, _beta)); + _f1 = _mm_add_ps(_f1, _mm_mul_ps(_c1, _beta)); + _f2 = _mm_add_ps(_f2, _mm_mul_ps(_c2, _beta)); + _f3 = _mm_add_ps(_f3, _mm_mul_ps(_c3, _beta)); + _f4 = _mm_add_ps(_f4, _mm_mul_ps(_c4, _beta)); + _f5 = _mm_add_ps(_f5, _mm_mul_ps(_c5, _beta)); + _f6 = _mm_add_ps(_f6, _mm_mul_ps(_c6, _beta)); + _f7 = _mm_add_ps(_f7, _mm_mul_ps(_c7, _beta)); +#endif + } + pC += 4; + } + if (broadcast_type_C == 4) + { + __m128 _c = _mm_loadu_ps(pC); + if (beta != 1.f) + _c = _mm_mul_ps(_c, _mm_set1_ps(beta)); + _f0 = _mm_add_ps(_f0, _c); + _f1 = _mm_add_ps(_f1, _c); + _f2 = _mm_add_ps(_f2, _c); + _f3 = _mm_add_ps(_f3, _c); + _f4 = _mm_add_ps(_f4, _c); + _f5 = _mm_add_ps(_f5, _c); + _f6 = _mm_add_ps(_f6, _c); + _f7 = _mm_add_ps(_f7, _c); + pC += 4; + } + } + if (alpha != 1.f) + { + __m128 _alpha = _mm_set1_ps(alpha); + _f0 = _mm_mul_ps(_f0, _alpha); + _f1 = _mm_mul_ps(_f1, _alpha); + _f2 = _mm_mul_ps(_f2, _alpha); + _f3 = _mm_mul_ps(_f3, _alpha); + _f4 = _mm_mul_ps(_f4, _alpha); + _f5 = _mm_mul_ps(_f5, _alpha); + _f6 = _mm_mul_ps(_f6, _alpha); + _f7 = _mm_mul_ps(_f7, _alpha); + } + _mm_storeu_ps(p0, _f0); + _mm_storeu_ps(p0 + out_hstep, _f1); + _mm_storeu_ps(p0 + out_hstep * 2, _f2); + _mm_storeu_ps(p0 + out_hstep * 3, _f3); + _mm_storeu_ps(p0 + out_hstep * 4, _f4); + _mm_storeu_ps(p0 + out_hstep * 5, _f5); + _mm_storeu_ps(p0 + out_hstep * 6, _f6); + _mm_storeu_ps(p0 + out_hstep * 7, _f7); + p0 += 4; +#if __AVX2__ + pp += 32; +#else + pp += 16; + pp1 += 16; +#endif + } +#endif // defined(__x86_64__) || defined(_M_X64) + for (; jj + 1 < max_jj; jj += 2) + { +#if __AVX2__ + __m256 _t0 = _mm256_loadu_ps(pp); + __m256 _t1 = _mm256_loadu_ps(pp + 8); +#else + __m256 _t0 = combine4x2_ps(_mm_loadu_ps(pp), _mm_loadu_ps(pp1)); + __m256 _t1 = combine4x2_ps(_mm_loadu_ps(pp + 4), _mm_loadu_ps(pp1 + 4)); +#endif + __m256 _tmp0 = _mm256_shuffle_ps(_t0, _t0, _MM_SHUFFLE(3, 1, 2, 0)); + __m256 _tmp1 = _mm256_shuffle_ps(_t1, _t1, _MM_SHUFFLE(0, 2, 3, 1)); + _t0 = _mm256_unpacklo_ps(_tmp0, _tmp1); + _t1 = _mm256_unpackhi_ps(_tmp0, _tmp1); + _t1 = _mm256_shuffle_ps(_t1, _t1, _MM_SHUFFLE(2, 1, 0, 3)); + transpose8x2_ps(_t0, _t1); + __m128 _r01 = _mm256_extractf128_ps(_t0, 0); + __m128 _r23 = _mm256_extractf128_ps(_t0, 1); + __m128 _r45 = _mm256_extractf128_ps(_t1, 0); + __m128 _r67 = _mm256_extractf128_ps(_t1, 1); + __m128 _f0 = _r01; + __m128 _f1 = _mm_movehl_ps(_r01, _r01); + __m128 _f2 = _r23; + __m128 _f3 = _mm_movehl_ps(_r23, _r23); + __m128 _f4 = _r45; + __m128 _f5 = _mm_movehl_ps(_r45, _r45); + __m128 _f6 = _r67; + __m128 _f7 = _mm_movehl_ps(_r67, _r67); + if (pC) + { + if (broadcast_type_C == 0) + { + __m128 _c = _mm_set1_ps(c0); + _f0 = _mm_add_ps(_f0, _c); + _f1 = _mm_add_ps(_f1, _c); + _f2 = _mm_add_ps(_f2, _c); + _f3 = _mm_add_ps(_f3, _c); + _f4 = _mm_add_ps(_f4, _c); + _f5 = _mm_add_ps(_f5, _c); + _f6 = _mm_add_ps(_f6, _c); + _f7 = _mm_add_ps(_f7, _c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = _mm_add_ps(_f0, _mm_set1_ps(c0)); + _f1 = _mm_add_ps(_f1, _mm_set1_ps(c1)); + _f2 = _mm_add_ps(_f2, _mm_set1_ps(c2)); + _f3 = _mm_add_ps(_f3, _mm_set1_ps(c3)); + _f4 = _mm_add_ps(_f4, _mm_set1_ps(c4)); + _f5 = _mm_add_ps(_f5, _mm_set1_ps(c5)); + _f6 = _mm_add_ps(_f6, _mm_set1_ps(c6)); + _f7 = _mm_add_ps(_f7, _mm_set1_ps(c7)); + } + if (broadcast_type_C == 3) + { + __m128 _c0 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pC); + __m128 _c1 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N)); + __m128 _c2 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N * 2)); + __m128 _c3 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N * 3)); + __m128 _c4 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N * 4)); + __m128 _c5 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N * 5)); + __m128 _c6 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N * 6)); + __m128 _c7 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N * 7)); + if (beta == 1.f) + { + _f0 = _mm_add_ps(_f0, _c0); + _f1 = _mm_add_ps(_f1, _c1); + _f2 = _mm_add_ps(_f2, _c2); + _f3 = _mm_add_ps(_f3, _c3); + _f4 = _mm_add_ps(_f4, _c4); + _f5 = _mm_add_ps(_f5, _c5); + _f6 = _mm_add_ps(_f6, _c6); + _f7 = _mm_add_ps(_f7, _c7); + } + else + { + __m128 _beta = _mm_set1_ps(beta); +#if __FMA__ + _f0 = _mm_fmadd_ps(_c0, _beta, _f0); + _f1 = _mm_fmadd_ps(_c1, _beta, _f1); + _f2 = _mm_fmadd_ps(_c2, _beta, _f2); + _f3 = _mm_fmadd_ps(_c3, _beta, _f3); + _f4 = _mm_fmadd_ps(_c4, _beta, _f4); + _f5 = _mm_fmadd_ps(_c5, _beta, _f5); + _f6 = _mm_fmadd_ps(_c6, _beta, _f6); + _f7 = _mm_fmadd_ps(_c7, _beta, _f7); +#else + _f0 = _mm_add_ps(_f0, _mm_mul_ps(_c0, _beta)); + _f1 = _mm_add_ps(_f1, _mm_mul_ps(_c1, _beta)); + _f2 = _mm_add_ps(_f2, _mm_mul_ps(_c2, _beta)); + _f3 = _mm_add_ps(_f3, _mm_mul_ps(_c3, _beta)); + _f4 = _mm_add_ps(_f4, _mm_mul_ps(_c4, _beta)); + _f5 = _mm_add_ps(_f5, _mm_mul_ps(_c5, _beta)); + _f6 = _mm_add_ps(_f6, _mm_mul_ps(_c6, _beta)); + _f7 = _mm_add_ps(_f7, _mm_mul_ps(_c7, _beta)); +#endif + } + pC += 2; + } + if (broadcast_type_C == 4) + { + __m128 _c = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pC); + if (beta != 1.f) + _c = _mm_mul_ps(_c, _mm_set1_ps(beta)); + _f0 = _mm_add_ps(_f0, _c); + _f1 = _mm_add_ps(_f1, _c); + _f2 = _mm_add_ps(_f2, _c); + _f3 = _mm_add_ps(_f3, _c); + _f4 = _mm_add_ps(_f4, _c); + _f5 = _mm_add_ps(_f5, _c); + _f6 = _mm_add_ps(_f6, _c); + _f7 = _mm_add_ps(_f7, _c); + pC += 2; + } + } + if (alpha != 1.f) + { + __m128 _alpha = _mm_set1_ps(alpha); + _f0 = _mm_mul_ps(_f0, _alpha); + _f1 = _mm_mul_ps(_f1, _alpha); + _f2 = _mm_mul_ps(_f2, _alpha); + _f3 = _mm_mul_ps(_f3, _alpha); + _f4 = _mm_mul_ps(_f4, _alpha); + _f5 = _mm_mul_ps(_f5, _alpha); + _f6 = _mm_mul_ps(_f6, _alpha); + _f7 = _mm_mul_ps(_f7, _alpha); + } + _mm_storel_pi((__m64*)p0, _f0); + _mm_storel_pi((__m64*)(p0 + out_hstep), _f1); + _mm_storel_pi((__m64*)(p0 + out_hstep * 2), _f2); + _mm_storel_pi((__m64*)(p0 + out_hstep * 3), _f3); + _mm_storel_pi((__m64*)(p0 + out_hstep * 4), _f4); + _mm_storel_pi((__m64*)(p0 + out_hstep * 5), _f5); + _mm_storel_pi((__m64*)(p0 + out_hstep * 6), _f6); + _mm_storel_pi((__m64*)(p0 + out_hstep * 7), _f7); + p0 += 2; +#if __AVX2__ + pp += 16; +#else + pp += 8; + pp1 += 8; +#endif + } + for (; jj < max_jj; jj++) + { +#if __AVX2__ + const float* pp4 = pp + 4; +#else + const float* pp4 = pp1; +#endif + float f0 = pp[0]; + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) f0 += c0; + if (broadcast_type_C == 3) f0 += pC[0] * beta; + if (broadcast_type_C == 4) f0 += pC[0] * beta; + } + if (alpha != 1.f) f0 *= alpha; + p0[0] = f0; + float f1 = pp[1]; + if (pC) + { + if (broadcast_type_C == 0) f1 += c0; + if (broadcast_type_C == 1 || broadcast_type_C == 2) f1 += c1; + if (broadcast_type_C == 3) f1 += pC[N] * beta; + if (broadcast_type_C == 4) f1 += pC[0] * beta; + } + if (alpha != 1.f) f1 *= alpha; + p0[out_hstep] = f1; + float f2 = pp[2]; + if (pC) + { + if (broadcast_type_C == 0) f2 += c0; + if (broadcast_type_C == 1 || broadcast_type_C == 2) f2 += c2; + if (broadcast_type_C == 3) f2 += pC[N * 2] * beta; + if (broadcast_type_C == 4) f2 += pC[0] * beta; + } + if (alpha != 1.f) f2 *= alpha; + p0[out_hstep * 2] = f2; + float f3 = pp[3]; + if (pC) + { + if (broadcast_type_C == 0) f3 += c0; + if (broadcast_type_C == 1 || broadcast_type_C == 2) f3 += c3; + if (broadcast_type_C == 3) f3 += pC[N * 3] * beta; + if (broadcast_type_C == 4) f3 += pC[0] * beta; + } + if (alpha != 1.f) f3 *= alpha; + p0[out_hstep * 3] = f3; + float f4 = pp4[0]; + if (pC) + { + if (broadcast_type_C == 0) f4 += c0; + if (broadcast_type_C == 1 || broadcast_type_C == 2) f4 += c4; + if (broadcast_type_C == 3) f4 += pC[N * 4] * beta; + if (broadcast_type_C == 4) f4 += pC[0] * beta; + } + if (alpha != 1.f) f4 *= alpha; + p0[out_hstep * 4] = f4; + float f5 = pp4[1]; + if (pC) + { + if (broadcast_type_C == 0) f5 += c0; + if (broadcast_type_C == 1 || broadcast_type_C == 2) f5 += c5; + if (broadcast_type_C == 3) f5 += pC[N * 5] * beta; + if (broadcast_type_C == 4) f5 += pC[0] * beta; + } + if (alpha != 1.f) f5 *= alpha; + p0[out_hstep * 5] = f5; + float f6 = pp4[2]; + if (pC) + { + if (broadcast_type_C == 0) f6 += c0; + if (broadcast_type_C == 1 || broadcast_type_C == 2) f6 += c6; + if (broadcast_type_C == 3) f6 += pC[N * 6] * beta; + if (broadcast_type_C == 4) f6 += pC[0] * beta; + } + if (alpha != 1.f) f6 *= alpha; + p0[out_hstep * 6] = f6; + float f7 = pp4[3]; + if (pC) + { + if (broadcast_type_C == 0) f7 += c0; + if (broadcast_type_C == 1 || broadcast_type_C == 2) f7 += c7; + if (broadcast_type_C == 3) + { + f7 += pC[N * 7] * beta; + pC++; + } + if (broadcast_type_C == 4) + { + f7 += pC[0] * beta; + pC++; + } + } + if (alpha != 1.f) f7 *= alpha; + p0[out_hstep * 7] = f7; + p0++; +#if __AVX2__ + pp += 8; +#else + pp += 4; + pp1 += 4; +#endif + } +#if !__AVX2__ + pp = pp1; + pp1 = pp + max_jj * 4; +#endif + } +#endif // __AVX__ + for (; ii + 3 < max_ii; ii += 4) + { + float* p0 = (float*)top_blob + (i + ii) * out_hstep + j; + + float c0 = 0.f; + float c1 = c0; + float c2 = c0; + float c3 = c0; + if (pC) + { + if (broadcast_type_C == 0) + { + c0 = pC[0]; + c1 = c0; + c2 = c0; + c3 = c0; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + c0 = pC[i + ii]; + c1 = pC[i + ii + 1]; + c2 = pC[i + ii + 2]; + c3 = pC[i + ii + 3]; + } + if (broadcast_type_C == 3) + pC = (const float*)C + (i + ii) * N + j; + if (broadcast_type_C == 4) + pC = (const float*)C + j; + if (broadcast_type_C == 0 && beta != 1.f) + c0 *= beta; + if ((broadcast_type_C == 1 || broadcast_type_C == 2) && beta != 1.f) + { + c0 *= beta; + c1 *= beta; + c2 *= beta; + c3 *= beta; + } + } + + int jj = 0; +#if defined(__x86_64__) || defined(_M_X64) +#if __AVX512F__ + for (; jj + 7 < max_jj; jj += 8) + { + __m256 _f0 = _mm256_loadu_ps(pp + 0); + __m256 _f1 = _mm256_loadu_ps(pp + 8); + __m256 _f2 = _mm256_loadu_ps(pp + 16); + __m256 _f3 = _mm256_loadu_ps(pp + 24); + __m128 _f00 = _mm256_castps256_ps128(_f0); + __m128 _f01 = _mm256_castps256_ps128(_f1); + __m128 _f02 = _mm256_castps256_ps128(_f2); + __m128 _f03 = _mm256_castps256_ps128(_f3); + __m128 _f10 = _mm256_extractf128_ps(_f0, 1); + __m128 _f11 = _mm256_extractf128_ps(_f1, 1); + __m128 _f12 = _mm256_extractf128_ps(_f2, 1); + __m128 _f13 = _mm256_extractf128_ps(_f3, 1); + { + _f01 = _mm_shuffle_ps(_f01, _f01, _MM_SHUFFLE(2, 1, 0, 3)); + _f03 = _mm_shuffle_ps(_f03, _f03, _MM_SHUFFLE(2, 1, 0, 3)); + __m128 _tmp0 = _mm_unpacklo_ps(_f00, _f03); + __m128 _tmp1 = _mm_unpackhi_ps(_f00, _f03); + __m128 _tmp2 = _mm_unpacklo_ps(_f02, _f01); + __m128 _tmp3 = _mm_unpackhi_ps(_f02, _f01); + _f00 = _mm_castpd_ps(_mm_unpacklo_pd(_mm_castps_pd(_tmp0), _mm_castps_pd(_tmp2))); + _f01 = _mm_castpd_ps(_mm_unpackhi_pd(_mm_castps_pd(_tmp0), _mm_castps_pd(_tmp2))); + _f02 = _mm_castpd_ps(_mm_unpacklo_pd(_mm_castps_pd(_tmp3), _mm_castps_pd(_tmp1))); + _f03 = _mm_castpd_ps(_mm_unpackhi_pd(_mm_castps_pd(_tmp3), _mm_castps_pd(_tmp1))); + _f01 = _mm_shuffle_ps(_f01, _f01, _MM_SHUFFLE(2, 1, 0, 3)); + _f03 = _mm_shuffle_ps(_f03, _f03, _MM_SHUFFLE(2, 1, 0, 3)); + _MM_TRANSPOSE4_PS(_f00, _f01, _f02, _f03); + } + { + _f11 = _mm_shuffle_ps(_f11, _f11, _MM_SHUFFLE(2, 1, 0, 3)); + _f13 = _mm_shuffle_ps(_f13, _f13, _MM_SHUFFLE(2, 1, 0, 3)); + __m128 _tmp0 = _mm_unpacklo_ps(_f10, _f13); + __m128 _tmp1 = _mm_unpackhi_ps(_f10, _f13); + __m128 _tmp2 = _mm_unpacklo_ps(_f12, _f11); + __m128 _tmp3 = _mm_unpackhi_ps(_f12, _f11); + _f10 = _mm_castpd_ps(_mm_unpacklo_pd(_mm_castps_pd(_tmp0), _mm_castps_pd(_tmp2))); + _f11 = _mm_castpd_ps(_mm_unpackhi_pd(_mm_castps_pd(_tmp0), _mm_castps_pd(_tmp2))); + _f12 = _mm_castpd_ps(_mm_unpacklo_pd(_mm_castps_pd(_tmp3), _mm_castps_pd(_tmp1))); + _f13 = _mm_castpd_ps(_mm_unpackhi_pd(_mm_castps_pd(_tmp3), _mm_castps_pd(_tmp1))); + _f11 = _mm_shuffle_ps(_f11, _f11, _MM_SHUFFLE(2, 1, 0, 3)); + _f13 = _mm_shuffle_ps(_f13, _f13, _MM_SHUFFLE(2, 1, 0, 3)); + _MM_TRANSPOSE4_PS(_f10, _f11, _f12, _f13); + } + _f0 = combine4x2_ps(_f00, _f10); + _f1 = combine4x2_ps(_f01, _f11); + _f2 = combine4x2_ps(_f02, _f12); + _f3 = combine4x2_ps(_f03, _f13); + if (pC) + { + if (broadcast_type_C == 0) + { + __m256 _c = _mm256_set1_ps(c0); + _f0 = _mm256_add_ps(_f0, _c); + _f1 = _mm256_add_ps(_f1, _c); + _f2 = _mm256_add_ps(_f2, _c); + _f3 = _mm256_add_ps(_f3, _c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = _mm256_add_ps(_f0, _mm256_set1_ps(c0)); + _f1 = _mm256_add_ps(_f1, _mm256_set1_ps(c1)); + _f2 = _mm256_add_ps(_f2, _mm256_set1_ps(c2)); + _f3 = _mm256_add_ps(_f3, _mm256_set1_ps(c3)); + } + if (broadcast_type_C == 3) + { + __m256 _c0 = _mm256_loadu_ps(pC); + __m256 _c1 = _mm256_loadu_ps(pC + N); + __m256 _c2 = _mm256_loadu_ps(pC + N * 2); + __m256 _c3 = _mm256_loadu_ps(pC + N * 3); + if (beta == 1.f) + { + _f0 = _mm256_add_ps(_f0, _c0); + _f1 = _mm256_add_ps(_f1, _c1); + _f2 = _mm256_add_ps(_f2, _c2); + _f3 = _mm256_add_ps(_f3, _c3); + } + else + { + __m256 _beta = _mm256_set1_ps(beta); +#if __FMA__ + _f0 = _mm256_fmadd_ps(_c0, _beta, _f0); + _f1 = _mm256_fmadd_ps(_c1, _beta, _f1); + _f2 = _mm256_fmadd_ps(_c2, _beta, _f2); + _f3 = _mm256_fmadd_ps(_c3, _beta, _f3); +#else + _f0 = _mm256_add_ps(_f0, _mm256_mul_ps(_c0, _beta)); + _f1 = _mm256_add_ps(_f1, _mm256_mul_ps(_c1, _beta)); + _f2 = _mm256_add_ps(_f2, _mm256_mul_ps(_c2, _beta)); + _f3 = _mm256_add_ps(_f3, _mm256_mul_ps(_c3, _beta)); +#endif + } + pC += 8; + } + if (broadcast_type_C == 4) + { + __m256 _c = _mm256_loadu_ps(pC); + if (beta != 1.f) + _c = _mm256_mul_ps(_c, _mm256_set1_ps(beta)); + _f0 = _mm256_add_ps(_f0, _c); + _f1 = _mm256_add_ps(_f1, _c); + _f2 = _mm256_add_ps(_f2, _c); + _f3 = _mm256_add_ps(_f3, _c); + pC += 8; + } + } + if (alpha != 1.f) + { + __m256 _alpha = _mm256_set1_ps(alpha); + _f0 = _mm256_mul_ps(_f0, _alpha); + _f1 = _mm256_mul_ps(_f1, _alpha); + _f2 = _mm256_mul_ps(_f2, _alpha); + _f3 = _mm256_mul_ps(_f3, _alpha); + } + _mm256_storeu_ps(p0, _f0); + _mm256_storeu_ps(p0 + out_hstep, _f1); + _mm256_storeu_ps(p0 + out_hstep * 2, _f2); + _mm256_storeu_ps(p0 + out_hstep * 3, _f3); + p0 += 8; + pp += 32; + } +#endif // __AVX512F__ + for (; jj + 3 < max_jj; jj += 4) + { + __m128 _f0 = _mm_loadu_ps(pp + 0); + __m128 _f1 = _mm_loadu_ps(pp + 4); + __m128 _f2 = _mm_loadu_ps(pp + 8); + __m128 _f3 = _mm_loadu_ps(pp + 12); + { + _f1 = _mm_shuffle_ps(_f1, _f1, _MM_SHUFFLE(2, 1, 0, 3)); + _f3 = _mm_shuffle_ps(_f3, _f3, _MM_SHUFFLE(2, 1, 0, 3)); + __m128 _tmp0 = _mm_unpacklo_ps(_f0, _f3); + __m128 _tmp1 = _mm_unpackhi_ps(_f0, _f3); + __m128 _tmp2 = _mm_unpacklo_ps(_f2, _f1); + __m128 _tmp3 = _mm_unpackhi_ps(_f2, _f1); + _f0 = _mm_castpd_ps(_mm_unpacklo_pd(_mm_castps_pd(_tmp0), _mm_castps_pd(_tmp2))); + _f1 = _mm_castpd_ps(_mm_unpackhi_pd(_mm_castps_pd(_tmp0), _mm_castps_pd(_tmp2))); + _f2 = _mm_castpd_ps(_mm_unpacklo_pd(_mm_castps_pd(_tmp3), _mm_castps_pd(_tmp1))); + _f3 = _mm_castpd_ps(_mm_unpackhi_pd(_mm_castps_pd(_tmp3), _mm_castps_pd(_tmp1))); + _f1 = _mm_shuffle_ps(_f1, _f1, _MM_SHUFFLE(2, 1, 0, 3)); + _f3 = _mm_shuffle_ps(_f3, _f3, _MM_SHUFFLE(2, 1, 0, 3)); + _MM_TRANSPOSE4_PS(_f0, _f1, _f2, _f3); + } + if (pC) + { + if (broadcast_type_C == 0) + { + __m128 _c = _mm_set1_ps(c0); + _f0 = _mm_add_ps(_f0, _c); + _f1 = _mm_add_ps(_f1, _c); + _f2 = _mm_add_ps(_f2, _c); + _f3 = _mm_add_ps(_f3, _c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = _mm_add_ps(_f0, _mm_set1_ps(c0)); + _f1 = _mm_add_ps(_f1, _mm_set1_ps(c1)); + _f2 = _mm_add_ps(_f2, _mm_set1_ps(c2)); + _f3 = _mm_add_ps(_f3, _mm_set1_ps(c3)); + } + if (broadcast_type_C == 3) + { + __m128 _c0 = _mm_loadu_ps(pC); + __m128 _c1 = _mm_loadu_ps(pC + N); + __m128 _c2 = _mm_loadu_ps(pC + N * 2); + __m128 _c3 = _mm_loadu_ps(pC + N * 3); + if (beta == 1.f) + { + _f0 = _mm_add_ps(_f0, _c0); + _f1 = _mm_add_ps(_f1, _c1); + _f2 = _mm_add_ps(_f2, _c2); + _f3 = _mm_add_ps(_f3, _c3); + } + else + { + __m128 _beta = _mm_set1_ps(beta); +#if __FMA__ + _f0 = _mm_fmadd_ps(_c0, _beta, _f0); + _f1 = _mm_fmadd_ps(_c1, _beta, _f1); + _f2 = _mm_fmadd_ps(_c2, _beta, _f2); + _f3 = _mm_fmadd_ps(_c3, _beta, _f3); +#else + _f0 = _mm_add_ps(_f0, _mm_mul_ps(_c0, _beta)); + _f1 = _mm_add_ps(_f1, _mm_mul_ps(_c1, _beta)); + _f2 = _mm_add_ps(_f2, _mm_mul_ps(_c2, _beta)); + _f3 = _mm_add_ps(_f3, _mm_mul_ps(_c3, _beta)); +#endif + } + pC += 4; + } + if (broadcast_type_C == 4) + { + __m128 _c = _mm_loadu_ps(pC); + if (beta != 1.f) + _c = _mm_mul_ps(_c, _mm_set1_ps(beta)); + _f0 = _mm_add_ps(_f0, _c); + _f1 = _mm_add_ps(_f1, _c); + _f2 = _mm_add_ps(_f2, _c); + _f3 = _mm_add_ps(_f3, _c); + pC += 4; + } + } + if (alpha != 1.f) + { + __m128 _alpha = _mm_set1_ps(alpha); + _f0 = _mm_mul_ps(_f0, _alpha); + _f1 = _mm_mul_ps(_f1, _alpha); + _f2 = _mm_mul_ps(_f2, _alpha); + _f3 = _mm_mul_ps(_f3, _alpha); + } + _mm_storeu_ps(p0, _f0); + _mm_storeu_ps(p0 + out_hstep, _f1); + _mm_storeu_ps(p0 + out_hstep * 2, _f2); + _mm_storeu_ps(p0 + out_hstep * 3, _f3); + p0 += 4; + pp += 16; + } +#endif // defined(__x86_64__) || defined(_M_X64) + for (; jj + 1 < max_jj; jj += 2) + { + __m128 _t0 = _mm_loadu_ps(pp); + __m128 _t1 = _mm_loadu_ps(pp + 4); + __m128 _tmp0 = _mm_shuffle_ps(_t0, _t0, _MM_SHUFFLE(3, 1, 2, 0)); + __m128 _tmp1 = _mm_shuffle_ps(_t1, _t1, _MM_SHUFFLE(0, 2, 3, 1)); + __m128 _c0v = _mm_unpacklo_ps(_tmp0, _tmp1); + __m128 _c1v = _mm_unpackhi_ps(_tmp0, _tmp1); + _c1v = _mm_shuffle_ps(_c1v, _c1v, _MM_SHUFFLE(2, 1, 0, 3)); + __m128 _f01 = _mm_unpacklo_ps(_c0v, _c1v); + __m128 _f23 = _mm_unpackhi_ps(_c0v, _c1v); + __m128 _f0 = _f01; + __m128 _f1 = _mm_movehl_ps(_f01, _f01); + __m128 _f2 = _f23; + __m128 _f3 = _mm_movehl_ps(_f23, _f23); + if (pC) + { + if (broadcast_type_C == 0) + { + __m128 _c = _mm_set1_ps(c0); + _f0 = _mm_add_ps(_f0, _c); + _f1 = _mm_add_ps(_f1, _c); + _f2 = _mm_add_ps(_f2, _c); + _f3 = _mm_add_ps(_f3, _c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = _mm_add_ps(_f0, _mm_set1_ps(c0)); + _f1 = _mm_add_ps(_f1, _mm_set1_ps(c1)); + _f2 = _mm_add_ps(_f2, _mm_set1_ps(c2)); + _f3 = _mm_add_ps(_f3, _mm_set1_ps(c3)); + } + if (broadcast_type_C == 3) + { + __m128 _c0 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pC); + __m128 _c1 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N)); + __m128 _c2 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N * 2)); + __m128 _c3 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N * 3)); + if (beta == 1.f) + { + _f0 = _mm_add_ps(_f0, _c0); + _f1 = _mm_add_ps(_f1, _c1); + _f2 = _mm_add_ps(_f2, _c2); + _f3 = _mm_add_ps(_f3, _c3); + } + else + { + __m128 _beta = _mm_set1_ps(beta); +#if __FMA__ + _f0 = _mm_fmadd_ps(_c0, _beta, _f0); + _f1 = _mm_fmadd_ps(_c1, _beta, _f1); + _f2 = _mm_fmadd_ps(_c2, _beta, _f2); + _f3 = _mm_fmadd_ps(_c3, _beta, _f3); +#else + _f0 = _mm_add_ps(_f0, _mm_mul_ps(_c0, _beta)); + _f1 = _mm_add_ps(_f1, _mm_mul_ps(_c1, _beta)); + _f2 = _mm_add_ps(_f2, _mm_mul_ps(_c2, _beta)); + _f3 = _mm_add_ps(_f3, _mm_mul_ps(_c3, _beta)); +#endif + } + pC += 2; + } + if (broadcast_type_C == 4) + { + __m128 _c = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pC); + if (beta != 1.f) + _c = _mm_mul_ps(_c, _mm_set1_ps(beta)); + _f0 = _mm_add_ps(_f0, _c); + _f1 = _mm_add_ps(_f1, _c); + _f2 = _mm_add_ps(_f2, _c); + _f3 = _mm_add_ps(_f3, _c); + pC += 2; + } + } + if (alpha != 1.f) + { + __m128 _alpha = _mm_set1_ps(alpha); + _f0 = _mm_mul_ps(_f0, _alpha); + _f1 = _mm_mul_ps(_f1, _alpha); + _f2 = _mm_mul_ps(_f2, _alpha); + _f3 = _mm_mul_ps(_f3, _alpha); + } + _mm_storel_pi((__m64*)(p0), _f0); + _mm_storel_pi((__m64*)(p0 + out_hstep), _f1); + _mm_storel_pi((__m64*)(p0 + out_hstep * 2), _f2); + _mm_storel_pi((__m64*)(p0 + out_hstep * 3), _f3); + p0 += 2; + pp += 8; + } + for (; jj < max_jj; jj += 1) + { + float f0_0 = pp[0]; + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) f0_0 += c0; + if (broadcast_type_C == 3 || broadcast_type_C == 4) f0_0 += pC[0] * beta; + } + if (alpha != 1.f) f0_0 *= alpha; + p0[0] = f0_0; + float f1_0 = pp[1]; + if (pC) + { + if (broadcast_type_C == 0) f1_0 += c0; + if (broadcast_type_C == 1 || broadcast_type_C == 2) f1_0 += c1; + if (broadcast_type_C == 3) f1_0 += pC[N] * beta; + if (broadcast_type_C == 4) f1_0 += pC[0] * beta; + } + if (alpha != 1.f) f1_0 *= alpha; + p0[out_hstep] = f1_0; + float f2_0 = pp[2]; + if (pC) + { + if (broadcast_type_C == 0) f2_0 += c0; + if (broadcast_type_C == 1 || broadcast_type_C == 2) f2_0 += c2; + if (broadcast_type_C == 3) f2_0 += pC[N * 2] * beta; + if (broadcast_type_C == 4) f2_0 += pC[0] * beta; + } + if (alpha != 1.f) f2_0 *= alpha; + p0[out_hstep * 2] = f2_0; + float f3_0 = pp[3]; + if (pC) + { + if (broadcast_type_C == 0) f3_0 += c0; + if (broadcast_type_C == 1 || broadcast_type_C == 2) f3_0 += c3; + if (broadcast_type_C == 3) + { + f3_0 += pC[N * 3] * beta; + pC++; + } + if (broadcast_type_C == 4) + { + f3_0 += pC[0] * beta; + pC++; + } + } + if (alpha != 1.f) f3_0 *= alpha; + p0[out_hstep * 3] = f3_0; + p0++; + pp += 4; + } + } + +#endif // __SSE2__ + + for (; ii + 1 < max_ii; ii += 2) + { + float* p0 = (float*)top_blob + (i + ii) * out_hstep + j; + + float c0 = 0.f; + float c1 = c0; + if (pC) + { + if (broadcast_type_C == 0) + { + c0 = pC[0]; + c1 = c0; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + c0 = pC[i + ii]; + c1 = pC[i + ii + 1]; + } + if (broadcast_type_C == 3) + pC = (const float*)C + (i + ii) * N + j; + if (broadcast_type_C == 4) + pC = (const float*)C + j; + if (broadcast_type_C == 0 && beta != 1.f) + c0 *= beta; + if ((broadcast_type_C == 1 || broadcast_type_C == 2) && beta != 1.f) + { + c0 *= beta; + c1 *= beta; + } + } + + int jj = 0; +#if __SSE2__ +#if defined(__x86_64__) || defined(_M_X64) +#if __AVX512F__ + for (; jj + 7 < max_jj; jj += 8) + { + __m128 _f0 = _mm_loadu_ps(pp); + __m128 _f1 = _mm_loadu_ps(pp + 4); + __m128 _f2 = _mm_loadu_ps(pp + 8); + __m128 _f3 = _mm_loadu_ps(pp + 12); + _f2 = _mm_shuffle_ps(_f2, _f2, _MM_SHUFFLE(2, 3, 0, 1)); + _f3 = _mm_shuffle_ps(_f3, _f3, _MM_SHUFFLE(2, 3, 0, 1)); + __m128 _tmp0 = _mm_unpacklo_ps(_f0, _f2); + __m128 _tmp1 = _mm_unpackhi_ps(_f0, _f2); + __m128 _tmp2 = _mm_unpacklo_ps(_f1, _f3); + __m128 _tmp3 = _mm_unpackhi_ps(_f1, _f3); + _f0 = _mm_castpd_ps(_mm_unpacklo_pd(_mm_castps_pd(_tmp0), _mm_castps_pd(_tmp1))); + _f1 = _mm_castpd_ps(_mm_unpacklo_pd(_mm_castps_pd(_tmp2), _mm_castps_pd(_tmp3))); + _f2 = _mm_castpd_ps(_mm_unpackhi_pd(_mm_castps_pd(_tmp0), _mm_castps_pd(_tmp1))); + _f3 = _mm_castpd_ps(_mm_unpackhi_pd(_mm_castps_pd(_tmp2), _mm_castps_pd(_tmp3))); + _f2 = _mm_shuffle_ps(_f2, _f2, _MM_SHUFFLE(2, 3, 0, 1)); + _f3 = _mm_shuffle_ps(_f3, _f3, _MM_SHUFFLE(2, 3, 0, 1)); + __m256 _f0x2 = combine4x2_ps(_f0, _f1); + __m256 _f1x2 = combine4x2_ps(_f2, _f3); + if (pC) + { + if (broadcast_type_C == 0) + { + __m256 _c = _mm256_set1_ps(c0); + _f0x2 = _mm256_add_ps(_f0x2, _c); + _f1x2 = _mm256_add_ps(_f1x2, _c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0x2 = _mm256_add_ps(_f0x2, _mm256_set1_ps(c0)); + _f1x2 = _mm256_add_ps(_f1x2, _mm256_set1_ps(c1)); + } + if (broadcast_type_C == 3) + { + __m256 _c0 = _mm256_loadu_ps(pC); + __m256 _c1 = _mm256_loadu_ps(pC + N); + if (beta == 1.f) + { + _f0x2 = _mm256_add_ps(_f0x2, _c0); + _f1x2 = _mm256_add_ps(_f1x2, _c1); + } + else + { + __m256 _beta = _mm256_set1_ps(beta); +#if __FMA__ + _f0x2 = _mm256_fmadd_ps(_c0, _beta, _f0x2); + _f1x2 = _mm256_fmadd_ps(_c1, _beta, _f1x2); +#else + _f0x2 = _mm256_add_ps(_f0x2, _mm256_mul_ps(_c0, _beta)); + _f1x2 = _mm256_add_ps(_f1x2, _mm256_mul_ps(_c1, _beta)); +#endif + } + pC += 8; + } + if (broadcast_type_C == 4) + { + __m256 _c = _mm256_loadu_ps(pC); + if (beta != 1.f) + _c = _mm256_mul_ps(_c, _mm256_set1_ps(beta)); + _f0x2 = _mm256_add_ps(_f0x2, _c); + _f1x2 = _mm256_add_ps(_f1x2, _c); + pC += 8; + } + } + if (alpha != 1.f) + { + __m256 _alpha = _mm256_set1_ps(alpha); + _f0x2 = _mm256_mul_ps(_f0x2, _alpha); + _f1x2 = _mm256_mul_ps(_f1x2, _alpha); + } + _mm256_storeu_ps(p0, _f0x2); + _mm256_storeu_ps(p0 + out_hstep, _f1x2); + p0 += 8; + pp += 16; + } +#endif // __AVX512F__ + for (; jj + 3 < max_jj; jj += 4) + { + __m128 _f0 = _mm_loadu_ps(pp + 0); + __m128 _f1 = _mm_loadu_ps(pp + 4); + __m128 _tmp0 = _mm_unpacklo_ps(_f0, _f1); + __m128 _tmp1 = _mm_unpackhi_ps(_f0, _f1); + _f0 = _mm_castpd_ps(_mm_unpacklo_pd(_mm_castps_pd(_tmp0), _mm_castps_pd(_tmp1))); + _f1 = _mm_castpd_ps(_mm_unpackhi_pd(_mm_castps_pd(_tmp1), _mm_castps_pd(_tmp0))); + _f1 = _mm_shuffle_ps(_f1, _f1, _MM_SHUFFLE(0, 3, 2, 1)); + if (pC) + { + if (broadcast_type_C == 0) + { + __m128 _c = _mm_set1_ps(c0); + _f0 = _mm_add_ps(_f0, _c); + _f1 = _mm_add_ps(_f1, _c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = _mm_add_ps(_f0, _mm_set1_ps(c0)); + _f1 = _mm_add_ps(_f1, _mm_set1_ps(c1)); + } + if (broadcast_type_C == 3) + { + __m128 _c0 = _mm_loadu_ps(pC); + __m128 _c1 = _mm_loadu_ps(pC + N); + if (beta == 1.f) + { + _f0 = _mm_add_ps(_f0, _c0); + _f1 = _mm_add_ps(_f1, _c1); + } + else + { + __m128 _beta = _mm_set1_ps(beta); +#if __FMA__ + _f0 = _mm_fmadd_ps(_c0, _beta, _f0); + _f1 = _mm_fmadd_ps(_c1, _beta, _f1); +#else + _f0 = _mm_add_ps(_f0, _mm_mul_ps(_c0, _beta)); + _f1 = _mm_add_ps(_f1, _mm_mul_ps(_c1, _beta)); +#endif + } + pC += 4; + } + if (broadcast_type_C == 4) + { + __m128 _c = _mm_loadu_ps(pC); + if (beta != 1.f) + _c = _mm_mul_ps(_c, _mm_set1_ps(beta)); + _f0 = _mm_add_ps(_f0, _c); + _f1 = _mm_add_ps(_f1, _c); + pC += 4; + } + } + if (alpha != 1.f) + { + __m128 _alpha = _mm_set1_ps(alpha); + _f0 = _mm_mul_ps(_f0, _alpha); + _f1 = _mm_mul_ps(_f1, _alpha); + } + _mm_storeu_ps(p0, _f0); + _mm_storeu_ps(p0 + out_hstep, _f1); + p0 += 4; + pp += 8; + } +#endif // defined(__x86_64__) || defined(_M_X64) + for (; jj + 1 < max_jj; jj += 2) + { + __m128 _f0 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pp + 0)); + __m128 _f1 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pp + 2)); + if (pC) + { + if (broadcast_type_C == 0) + { + __m128 _c = _mm_set1_ps(c0); + _f0 = _mm_add_ps(_f0, _c); + _f1 = _mm_add_ps(_f1, _c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = _mm_add_ps(_f0, _mm_set1_ps(c0)); + _f1 = _mm_add_ps(_f1, _mm_set1_ps(c1)); + } + if (broadcast_type_C == 3) + { + __m128 _c0 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pC); + __m128 _c1 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N)); + if (beta == 1.f) + { + _f0 = _mm_add_ps(_f0, _c0); + _f1 = _mm_add_ps(_f1, _c1); + } + else + { + __m128 _beta = _mm_set1_ps(beta); +#if __FMA__ + _f0 = _mm_fmadd_ps(_c0, _beta, _f0); + _f1 = _mm_fmadd_ps(_c1, _beta, _f1); +#else + _f0 = _mm_add_ps(_f0, _mm_mul_ps(_c0, _beta)); + _f1 = _mm_add_ps(_f1, _mm_mul_ps(_c1, _beta)); +#endif + } + pC += 2; + } + if (broadcast_type_C == 4) + { + __m128 _c = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pC); + if (beta != 1.f) + _c = _mm_mul_ps(_c, _mm_set1_ps(beta)); + _f0 = _mm_add_ps(_f0, _c); + _f1 = _mm_add_ps(_f1, _c); + pC += 2; + } + } + if (alpha != 1.f) + { + __m128 _alpha = _mm_set1_ps(alpha); + _f0 = _mm_mul_ps(_f0, _alpha); + _f1 = _mm_mul_ps(_f1, _alpha); + } + _mm_storel_pi((__m64*)(p0), _f0); + _mm_storel_pi((__m64*)(p0 + out_hstep), _f1); + p0 += 2; + pp += 4; + } + for (; jj < max_jj; jj += 1) + { + float f0_0 = pp[0]; + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) f0_0 += c0; + if (broadcast_type_C == 3 || broadcast_type_C == 4) f0_0 += pC[0] * beta; + } + if (alpha != 1.f) f0_0 *= alpha; + p0[0] = f0_0; + float f1_0 = pp[1]; + if (pC) + { + if (broadcast_type_C == 0) f1_0 += c0; + if (broadcast_type_C == 1 || broadcast_type_C == 2) f1_0 += c1; + if (broadcast_type_C == 3) + { + f1_0 += pC[N] * beta; + pC++; + } + if (broadcast_type_C == 4) + { + f1_0 += pC[0] * beta; + pC++; + } + } + if (alpha != 1.f) f1_0 *= alpha; + p0[out_hstep] = f1_0; + p0++; + pp += 2; + } +#else + for (; jj + 1 < max_jj; jj += 2) + { + float f0_0 = pp[0]; + float f0_1 = pp[1]; + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + f0_0 += c0; + f0_1 += c0; + } + if (broadcast_type_C == 3 || broadcast_type_C == 4) + { + f0_0 += pC[0] * beta; + f0_1 += pC[1] * beta; + } + } + if (alpha != 1.f) + { + f0_0 *= alpha; + f0_1 *= alpha; + } + p0[0] = f0_0; + p0[1] = f0_1; + + float f1_0 = pp[2]; + float f1_1 = pp[3]; + if (pC) + { + if (broadcast_type_C == 0) + { + f1_0 += c0; + f1_1 += c0; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + f1_0 += c1; + f1_1 += c1; + } + if (broadcast_type_C == 3) + { + f1_0 += pC[N] * beta; + f1_1 += pC[N + 1] * beta; + pC += 2; + } + if (broadcast_type_C == 4) + { + f1_0 += pC[0] * beta; + f1_1 += pC[1] * beta; + pC += 2; + } + } + if (alpha != 1.f) + { + f1_0 *= alpha; + f1_1 *= alpha; + } + p0[out_hstep] = f1_0; + p0[out_hstep + 1] = f1_1; + + p0 += 2; + pp += 4; + } + for (; jj < max_jj; jj += 1) + { + float f0_0 = pp[0]; + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) f0_0 += c0; + if (broadcast_type_C == 3 || broadcast_type_C == 4) f0_0 += pC[0] * beta; + } + if (alpha != 1.f) f0_0 *= alpha; + p0[0] = f0_0; + + float f1_0 = pp[1]; + if (pC) + { + if (broadcast_type_C == 0) f1_0 += c0; + if (broadcast_type_C == 1 || broadcast_type_C == 2) f1_0 += c1; + if (broadcast_type_C == 3) + { + f1_0 += pC[N] * beta; + pC++; + } + if (broadcast_type_C == 4) + { + f1_0 += pC[0] * beta; + pC++; + } + } + if (alpha != 1.f) f1_0 *= alpha; + p0[out_hstep] = f1_0; + + p0++; + pp += 2; + } +#endif // __SSE2__ + } + + for (; ii < max_ii; ii += 1) + { + float* p0 = (float*)top_blob + (i + ii) * out_hstep + j; + + float c0 = 0.f; + if (pC) + { + if (broadcast_type_C == 0) + c0 = pC[0]; + if (broadcast_type_C == 1 || broadcast_type_C == 2) + c0 = pC[i + ii]; + if (broadcast_type_C == 3) + pC = (const float*)C + (i + ii) * N + j; + if (broadcast_type_C == 4) + pC = (const float*)C + j; + if ((broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) && beta != 1.f) + c0 *= beta; + } + + int jj = 0; +#if __SSE2__ +#if defined(__x86_64__) || defined(_M_X64) +#if __AVX512F__ + for (; jj + 7 < max_jj; jj += 8) + { + __m256 _f0 = _mm256_loadu_ps(pp + 0); + if (pC) + { + if (broadcast_type_C == 0) + { + __m256 _c = _mm256_set1_ps(c0); + _f0 = _mm256_add_ps(_f0, _c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = _mm256_add_ps(_f0, _mm256_set1_ps(c0)); + } + if (broadcast_type_C == 3) + { + __m256 _c0 = _mm256_loadu_ps(pC); + if (beta == 1.f) + { + _f0 = _mm256_add_ps(_f0, _c0); + } + else + { + __m256 _beta = _mm256_set1_ps(beta); +#if __FMA__ + _f0 = _mm256_fmadd_ps(_c0, _beta, _f0); +#else + _f0 = _mm256_add_ps(_f0, _mm256_mul_ps(_c0, _beta)); +#endif + } + pC += 8; + } + if (broadcast_type_C == 4) + { + __m256 _c = _mm256_loadu_ps(pC); + if (beta != 1.f) + _c = _mm256_mul_ps(_c, _mm256_set1_ps(beta)); + _f0 = _mm256_add_ps(_f0, _c); + pC += 8; + } + } + if (alpha != 1.f) + _f0 = _mm256_mul_ps(_f0, _mm256_set1_ps(alpha)); + _mm256_storeu_ps(p0, _f0); + p0 += 8; + pp += 8; + } +#endif // __AVX512F__ + for (; jj + 3 < max_jj; jj += 4) + { + __m128 _f0 = _mm_loadu_ps(pp + 0); + if (pC) + { + if (broadcast_type_C == 0) + { + __m128 _c = _mm_set1_ps(c0); + _f0 = _mm_add_ps(_f0, _c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = _mm_add_ps(_f0, _mm_set1_ps(c0)); + } + if (broadcast_type_C == 3) + { + __m128 _c0 = _mm_loadu_ps(pC); + if (beta == 1.f) + { + _f0 = _mm_add_ps(_f0, _c0); + } + else + { + __m128 _beta = _mm_set1_ps(beta); +#if __FMA__ + _f0 = _mm_fmadd_ps(_c0, _beta, _f0); +#else + _f0 = _mm_add_ps(_f0, _mm_mul_ps(_c0, _beta)); +#endif + } + pC += 4; + } + if (broadcast_type_C == 4) + { + __m128 _c = _mm_loadu_ps(pC); + if (beta != 1.f) + _c = _mm_mul_ps(_c, _mm_set1_ps(beta)); + _f0 = _mm_add_ps(_f0, _c); + pC += 4; + } + } + if (alpha != 1.f) + _f0 = _mm_mul_ps(_f0, _mm_set1_ps(alpha)); + _mm_storeu_ps(p0, _f0); + p0 += 4; + pp += 4; + } +#endif // defined(__x86_64__) || defined(_M_X64) + for (; jj + 1 < max_jj; jj += 2) + { + __m128 _f0 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pp + 0)); + if (pC) + { + if (broadcast_type_C == 0) + { + __m128 _c = _mm_set1_ps(c0); + _f0 = _mm_add_ps(_f0, _c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = _mm_add_ps(_f0, _mm_set1_ps(c0)); + } + if (broadcast_type_C == 3) + { + __m128 _c0 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pC); + if (beta == 1.f) + { + _f0 = _mm_add_ps(_f0, _c0); + } + else + { + __m128 _beta = _mm_set1_ps(beta); +#if __FMA__ + _f0 = _mm_fmadd_ps(_c0, _beta, _f0); +#else + _f0 = _mm_add_ps(_f0, _mm_mul_ps(_c0, _beta)); +#endif + } + pC += 2; + } + if (broadcast_type_C == 4) + { + __m128 _c = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pC); + if (beta != 1.f) + _c = _mm_mul_ps(_c, _mm_set1_ps(beta)); + _f0 = _mm_add_ps(_f0, _c); + pC += 2; + } + } + if (alpha != 1.f) + _f0 = _mm_mul_ps(_f0, _mm_set1_ps(alpha)); + _mm_storel_pi((__m64*)(p0), _f0); + p0 += 2; + pp += 2; + } +#else + for (; jj + 1 < max_jj; jj += 2) + { + float f0_0 = pp[0]; + float f0_1 = pp[1]; + if (pC) + { + if (broadcast_type_C == 3 || broadcast_type_C == 4) + { + f0_0 += pC[0] * beta; + f0_1 += pC[1] * beta; + pC += 2; + } + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + f0_0 += c0; + f0_1 += c0; + } + } + if (alpha != 1.f) + { + f0_0 *= alpha; + f0_1 *= alpha; + } + p0[0] = f0_0; + p0[1] = f0_1; + p0 += 2; + pp += 2; + } +#endif // __SSE2__ + for (; jj < max_jj; jj += 1) + { + float f0_0 = pp[0]; + if (pC) + { + if (broadcast_type_C == 3 || broadcast_type_C == 4) + { + f0_0 += pC[0] * beta; + pC++; + } + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + f0_0 += c0; + } + if (alpha != 1.f) f0_0 *= alpha; + p0[0] = f0_0; + p0++; + pp += 1; + } + } +} + +static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta) +{ +#if NCNN_RUNTIME_CPU && NCNN_AVX512VNNI && __AVX512F__ && !__AVX512VNNI__ + if (ncnn::cpu_support_x86_avx512_vnni()) + { + transpose_unpack_output_tile_wq_int8_avx512vnni(topT, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, N, alpha, beta); + return; + } +#endif +#if NCNN_RUNTIME_CPU && NCNN_AVXVNNIINT8 && __AVX__ && !__AVXVNNIINT8__ && !__AVX512VNNI__ + if (ncnn::cpu_support_x86_avx_vnni_int8()) + { + transpose_unpack_output_tile_wq_int8_avxvnniint8(topT, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, N, alpha, beta); + return; + } +#endif +#if NCNN_RUNTIME_CPU && NCNN_AVXVNNI && __AVX__ && !__AVXVNNI__ && !__AVXVNNIINT8__ && !__AVX512VNNI__ + if (ncnn::cpu_support_x86_avx_vnni()) + { + transpose_unpack_output_tile_wq_int8_avxvnni(topT, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, N, alpha, beta); + return; + } +#endif +#if NCNN_RUNTIME_CPU && NCNN_AVX2 && __AVX__ && !__AVX2__ && !__AVXVNNI__ && !__AVXVNNIINT8__ && !__AVX512VNNI__ + if (ncnn::cpu_support_x86_avx2()) + { + transpose_unpack_output_tile_wq_int8_avx2(topT, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, N, alpha, beta); + return; + } +#endif + + const float* pC = C; + const float* pp = topT; + const size_t out_hstep = top_blob.dims == 3 ? top_blob.cstep : (size_t)top_blob.w; + + int ii = 0; +#if __SSE2__ +#if __AVX__ +#if __AVX512F__ + for (; ii + 15 < max_ii; ii += 16) + { + float* p0 = (float*)top_blob + j * out_hstep + i + ii; + + __m512 _c = _mm512_setzero_ps(); + __m512i _c_vindex = _mm512_setzero_si512(); + if (pC) + { + if (broadcast_type_C == 0) + _c = _mm512_set1_ps(pC[0]); + if (broadcast_type_C == 1 || broadcast_type_C == 2) + _c = _mm512_loadu_ps(pC + i + ii); + if (broadcast_type_C == 3) + { + pC = (const float*)C + (i + ii) * N + j; + _c_vindex = _mm512_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15); + _c_vindex = _mm512_mullo_epi32(_c_vindex, _mm512_set1_epi32(N)); + } + if (broadcast_type_C == 4) + pC = (const float*)C + j; + if ((broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) && beta != 1.f) + _c = _mm512_mul_ps(_c, _mm512_set1_ps(beta)); + } + + int jj = 0; +#if defined(__x86_64__) || defined(_M_X64) + for (; jj + 7 < max_jj; jj += 8) + { + __m512 _f0 = _mm512_loadu_ps(pp + 0); + __m512 _f1 = _mm512_loadu_ps(pp + 16); + __m512 _f2 = _mm512_loadu_ps(pp + 32); + __m512 _f3 = _mm512_loadu_ps(pp + 48); + __m512 _f4 = _mm512_loadu_ps(pp + 64); + __m512 _f5 = _mm512_loadu_ps(pp + 80); + __m512 _f6 = _mm512_loadu_ps(pp + 96); + __m512 _f7 = _mm512_loadu_ps(pp + 112); + pp += 128; + _f1 = _mm512_permute_ps(_f1, _MM_SHUFFLE(2, 1, 0, 3)); + _f3 = _mm512_permute_ps(_f3, _MM_SHUFFLE(2, 1, 0, 3)); + _f5 = _mm512_permute_ps(_f5, _MM_SHUFFLE(2, 1, 0, 3)); + _f7 = _mm512_permute_ps(_f7, _MM_SHUFFLE(2, 1, 0, 3)); + __m512 _tmp0 = _mm512_unpacklo_ps(_f0, _f3); + __m512 _tmp1 = _mm512_unpackhi_ps(_f0, _f3); + __m512 _tmp2 = _mm512_unpacklo_ps(_f2, _f1); + __m512 _tmp3 = _mm512_unpackhi_ps(_f2, _f1); + __m512 _tmp4 = _mm512_unpacklo_ps(_f4, _f7); + __m512 _tmp5 = _mm512_unpackhi_ps(_f4, _f7); + __m512 _tmp6 = _mm512_unpacklo_ps(_f6, _f5); + __m512 _tmp7 = _mm512_unpackhi_ps(_f6, _f5); + _f0 = _mm512_castpd_ps(_mm512_unpacklo_pd(_mm512_castps_pd(_tmp0), _mm512_castps_pd(_tmp2))); + _f1 = _mm512_castpd_ps(_mm512_unpackhi_pd(_mm512_castps_pd(_tmp0), _mm512_castps_pd(_tmp2))); + _f2 = _mm512_castpd_ps(_mm512_unpacklo_pd(_mm512_castps_pd(_tmp3), _mm512_castps_pd(_tmp1))); + _f3 = _mm512_castpd_ps(_mm512_unpackhi_pd(_mm512_castps_pd(_tmp3), _mm512_castps_pd(_tmp1))); + _f4 = _mm512_castpd_ps(_mm512_unpacklo_pd(_mm512_castps_pd(_tmp4), _mm512_castps_pd(_tmp6))); + _f5 = _mm512_castpd_ps(_mm512_unpackhi_pd(_mm512_castps_pd(_tmp4), _mm512_castps_pd(_tmp6))); + _f6 = _mm512_castpd_ps(_mm512_unpacklo_pd(_mm512_castps_pd(_tmp7), _mm512_castps_pd(_tmp5))); + _f7 = _mm512_castpd_ps(_mm512_unpackhi_pd(_mm512_castps_pd(_tmp7), _mm512_castps_pd(_tmp5))); + _f1 = _mm512_permute_ps(_f1, _MM_SHUFFLE(2, 1, 0, 3)); + _f3 = _mm512_permute_ps(_f3, _MM_SHUFFLE(2, 1, 0, 3)); + _f5 = _mm512_permute_ps(_f5, _MM_SHUFFLE(2, 1, 0, 3)); + _f7 = _mm512_permute_ps(_f7, _MM_SHUFFLE(2, 1, 0, 3)); + _tmp0 = _mm512_shuffle_f32x4(_f0, _f4, _MM_SHUFFLE(0, 1, 1, 0)); + _tmp1 = _mm512_shuffle_f32x4(_f1, _f5, _MM_SHUFFLE(0, 1, 1, 0)); + _tmp2 = _mm512_shuffle_f32x4(_f2, _f6, _MM_SHUFFLE(0, 1, 1, 0)); + _tmp3 = _mm512_shuffle_f32x4(_f3, _f7, _MM_SHUFFLE(0, 1, 1, 0)); + _tmp4 = _mm512_shuffle_f32x4(_f0, _f4, _MM_SHUFFLE(2, 3, 3, 2)); + _tmp5 = _mm512_shuffle_f32x4(_f1, _f5, _MM_SHUFFLE(2, 3, 3, 2)); + _tmp6 = _mm512_shuffle_f32x4(_f2, _f6, _MM_SHUFFLE(2, 3, 3, 2)); + _tmp7 = _mm512_shuffle_f32x4(_f3, _f7, _MM_SHUFFLE(2, 3, 3, 2)); + _f0 = _mm512_shuffle_f32x4(_tmp0, _tmp4, _MM_SHUFFLE(2, 0, 2, 0)); + _f1 = _mm512_shuffle_f32x4(_tmp1, _tmp5, _MM_SHUFFLE(2, 0, 2, 0)); + _f2 = _mm512_shuffle_f32x4(_tmp2, _tmp6, _MM_SHUFFLE(2, 0, 2, 0)); + _f3 = _mm512_shuffle_f32x4(_tmp3, _tmp7, _MM_SHUFFLE(2, 0, 2, 0)); + _f4 = _mm512_shuffle_f32x4(_tmp0, _tmp4, _MM_SHUFFLE(1, 3, 1, 3)); + _f5 = _mm512_shuffle_f32x4(_tmp1, _tmp5, _MM_SHUFFLE(1, 3, 1, 3)); + _f6 = _mm512_shuffle_f32x4(_tmp2, _tmp6, _MM_SHUFFLE(1, 3, 1, 3)); + _f7 = _mm512_shuffle_f32x4(_tmp3, _tmp7, _MM_SHUFFLE(1, 3, 1, 3)); + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = _mm512_add_ps(_f0, _c); + _f1 = _mm512_add_ps(_f1, _c); + _f2 = _mm512_add_ps(_f2, _c); + _f3 = _mm512_add_ps(_f3, _c); + _f4 = _mm512_add_ps(_f4, _c); + _f5 = _mm512_add_ps(_f5, _c); + _f6 = _mm512_add_ps(_f6, _c); + _f7 = _mm512_add_ps(_f7, _c); + } + if (broadcast_type_C == 3) + { + __m512 _c0 = _mm512_i32gather_ps(_c_vindex, pC, sizeof(float)); + if (beta != 1.f) _c0 = _mm512_mul_ps(_c0, _mm512_set1_ps(beta)); + _f0 = _mm512_add_ps(_f0, _c0); + __m512 _c1 = _mm512_i32gather_ps(_c_vindex, pC + 1, sizeof(float)); + if (beta != 1.f) _c1 = _mm512_mul_ps(_c1, _mm512_set1_ps(beta)); + _f1 = _mm512_add_ps(_f1, _c1); + __m512 _c2 = _mm512_i32gather_ps(_c_vindex, pC + 2, sizeof(float)); + if (beta != 1.f) _c2 = _mm512_mul_ps(_c2, _mm512_set1_ps(beta)); + _f2 = _mm512_add_ps(_f2, _c2); + __m512 _c3 = _mm512_i32gather_ps(_c_vindex, pC + 3, sizeof(float)); + if (beta != 1.f) _c3 = _mm512_mul_ps(_c3, _mm512_set1_ps(beta)); + _f3 = _mm512_add_ps(_f3, _c3); + __m512 _c4 = _mm512_i32gather_ps(_c_vindex, pC + 4, sizeof(float)); + if (beta != 1.f) _c4 = _mm512_mul_ps(_c4, _mm512_set1_ps(beta)); + _f4 = _mm512_add_ps(_f4, _c4); + __m512 _c5 = _mm512_i32gather_ps(_c_vindex, pC + 5, sizeof(float)); + if (beta != 1.f) _c5 = _mm512_mul_ps(_c5, _mm512_set1_ps(beta)); + _f5 = _mm512_add_ps(_f5, _c5); + __m512 _c6 = _mm512_i32gather_ps(_c_vindex, pC + 6, sizeof(float)); + if (beta != 1.f) _c6 = _mm512_mul_ps(_c6, _mm512_set1_ps(beta)); + _f6 = _mm512_add_ps(_f6, _c6); + __m512 _c7 = _mm512_i32gather_ps(_c_vindex, pC + 7, sizeof(float)); + if (beta != 1.f) _c7 = _mm512_mul_ps(_c7, _mm512_set1_ps(beta)); + _f7 = _mm512_add_ps(_f7, _c7); + pC += 8; + } + if (broadcast_type_C == 4) + { + __m512 _c0 = _mm512_set1_ps(pC[0]); + if (beta != 1.f) _c0 = _mm512_mul_ps(_c0, _mm512_set1_ps(beta)); + _f0 = _mm512_add_ps(_f0, _c0); + __m512 _c1 = _mm512_set1_ps(pC[1]); + if (beta != 1.f) _c1 = _mm512_mul_ps(_c1, _mm512_set1_ps(beta)); + _f1 = _mm512_add_ps(_f1, _c1); + __m512 _c2 = _mm512_set1_ps(pC[2]); + if (beta != 1.f) _c2 = _mm512_mul_ps(_c2, _mm512_set1_ps(beta)); + _f2 = _mm512_add_ps(_f2, _c2); + __m512 _c3 = _mm512_set1_ps(pC[3]); + if (beta != 1.f) _c3 = _mm512_mul_ps(_c3, _mm512_set1_ps(beta)); + _f3 = _mm512_add_ps(_f3, _c3); + __m512 _c4 = _mm512_set1_ps(pC[4]); + if (beta != 1.f) _c4 = _mm512_mul_ps(_c4, _mm512_set1_ps(beta)); + _f4 = _mm512_add_ps(_f4, _c4); + __m512 _c5 = _mm512_set1_ps(pC[5]); + if (beta != 1.f) _c5 = _mm512_mul_ps(_c5, _mm512_set1_ps(beta)); + _f5 = _mm512_add_ps(_f5, _c5); + __m512 _c6 = _mm512_set1_ps(pC[6]); + if (beta != 1.f) _c6 = _mm512_mul_ps(_c6, _mm512_set1_ps(beta)); + _f6 = _mm512_add_ps(_f6, _c6); + __m512 _c7 = _mm512_set1_ps(pC[7]); + if (beta != 1.f) _c7 = _mm512_mul_ps(_c7, _mm512_set1_ps(beta)); + _f7 = _mm512_add_ps(_f7, _c7); + pC += 8; + } + } + if (alpha != 1.f) + { + const __m512 _alpha = _mm512_set1_ps(alpha); + _f0 = _mm512_mul_ps(_f0, _alpha); + _f1 = _mm512_mul_ps(_f1, _alpha); + _f2 = _mm512_mul_ps(_f2, _alpha); + _f3 = _mm512_mul_ps(_f3, _alpha); + _f4 = _mm512_mul_ps(_f4, _alpha); + _f5 = _mm512_mul_ps(_f5, _alpha); + _f6 = _mm512_mul_ps(_f6, _alpha); + _f7 = _mm512_mul_ps(_f7, _alpha); + } + _mm512_storeu_ps(p0, _f0); + _mm512_storeu_ps(p0 + out_hstep, _f1); + _mm512_storeu_ps(p0 + out_hstep * 2, _f2); + _mm512_storeu_ps(p0 + out_hstep * 3, _f3); + _mm512_storeu_ps(p0 + out_hstep * 4, _f4); + _mm512_storeu_ps(p0 + out_hstep * 5, _f5); + _mm512_storeu_ps(p0 + out_hstep * 6, _f6); + _mm512_storeu_ps(p0 + out_hstep * 7, _f7); + p0 += out_hstep * 8; + + } + for (; jj + 3 < max_jj; jj += 4) + { + __m512 _f0 = _mm512_loadu_ps(pp + 0); + __m512 _f1 = _mm512_loadu_ps(pp + 16); + __m512 _f2 = _mm512_loadu_ps(pp + 32); + __m512 _f3 = _mm512_loadu_ps(pp + 48); + pp += 64; + _f1 = _mm512_permute_ps(_f1, _MM_SHUFFLE(2, 1, 0, 3)); + _f3 = _mm512_permute_ps(_f3, _MM_SHUFFLE(2, 1, 0, 3)); + __m512 _tmp0 = _mm512_unpacklo_ps(_f0, _f3); + __m512 _tmp1 = _mm512_unpackhi_ps(_f0, _f3); + __m512 _tmp2 = _mm512_unpacklo_ps(_f2, _f1); + __m512 _tmp3 = _mm512_unpackhi_ps(_f2, _f1); + _f0 = _mm512_castpd_ps(_mm512_unpacklo_pd(_mm512_castps_pd(_tmp0), _mm512_castps_pd(_tmp2))); + _f1 = _mm512_castpd_ps(_mm512_unpackhi_pd(_mm512_castps_pd(_tmp0), _mm512_castps_pd(_tmp2))); + _f2 = _mm512_castpd_ps(_mm512_unpacklo_pd(_mm512_castps_pd(_tmp3), _mm512_castps_pd(_tmp1))); + _f3 = _mm512_castpd_ps(_mm512_unpackhi_pd(_mm512_castps_pd(_tmp3), _mm512_castps_pd(_tmp1))); + _f1 = _mm512_permute_ps(_f1, _MM_SHUFFLE(2, 1, 0, 3)); + _f3 = _mm512_permute_ps(_f3, _MM_SHUFFLE(2, 1, 0, 3)); + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = _mm512_add_ps(_f0, _c); + _f1 = _mm512_add_ps(_f1, _c); + _f2 = _mm512_add_ps(_f2, _c); + _f3 = _mm512_add_ps(_f3, _c); + } + if (broadcast_type_C == 3) + { + __m512 _c0 = _mm512_i32gather_ps(_c_vindex, pC, sizeof(float)); + if (beta != 1.f) _c0 = _mm512_mul_ps(_c0, _mm512_set1_ps(beta)); + _f0 = _mm512_add_ps(_f0, _c0); + __m512 _c1 = _mm512_i32gather_ps(_c_vindex, pC + 1, sizeof(float)); + if (beta != 1.f) _c1 = _mm512_mul_ps(_c1, _mm512_set1_ps(beta)); + _f1 = _mm512_add_ps(_f1, _c1); + __m512 _c2 = _mm512_i32gather_ps(_c_vindex, pC + 2, sizeof(float)); + if (beta != 1.f) _c2 = _mm512_mul_ps(_c2, _mm512_set1_ps(beta)); + _f2 = _mm512_add_ps(_f2, _c2); + __m512 _c3 = _mm512_i32gather_ps(_c_vindex, pC + 3, sizeof(float)); + if (beta != 1.f) _c3 = _mm512_mul_ps(_c3, _mm512_set1_ps(beta)); + _f3 = _mm512_add_ps(_f3, _c3); + pC += 4; + } + if (broadcast_type_C == 4) + { + __m512 _c0 = _mm512_set1_ps(pC[0]); + if (beta != 1.f) _c0 = _mm512_mul_ps(_c0, _mm512_set1_ps(beta)); + _f0 = _mm512_add_ps(_f0, _c0); + __m512 _c1 = _mm512_set1_ps(pC[1]); + if (beta != 1.f) _c1 = _mm512_mul_ps(_c1, _mm512_set1_ps(beta)); + _f1 = _mm512_add_ps(_f1, _c1); + __m512 _c2 = _mm512_set1_ps(pC[2]); + if (beta != 1.f) _c2 = _mm512_mul_ps(_c2, _mm512_set1_ps(beta)); + _f2 = _mm512_add_ps(_f2, _c2); + __m512 _c3 = _mm512_set1_ps(pC[3]); + if (beta != 1.f) _c3 = _mm512_mul_ps(_c3, _mm512_set1_ps(beta)); + _f3 = _mm512_add_ps(_f3, _c3); + pC += 4; + } + } + if (alpha != 1.f) + { + const __m512 _alpha = _mm512_set1_ps(alpha); + _f0 = _mm512_mul_ps(_f0, _alpha); + _f1 = _mm512_mul_ps(_f1, _alpha); + _f2 = _mm512_mul_ps(_f2, _alpha); + _f3 = _mm512_mul_ps(_f3, _alpha); + } + _mm512_storeu_ps(p0, _f0); + _mm512_storeu_ps(p0 + out_hstep, _f1); + _mm512_storeu_ps(p0 + out_hstep * 2, _f2); + _mm512_storeu_ps(p0 + out_hstep * 3, _f3); + p0 += out_hstep * 4; + + } +#endif // defined(__x86_64__) || defined(_M_X64) + for (; jj + 1 < max_jj; jj += 2) + { + __m512 _f0 = _mm512_loadu_ps(pp + 0); + __m512 _f1 = _mm512_loadu_ps(pp + 16); + pp += 32; + __m512 _tmp0 = _mm512_permute_ps(_f0, _MM_SHUFFLE(3, 1, 2, 0)); + __m512 _tmp1 = _mm512_permute_ps(_f1, _MM_SHUFFLE(0, 2, 3, 1)); + _f0 = _mm512_unpacklo_ps(_tmp0, _tmp1); + _f1 = _mm512_unpackhi_ps(_tmp0, _tmp1); + _f1 = _mm512_permute_ps(_f1, _MM_SHUFFLE(2, 1, 0, 3)); + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = _mm512_add_ps(_f0, _c); + _f1 = _mm512_add_ps(_f1, _c); + } + if (broadcast_type_C == 3) + { + __m512 _c0 = _mm512_i32gather_ps(_c_vindex, pC, sizeof(float)); + if (beta != 1.f) _c0 = _mm512_mul_ps(_c0, _mm512_set1_ps(beta)); + _f0 = _mm512_add_ps(_f0, _c0); + __m512 _c1 = _mm512_i32gather_ps(_c_vindex, pC + 1, sizeof(float)); + if (beta != 1.f) _c1 = _mm512_mul_ps(_c1, _mm512_set1_ps(beta)); + _f1 = _mm512_add_ps(_f1, _c1); + pC += 2; + } + if (broadcast_type_C == 4) + { + __m512 _c0 = _mm512_set1_ps(pC[0]); + if (beta != 1.f) _c0 = _mm512_mul_ps(_c0, _mm512_set1_ps(beta)); + _f0 = _mm512_add_ps(_f0, _c0); + __m512 _c1 = _mm512_set1_ps(pC[1]); + if (beta != 1.f) _c1 = _mm512_mul_ps(_c1, _mm512_set1_ps(beta)); + _f1 = _mm512_add_ps(_f1, _c1); + pC += 2; + } + } + if (alpha != 1.f) + { + const __m512 _alpha = _mm512_set1_ps(alpha); + _f0 = _mm512_mul_ps(_f0, _alpha); + _f1 = _mm512_mul_ps(_f1, _alpha); + } + _mm512_storeu_ps(p0, _f0); + _mm512_storeu_ps(p0 + out_hstep, _f1); + p0 += out_hstep * 2; + + } + for (; jj < max_jj; jj++) + { + __m512 _f0 = _mm512_loadu_ps(pp + 0); + pp += 16; + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = _mm512_add_ps(_f0, _c); + } + if (broadcast_type_C == 3) + { + __m512 _c0 = _mm512_i32gather_ps(_c_vindex, pC, sizeof(float)); + if (beta != 1.f) _c0 = _mm512_mul_ps(_c0, _mm512_set1_ps(beta)); + _f0 = _mm512_add_ps(_f0, _c0); + pC++; + } + if (broadcast_type_C == 4) + { + __m512 _c0 = _mm512_set1_ps(pC[0]); + if (beta != 1.f) _c0 = _mm512_mul_ps(_c0, _mm512_set1_ps(beta)); + _f0 = _mm512_add_ps(_f0, _c0); + pC++; + } + } + if (alpha != 1.f) + { + const __m512 _alpha = _mm512_set1_ps(alpha); + _f0 = _mm512_mul_ps(_f0, _alpha); + } + _mm512_storeu_ps(p0, _f0); + p0 += out_hstep; + + } + } +#endif // __AVX512F__ +#if !__AVX2__ + const float* pp1 = pp + max_jj * 4; +#endif + for (; ii + 7 < max_ii; ii += 8) + { + float* p0 = (float*)top_blob + j * out_hstep + i + ii; + + float c0 = 0.f; + float c1 = c0; + float c2 = c0; + float c3 = c0; + float c4 = c0; + float c5 = c0; + float c6 = c0; + float c7 = c0; + if (pC) + { + if (broadcast_type_C == 0) + { + c0 = pC[0]; + c1 = c0; + c2 = c0; + c3 = c0; + c4 = c0; + c5 = c0; + c6 = c0; + c7 = c0; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + c0 = pC[i + ii]; + c1 = pC[i + ii + 1]; + c2 = pC[i + ii + 2]; + c3 = pC[i + ii + 3]; + c4 = pC[i + ii + 4]; + c5 = pC[i + ii + 5]; + c6 = pC[i + ii + 6]; + c7 = pC[i + ii + 7]; + } + if (broadcast_type_C == 3) + pC = (const float*)C + (i + ii) * N + j; + if (broadcast_type_C == 4) + pC = (const float*)C + j; + if (broadcast_type_C == 0 && beta != 1.f) + c0 *= beta; + if ((broadcast_type_C == 1 || broadcast_type_C == 2) && beta != 1.f) + { + c0 *= beta; + c1 *= beta; + c2 *= beta; + c3 *= beta; + c4 *= beta; + c5 *= beta; + c6 *= beta; + c7 *= beta; + } + } + + int jj = 0; +#if defined(__x86_64__) || defined(_M_X64) +#if __AVX512F__ + for (; jj + 7 < max_jj; jj += 8) + { + __m256 _f0 = _mm256_loadu_ps(pp + 0); + __m256 _f1 = _mm256_loadu_ps(pp + 8); + __m256 _f2 = _mm256_loadu_ps(pp + 16); + __m256 _f3 = _mm256_loadu_ps(pp + 24); + __m256 _f4 = _mm256_loadu_ps(pp + 32); + __m256 _f5 = _mm256_loadu_ps(pp + 40); + __m256 _f6 = _mm256_loadu_ps(pp + 48); + __m256 _f7 = _mm256_loadu_ps(pp + 56); + { + __m256 _tmp0 = _f0; + __m256 _tmp1 = _mm256_shuffle_ps(_f1, _f1, _MM_SHUFFLE(2, 1, 0, 3)); + __m256 _tmp2 = _f2; + __m256 _tmp3 = _mm256_shuffle_ps(_f3, _f3, _MM_SHUFFLE(2, 1, 0, 3)); + __m256 _tmp4 = _f4; + __m256 _tmp5 = _mm256_shuffle_ps(_f5, _f5, _MM_SHUFFLE(2, 1, 0, 3)); + __m256 _tmp6 = _f6; + __m256 _tmp7 = _mm256_shuffle_ps(_f7, _f7, _MM_SHUFFLE(2, 1, 0, 3)); + _f0 = _mm256_unpacklo_ps(_tmp0, _tmp3); + _f1 = _mm256_unpackhi_ps(_tmp0, _tmp3); + _f2 = _mm256_unpacklo_ps(_tmp2, _tmp1); + _f3 = _mm256_unpackhi_ps(_tmp2, _tmp1); + _f4 = _mm256_unpacklo_ps(_tmp4, _tmp7); + _f5 = _mm256_unpackhi_ps(_tmp4, _tmp7); + _f6 = _mm256_unpacklo_ps(_tmp6, _tmp5); + _f7 = _mm256_unpackhi_ps(_tmp6, _tmp5); + _tmp0 = _mm256_castpd_ps(_mm256_unpacklo_pd(_mm256_castps_pd(_f0), _mm256_castps_pd(_f2))); + _tmp1 = _mm256_castpd_ps(_mm256_unpackhi_pd(_mm256_castps_pd(_f0), _mm256_castps_pd(_f2))); + _tmp2 = _mm256_castpd_ps(_mm256_unpacklo_pd(_mm256_castps_pd(_f3), _mm256_castps_pd(_f1))); + _tmp3 = _mm256_castpd_ps(_mm256_unpackhi_pd(_mm256_castps_pd(_f3), _mm256_castps_pd(_f1))); + _tmp4 = _mm256_castpd_ps(_mm256_unpacklo_pd(_mm256_castps_pd(_f4), _mm256_castps_pd(_f6))); + _tmp5 = _mm256_castpd_ps(_mm256_unpackhi_pd(_mm256_castps_pd(_f4), _mm256_castps_pd(_f6))); + _tmp6 = _mm256_castpd_ps(_mm256_unpacklo_pd(_mm256_castps_pd(_f7), _mm256_castps_pd(_f5))); + _tmp7 = _mm256_castpd_ps(_mm256_unpackhi_pd(_mm256_castps_pd(_f7), _mm256_castps_pd(_f5))); + _tmp1 = _mm256_shuffle_ps(_tmp1, _tmp1, _MM_SHUFFLE(2, 1, 0, 3)); + _tmp3 = _mm256_shuffle_ps(_tmp3, _tmp3, _MM_SHUFFLE(2, 1, 0, 3)); + _tmp5 = _mm256_shuffle_ps(_tmp5, _tmp5, _MM_SHUFFLE(2, 1, 0, 3)); + _tmp7 = _mm256_shuffle_ps(_tmp7, _tmp7, _MM_SHUFFLE(2, 1, 0, 3)); + _f0 = _mm256_permute2f128_ps(_tmp0, _tmp4, _MM_SHUFFLE(0, 3, 0, 0)); + _f1 = _mm256_permute2f128_ps(_tmp1, _tmp5, _MM_SHUFFLE(0, 3, 0, 0)); + _f2 = _mm256_permute2f128_ps(_tmp2, _tmp6, _MM_SHUFFLE(0, 3, 0, 0)); + _f3 = _mm256_permute2f128_ps(_tmp3, _tmp7, _MM_SHUFFLE(0, 3, 0, 0)); + _f4 = _mm256_permute2f128_ps(_tmp4, _tmp0, _MM_SHUFFLE(0, 3, 0, 0)); + _f5 = _mm256_permute2f128_ps(_tmp5, _tmp1, _MM_SHUFFLE(0, 3, 0, 0)); + _f6 = _mm256_permute2f128_ps(_tmp6, _tmp2, _MM_SHUFFLE(0, 3, 0, 0)); + _f7 = _mm256_permute2f128_ps(_tmp7, _tmp3, _MM_SHUFFLE(0, 3, 0, 0)); + } + transpose8x8_ps(_f0, _f1, _f2, _f3, _f4, _f5, _f6, _f7); + if (pC) + { + if (broadcast_type_C == 0) + { + __m256 _c = _mm256_set1_ps(c0); + _f0 = _mm256_add_ps(_f0, _c); + _f1 = _mm256_add_ps(_f1, _c); + _f2 = _mm256_add_ps(_f2, _c); + _f3 = _mm256_add_ps(_f3, _c); + _f4 = _mm256_add_ps(_f4, _c); + _f5 = _mm256_add_ps(_f5, _c); + _f6 = _mm256_add_ps(_f6, _c); + _f7 = _mm256_add_ps(_f7, _c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = _mm256_add_ps(_f0, _mm256_set1_ps(c0)); + _f1 = _mm256_add_ps(_f1, _mm256_set1_ps(c1)); + _f2 = _mm256_add_ps(_f2, _mm256_set1_ps(c2)); + _f3 = _mm256_add_ps(_f3, _mm256_set1_ps(c3)); + _f4 = _mm256_add_ps(_f4, _mm256_set1_ps(c4)); + _f5 = _mm256_add_ps(_f5, _mm256_set1_ps(c5)); + _f6 = _mm256_add_ps(_f6, _mm256_set1_ps(c6)); + _f7 = _mm256_add_ps(_f7, _mm256_set1_ps(c7)); + } + if (broadcast_type_C == 3) + { + __m256 _c0 = _mm256_loadu_ps(pC); + __m256 _c1 = _mm256_loadu_ps(pC + N); + __m256 _c2 = _mm256_loadu_ps(pC + N * 2); + __m256 _c3 = _mm256_loadu_ps(pC + N * 3); + __m256 _c4 = _mm256_loadu_ps(pC + N * 4); + __m256 _c5 = _mm256_loadu_ps(pC + N * 5); + __m256 _c6 = _mm256_loadu_ps(pC + N * 6); + __m256 _c7 = _mm256_loadu_ps(pC + N * 7); + if (beta == 1.f) + { + _f0 = _mm256_add_ps(_f0, _c0); + _f1 = _mm256_add_ps(_f1, _c1); + _f2 = _mm256_add_ps(_f2, _c2); + _f3 = _mm256_add_ps(_f3, _c3); + _f4 = _mm256_add_ps(_f4, _c4); + _f5 = _mm256_add_ps(_f5, _c5); + _f6 = _mm256_add_ps(_f6, _c6); + _f7 = _mm256_add_ps(_f7, _c7); + } + else + { + __m256 _beta = _mm256_set1_ps(beta); +#if __FMA__ + _f0 = _mm256_fmadd_ps(_c0, _beta, _f0); + _f1 = _mm256_fmadd_ps(_c1, _beta, _f1); + _f2 = _mm256_fmadd_ps(_c2, _beta, _f2); + _f3 = _mm256_fmadd_ps(_c3, _beta, _f3); + _f4 = _mm256_fmadd_ps(_c4, _beta, _f4); + _f5 = _mm256_fmadd_ps(_c5, _beta, _f5); + _f6 = _mm256_fmadd_ps(_c6, _beta, _f6); + _f7 = _mm256_fmadd_ps(_c7, _beta, _f7); +#else + _f0 = _mm256_add_ps(_f0, _mm256_mul_ps(_c0, _beta)); + _f1 = _mm256_add_ps(_f1, _mm256_mul_ps(_c1, _beta)); + _f2 = _mm256_add_ps(_f2, _mm256_mul_ps(_c2, _beta)); + _f3 = _mm256_add_ps(_f3, _mm256_mul_ps(_c3, _beta)); + _f4 = _mm256_add_ps(_f4, _mm256_mul_ps(_c4, _beta)); + _f5 = _mm256_add_ps(_f5, _mm256_mul_ps(_c5, _beta)); + _f6 = _mm256_add_ps(_f6, _mm256_mul_ps(_c6, _beta)); + _f7 = _mm256_add_ps(_f7, _mm256_mul_ps(_c7, _beta)); +#endif + } + pC += 8; + } + if (broadcast_type_C == 4) + { + __m256 _c = _mm256_loadu_ps(pC); + if (beta != 1.f) + _c = _mm256_mul_ps(_c, _mm256_set1_ps(beta)); + _f0 = _mm256_add_ps(_f0, _c); + _f1 = _mm256_add_ps(_f1, _c); + _f2 = _mm256_add_ps(_f2, _c); + _f3 = _mm256_add_ps(_f3, _c); + _f4 = _mm256_add_ps(_f4, _c); + _f5 = _mm256_add_ps(_f5, _c); + _f6 = _mm256_add_ps(_f6, _c); + _f7 = _mm256_add_ps(_f7, _c); + pC += 8; + } + } + transpose8x8_ps(_f0, _f1, _f2, _f3, _f4, _f5, _f6, _f7); + if (alpha != 1.f) + { + __m256 _alpha = _mm256_set1_ps(alpha); + _f0 = _mm256_mul_ps(_f0, _alpha); + _f1 = _mm256_mul_ps(_f1, _alpha); + _f2 = _mm256_mul_ps(_f2, _alpha); + _f3 = _mm256_mul_ps(_f3, _alpha); + _f4 = _mm256_mul_ps(_f4, _alpha); + _f5 = _mm256_mul_ps(_f5, _alpha); + _f6 = _mm256_mul_ps(_f6, _alpha); + _f7 = _mm256_mul_ps(_f7, _alpha); + } + _mm256_storeu_ps(p0, _f0); + _mm256_storeu_ps(p0 + out_hstep, _f1); + _mm256_storeu_ps(p0 + out_hstep * 2, _f2); + _mm256_storeu_ps(p0 + out_hstep * 3, _f3); + _mm256_storeu_ps(p0 + out_hstep * 4, _f4); + _mm256_storeu_ps(p0 + out_hstep * 5, _f5); + _mm256_storeu_ps(p0 + out_hstep * 6, _f6); + _mm256_storeu_ps(p0 + out_hstep * 7, _f7); + p0 += out_hstep * 8; + pp += 64; + } +#endif // __AVX512F__ + for (; jj + 3 < max_jj; jj += 4) + { +#if __AVX2__ + __m256 _t0 = _mm256_loadu_ps(pp + 0); + __m256 _t1 = _mm256_loadu_ps(pp + 8); + __m256 _t2 = _mm256_loadu_ps(pp + 16); + __m256 _t3 = _mm256_loadu_ps(pp + 24); +#else + __m256 _t0 = combine4x2_ps(_mm_loadu_ps(pp + 0), _mm_loadu_ps(pp1 + 0)); + __m256 _t1 = combine4x2_ps(_mm_loadu_ps(pp + 4), _mm_loadu_ps(pp1 + 4)); + __m256 _t2 = combine4x2_ps(_mm_loadu_ps(pp + 8), _mm_loadu_ps(pp1 + 8)); + __m256 _t3 = combine4x2_ps(_mm_loadu_ps(pp + 12), _mm_loadu_ps(pp1 + 12)); +#endif + _t1 = _mm256_shuffle_ps(_t1, _t1, _MM_SHUFFLE(2, 1, 0, 3)); + _t3 = _mm256_shuffle_ps(_t3, _t3, _MM_SHUFFLE(2, 1, 0, 3)); + __m256 _tmp0 = _mm256_unpacklo_ps(_t0, _t3); + __m256 _tmp1 = _mm256_unpackhi_ps(_t0, _t3); + __m256 _tmp2 = _mm256_unpacklo_ps(_t2, _t1); + __m256 _tmp3 = _mm256_unpackhi_ps(_t2, _t1); + _t0 = _mm256_castpd_ps(_mm256_unpacklo_pd(_mm256_castps_pd(_tmp0), _mm256_castps_pd(_tmp2))); + _t1 = _mm256_castpd_ps(_mm256_unpackhi_pd(_mm256_castps_pd(_tmp0), _mm256_castps_pd(_tmp2))); + _t2 = _mm256_castpd_ps(_mm256_unpacklo_pd(_mm256_castps_pd(_tmp3), _mm256_castps_pd(_tmp1))); + _t3 = _mm256_castpd_ps(_mm256_unpackhi_pd(_mm256_castps_pd(_tmp3), _mm256_castps_pd(_tmp1))); + _t1 = _mm256_shuffle_ps(_t1, _t1, _MM_SHUFFLE(2, 1, 0, 3)); + _t3 = _mm256_shuffle_ps(_t3, _t3, _MM_SHUFFLE(2, 1, 0, 3)); + transpose8x4_ps(_t0, _t1, _t2, _t3); + __m128 _f0 = _mm256_extractf128_ps(_t0, 0); + __m128 _f1 = _mm256_extractf128_ps(_t0, 1); + __m128 _f2 = _mm256_extractf128_ps(_t1, 0); + __m128 _f3 = _mm256_extractf128_ps(_t1, 1); + __m128 _f4 = _mm256_extractf128_ps(_t2, 0); + __m128 _f5 = _mm256_extractf128_ps(_t2, 1); + __m128 _f6 = _mm256_extractf128_ps(_t3, 0); + __m128 _f7 = _mm256_extractf128_ps(_t3, 1); + if (pC) + { + if (broadcast_type_C == 0) + { + __m128 _c = _mm_set1_ps(c0); + _f0 = _mm_add_ps(_f0, _c); + _f1 = _mm_add_ps(_f1, _c); + _f2 = _mm_add_ps(_f2, _c); + _f3 = _mm_add_ps(_f3, _c); + _f4 = _mm_add_ps(_f4, _c); + _f5 = _mm_add_ps(_f5, _c); + _f6 = _mm_add_ps(_f6, _c); + _f7 = _mm_add_ps(_f7, _c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = _mm_add_ps(_f0, _mm_set1_ps(c0)); + _f1 = _mm_add_ps(_f1, _mm_set1_ps(c1)); + _f2 = _mm_add_ps(_f2, _mm_set1_ps(c2)); + _f3 = _mm_add_ps(_f3, _mm_set1_ps(c3)); + _f4 = _mm_add_ps(_f4, _mm_set1_ps(c4)); + _f5 = _mm_add_ps(_f5, _mm_set1_ps(c5)); + _f6 = _mm_add_ps(_f6, _mm_set1_ps(c6)); + _f7 = _mm_add_ps(_f7, _mm_set1_ps(c7)); + } + if (broadcast_type_C == 3) + { + __m128 _c0 = _mm_loadu_ps(pC); + __m128 _c1 = _mm_loadu_ps(pC + N); + __m128 _c2 = _mm_loadu_ps(pC + N * 2); + __m128 _c3 = _mm_loadu_ps(pC + N * 3); + __m128 _c4 = _mm_loadu_ps(pC + N * 4); + __m128 _c5 = _mm_loadu_ps(pC + N * 5); + __m128 _c6 = _mm_loadu_ps(pC + N * 6); + __m128 _c7 = _mm_loadu_ps(pC + N * 7); + if (beta == 1.f) + { + _f0 = _mm_add_ps(_f0, _c0); + _f1 = _mm_add_ps(_f1, _c1); + _f2 = _mm_add_ps(_f2, _c2); + _f3 = _mm_add_ps(_f3, _c3); + _f4 = _mm_add_ps(_f4, _c4); + _f5 = _mm_add_ps(_f5, _c5); + _f6 = _mm_add_ps(_f6, _c6); + _f7 = _mm_add_ps(_f7, _c7); + } + else + { + __m128 _beta = _mm_set1_ps(beta); +#if __FMA__ + _f0 = _mm_fmadd_ps(_c0, _beta, _f0); + _f1 = _mm_fmadd_ps(_c1, _beta, _f1); + _f2 = _mm_fmadd_ps(_c2, _beta, _f2); + _f3 = _mm_fmadd_ps(_c3, _beta, _f3); + _f4 = _mm_fmadd_ps(_c4, _beta, _f4); + _f5 = _mm_fmadd_ps(_c5, _beta, _f5); + _f6 = _mm_fmadd_ps(_c6, _beta, _f6); + _f7 = _mm_fmadd_ps(_c7, _beta, _f7); +#else + _f0 = _mm_add_ps(_f0, _mm_mul_ps(_c0, _beta)); + _f1 = _mm_add_ps(_f1, _mm_mul_ps(_c1, _beta)); + _f2 = _mm_add_ps(_f2, _mm_mul_ps(_c2, _beta)); + _f3 = _mm_add_ps(_f3, _mm_mul_ps(_c3, _beta)); + _f4 = _mm_add_ps(_f4, _mm_mul_ps(_c4, _beta)); + _f5 = _mm_add_ps(_f5, _mm_mul_ps(_c5, _beta)); + _f6 = _mm_add_ps(_f6, _mm_mul_ps(_c6, _beta)); + _f7 = _mm_add_ps(_f7, _mm_mul_ps(_c7, _beta)); +#endif + } + pC += 4; + } + if (broadcast_type_C == 4) + { + __m128 _c = _mm_loadu_ps(pC); + if (beta != 1.f) + _c = _mm_mul_ps(_c, _mm_set1_ps(beta)); + _f0 = _mm_add_ps(_f0, _c); + _f1 = _mm_add_ps(_f1, _c); + _f2 = _mm_add_ps(_f2, _c); + _f3 = _mm_add_ps(_f3, _c); + _f4 = _mm_add_ps(_f4, _c); + _f5 = _mm_add_ps(_f5, _c); + _f6 = _mm_add_ps(_f6, _c); + _f7 = _mm_add_ps(_f7, _c); + pC += 4; + } + } + if (alpha != 1.f) + { + __m128 _alpha = _mm_set1_ps(alpha); + _f0 = _mm_mul_ps(_f0, _alpha); + _f1 = _mm_mul_ps(_f1, _alpha); + _f2 = _mm_mul_ps(_f2, _alpha); + _f3 = _mm_mul_ps(_f3, _alpha); + _f4 = _mm_mul_ps(_f4, _alpha); + _f5 = _mm_mul_ps(_f5, _alpha); + _f6 = _mm_mul_ps(_f6, _alpha); + _f7 = _mm_mul_ps(_f7, _alpha); + } + { + __m128 _r0 = _f0; + __m128 _r1 = _f1; + __m128 _r2 = _f2; + __m128 _r3 = _f3; + _MM_TRANSPOSE4_PS(_r0, _r1, _r2, _r3); + __m128 _s0 = _f4; + __m128 _s1 = _f5; + __m128 _s2 = _f6; + __m128 _s3 = _f7; + _MM_TRANSPOSE4_PS(_s0, _s1, _s2, _s3); + _mm256_storeu_ps(p0, _mm256_insertf128_ps(_mm256_castps128_ps256(_r0), _s0, 1)); + _mm256_storeu_ps(p0 + out_hstep, _mm256_insertf128_ps(_mm256_castps128_ps256(_r1), _s1, 1)); + _mm256_storeu_ps(p0 + out_hstep * 2, _mm256_insertf128_ps(_mm256_castps128_ps256(_r2), _s2, 1)); + _mm256_storeu_ps(p0 + out_hstep * 3, _mm256_insertf128_ps(_mm256_castps128_ps256(_r3), _s3, 1)); + } + p0 += out_hstep * 4; +#if __AVX2__ + pp += 32; +#else + pp += 16; + pp1 += 16; +#endif + } +#endif // defined(__x86_64__) || defined(_M_X64) + for (; jj + 1 < max_jj; jj += 2) + { +#if __AVX2__ + __m256 _t0 = _mm256_loadu_ps(pp); + __m256 _t1 = _mm256_loadu_ps(pp + 8); +#else + __m256 _t0 = combine4x2_ps(_mm_loadu_ps(pp), _mm_loadu_ps(pp1)); + __m256 _t1 = combine4x2_ps(_mm_loadu_ps(pp + 4), _mm_loadu_ps(pp1 + 4)); +#endif + __m256 _tmp0 = _mm256_shuffle_ps(_t0, _t0, _MM_SHUFFLE(3, 1, 2, 0)); + __m256 _tmp1 = _mm256_shuffle_ps(_t1, _t1, _MM_SHUFFLE(0, 2, 3, 1)); + _t0 = _mm256_unpacklo_ps(_tmp0, _tmp1); + _t1 = _mm256_unpackhi_ps(_tmp0, _tmp1); + _t1 = _mm256_shuffle_ps(_t1, _t1, _MM_SHUFFLE(2, 1, 0, 3)); + transpose8x2_ps(_t0, _t1); + __m128 _r01 = _mm256_extractf128_ps(_t0, 0); + __m128 _r23 = _mm256_extractf128_ps(_t0, 1); + __m128 _r45 = _mm256_extractf128_ps(_t1, 0); + __m128 _r67 = _mm256_extractf128_ps(_t1, 1); + __m128 _f0 = _r01; + __m128 _f1 = _mm_movehl_ps(_r01, _r01); + __m128 _f2 = _r23; + __m128 _f3 = _mm_movehl_ps(_r23, _r23); + __m128 _f4 = _r45; + __m128 _f5 = _mm_movehl_ps(_r45, _r45); + __m128 _f6 = _r67; + __m128 _f7 = _mm_movehl_ps(_r67, _r67); + if (pC) + { + if (broadcast_type_C == 0) + { + __m128 _c = _mm_set1_ps(c0); + _f0 = _mm_add_ps(_f0, _c); + _f1 = _mm_add_ps(_f1, _c); + _f2 = _mm_add_ps(_f2, _c); + _f3 = _mm_add_ps(_f3, _c); + _f4 = _mm_add_ps(_f4, _c); + _f5 = _mm_add_ps(_f5, _c); + _f6 = _mm_add_ps(_f6, _c); + _f7 = _mm_add_ps(_f7, _c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = _mm_add_ps(_f0, _mm_set1_ps(c0)); + _f1 = _mm_add_ps(_f1, _mm_set1_ps(c1)); + _f2 = _mm_add_ps(_f2, _mm_set1_ps(c2)); + _f3 = _mm_add_ps(_f3, _mm_set1_ps(c3)); + _f4 = _mm_add_ps(_f4, _mm_set1_ps(c4)); + _f5 = _mm_add_ps(_f5, _mm_set1_ps(c5)); + _f6 = _mm_add_ps(_f6, _mm_set1_ps(c6)); + _f7 = _mm_add_ps(_f7, _mm_set1_ps(c7)); + } + if (broadcast_type_C == 3) + { + __m128 _c0 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pC); + __m128 _c1 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N)); + __m128 _c2 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N * 2)); + __m128 _c3 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N * 3)); + __m128 _c4 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N * 4)); + __m128 _c5 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N * 5)); + __m128 _c6 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N * 6)); + __m128 _c7 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N * 7)); + if (beta == 1.f) + { + _f0 = _mm_add_ps(_f0, _c0); + _f1 = _mm_add_ps(_f1, _c1); + _f2 = _mm_add_ps(_f2, _c2); + _f3 = _mm_add_ps(_f3, _c3); + _f4 = _mm_add_ps(_f4, _c4); + _f5 = _mm_add_ps(_f5, _c5); + _f6 = _mm_add_ps(_f6, _c6); + _f7 = _mm_add_ps(_f7, _c7); + } + else + { + __m128 _beta = _mm_set1_ps(beta); +#if __FMA__ + _f0 = _mm_fmadd_ps(_c0, _beta, _f0); + _f1 = _mm_fmadd_ps(_c1, _beta, _f1); + _f2 = _mm_fmadd_ps(_c2, _beta, _f2); + _f3 = _mm_fmadd_ps(_c3, _beta, _f3); + _f4 = _mm_fmadd_ps(_c4, _beta, _f4); + _f5 = _mm_fmadd_ps(_c5, _beta, _f5); + _f6 = _mm_fmadd_ps(_c6, _beta, _f6); + _f7 = _mm_fmadd_ps(_c7, _beta, _f7); +#else + _f0 = _mm_add_ps(_f0, _mm_mul_ps(_c0, _beta)); + _f1 = _mm_add_ps(_f1, _mm_mul_ps(_c1, _beta)); + _f2 = _mm_add_ps(_f2, _mm_mul_ps(_c2, _beta)); + _f3 = _mm_add_ps(_f3, _mm_mul_ps(_c3, _beta)); + _f4 = _mm_add_ps(_f4, _mm_mul_ps(_c4, _beta)); + _f5 = _mm_add_ps(_f5, _mm_mul_ps(_c5, _beta)); + _f6 = _mm_add_ps(_f6, _mm_mul_ps(_c6, _beta)); + _f7 = _mm_add_ps(_f7, _mm_mul_ps(_c7, _beta)); +#endif + } + pC += 2; + } + if (broadcast_type_C == 4) + { + __m128 _c = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pC); + if (beta != 1.f) + _c = _mm_mul_ps(_c, _mm_set1_ps(beta)); + _f0 = _mm_add_ps(_f0, _c); + _f1 = _mm_add_ps(_f1, _c); + _f2 = _mm_add_ps(_f2, _c); + _f3 = _mm_add_ps(_f3, _c); + _f4 = _mm_add_ps(_f4, _c); + _f5 = _mm_add_ps(_f5, _c); + _f6 = _mm_add_ps(_f6, _c); + _f7 = _mm_add_ps(_f7, _c); + pC += 2; + } + } + if (alpha != 1.f) + { + __m128 _alpha = _mm_set1_ps(alpha); + _f0 = _mm_mul_ps(_f0, _alpha); + _f1 = _mm_mul_ps(_f1, _alpha); + _f2 = _mm_mul_ps(_f2, _alpha); + _f3 = _mm_mul_ps(_f3, _alpha); + _f4 = _mm_mul_ps(_f4, _alpha); + _f5 = _mm_mul_ps(_f5, _alpha); + _f6 = _mm_mul_ps(_f6, _alpha); + _f7 = _mm_mul_ps(_f7, _alpha); + } + { + __m128 _r0 = _f0; + __m128 _r1 = _f1; + __m128 _r2 = _f2; + __m128 _r3 = _f3; + _MM_TRANSPOSE4_PS(_r0, _r1, _r2, _r3); + __m128 _s0 = _f4; + __m128 _s1 = _f5; + __m128 _s2 = _f6; + __m128 _s3 = _f7; + _MM_TRANSPOSE4_PS(_s0, _s1, _s2, _s3); + _mm256_storeu_ps(p0, _mm256_insertf128_ps(_mm256_castps128_ps256(_r0), _s0, 1)); + _mm256_storeu_ps(p0 + out_hstep, _mm256_insertf128_ps(_mm256_castps128_ps256(_r1), _s1, 1)); + } + p0 += out_hstep * 2; +#if __AVX2__ + pp += 16; +#else + pp += 8; + pp1 += 8; +#endif + } + for (; jj < max_jj; jj++) + { +#if __AVX2__ + __m128 _f03 = _mm_loadu_ps(pp); + __m128 _f47 = _mm_loadu_ps(pp + 4); +#else + __m128 _f03 = _mm_loadu_ps(pp); + __m128 _f47 = _mm_loadu_ps(pp1); +#endif + if (pC) + { + if (broadcast_type_C == 0) + { + __m128 _c = _mm_set1_ps(c0); + _f03 = _mm_add_ps(_f03, _c); + _f47 = _mm_add_ps(_f47, _c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + __m128 _c03 = _mm_setr_ps(c0, c1, c2, c3); + __m128 _c47 = _mm_setr_ps(c4, c5, c6, c7); + _f03 = _mm_add_ps(_f03, _c03); + _f47 = _mm_add_ps(_f47, _c47); + } + if (broadcast_type_C == 3 || broadcast_type_C == 4) + { + __m128 _c03; + __m128 _c47; + if (broadcast_type_C == 3) + { + _c03 = _mm_setr_ps(pC[0], pC[N], pC[N * 2], pC[N * 3]); + _c47 = _mm_setr_ps(pC[N * 4], pC[N * 5], pC[N * 6], pC[N * 7]); + } + if (broadcast_type_C == 4) + { + _c03 = _mm_set1_ps(pC[0]); + _c47 = _c03; + } + if (beta == 1.f) + { + _f03 = _mm_add_ps(_f03, _c03); + _f47 = _mm_add_ps(_f47, _c47); + } + else + { +#if __FMA__ + __m128 _beta = _mm_set1_ps(beta); + _f03 = _mm_fmadd_ps(_c03, _beta, _f03); + _f47 = _mm_fmadd_ps(_c47, _beta, _f47); +#else + _f03 = _mm_add_ps(_f03, _mm_mul_ps(_c03, _mm_set1_ps(beta))); + _f47 = _mm_add_ps(_f47, _mm_mul_ps(_c47, _mm_set1_ps(beta))); +#endif + } + pC++; + } + } + if (alpha != 1.f) + { + __m128 _alpha = _mm_set1_ps(alpha); + _f03 = _mm_mul_ps(_f03, _alpha); + _f47 = _mm_mul_ps(_f47, _alpha); + } + _mm256_storeu_ps(p0, _mm256_insertf128_ps(_mm256_castps128_ps256(_f03), _f47, 1)); + p0 += out_hstep; +#if __AVX2__ + pp += 8; +#else + pp += 4; + pp1 += 4; +#endif + } +#if !__AVX2__ + pp = pp1; + pp1 = pp + max_jj * 4; +#endif + } +#endif // __AVX__ + for (; ii + 3 < max_ii; ii += 4) + { + float* p0 = (float*)top_blob + j * out_hstep + i + ii; + + float c0 = 0.f; + float c1 = c0; + float c2 = c0; + float c3 = c0; + if (pC) + { + if (broadcast_type_C == 0) + { + c0 = pC[0]; + c1 = c0; + c2 = c0; + c3 = c0; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + c0 = pC[i + ii]; + c1 = pC[i + ii + 1]; + c2 = pC[i + ii + 2]; + c3 = pC[i + ii + 3]; + } + if (broadcast_type_C == 3) + pC = (const float*)C + (i + ii) * N + j; + if (broadcast_type_C == 4) + pC = (const float*)C + j; + if (broadcast_type_C == 0 && beta != 1.f) + c0 *= beta; + if ((broadcast_type_C == 1 || broadcast_type_C == 2) && beta != 1.f) + { + c0 *= beta; + c1 *= beta; + c2 *= beta; + c3 *= beta; + } + } + + int jj = 0; +#if defined(__x86_64__) || defined(_M_X64) +#if __AVX512F__ + for (; jj + 7 < max_jj; jj += 8) + { + __m256 _f0 = _mm256_loadu_ps(pp + 0); + __m256 _f1 = _mm256_loadu_ps(pp + 8); + __m256 _f2 = _mm256_loadu_ps(pp + 16); + __m256 _f3 = _mm256_loadu_ps(pp + 24); + __m128 _f00 = _mm256_castps256_ps128(_f0); + __m128 _f01 = _mm256_castps256_ps128(_f1); + __m128 _f02 = _mm256_castps256_ps128(_f2); + __m128 _f03 = _mm256_castps256_ps128(_f3); + __m128 _f10 = _mm256_extractf128_ps(_f0, 1); + __m128 _f11 = _mm256_extractf128_ps(_f1, 1); + __m128 _f12 = _mm256_extractf128_ps(_f2, 1); + __m128 _f13 = _mm256_extractf128_ps(_f3, 1); + { + _f01 = _mm_shuffle_ps(_f01, _f01, _MM_SHUFFLE(2, 1, 0, 3)); + _f03 = _mm_shuffle_ps(_f03, _f03, _MM_SHUFFLE(2, 1, 0, 3)); + __m128 _tmp0 = _mm_unpacklo_ps(_f00, _f03); + __m128 _tmp1 = _mm_unpackhi_ps(_f00, _f03); + __m128 _tmp2 = _mm_unpacklo_ps(_f02, _f01); + __m128 _tmp3 = _mm_unpackhi_ps(_f02, _f01); + _f00 = _mm_castpd_ps(_mm_unpacklo_pd(_mm_castps_pd(_tmp0), _mm_castps_pd(_tmp2))); + _f01 = _mm_castpd_ps(_mm_unpackhi_pd(_mm_castps_pd(_tmp0), _mm_castps_pd(_tmp2))); + _f02 = _mm_castpd_ps(_mm_unpacklo_pd(_mm_castps_pd(_tmp3), _mm_castps_pd(_tmp1))); + _f03 = _mm_castpd_ps(_mm_unpackhi_pd(_mm_castps_pd(_tmp3), _mm_castps_pd(_tmp1))); + _f01 = _mm_shuffle_ps(_f01, _f01, _MM_SHUFFLE(2, 1, 0, 3)); + _f03 = _mm_shuffle_ps(_f03, _f03, _MM_SHUFFLE(2, 1, 0, 3)); + _MM_TRANSPOSE4_PS(_f00, _f01, _f02, _f03); + } + { + _f11 = _mm_shuffle_ps(_f11, _f11, _MM_SHUFFLE(2, 1, 0, 3)); + _f13 = _mm_shuffle_ps(_f13, _f13, _MM_SHUFFLE(2, 1, 0, 3)); + __m128 _tmp0 = _mm_unpacklo_ps(_f10, _f13); + __m128 _tmp1 = _mm_unpackhi_ps(_f10, _f13); + __m128 _tmp2 = _mm_unpacklo_ps(_f12, _f11); + __m128 _tmp3 = _mm_unpackhi_ps(_f12, _f11); + _f10 = _mm_castpd_ps(_mm_unpacklo_pd(_mm_castps_pd(_tmp0), _mm_castps_pd(_tmp2))); + _f11 = _mm_castpd_ps(_mm_unpackhi_pd(_mm_castps_pd(_tmp0), _mm_castps_pd(_tmp2))); + _f12 = _mm_castpd_ps(_mm_unpacklo_pd(_mm_castps_pd(_tmp3), _mm_castps_pd(_tmp1))); + _f13 = _mm_castpd_ps(_mm_unpackhi_pd(_mm_castps_pd(_tmp3), _mm_castps_pd(_tmp1))); + _f11 = _mm_shuffle_ps(_f11, _f11, _MM_SHUFFLE(2, 1, 0, 3)); + _f13 = _mm_shuffle_ps(_f13, _f13, _MM_SHUFFLE(2, 1, 0, 3)); + _MM_TRANSPOSE4_PS(_f10, _f11, _f12, _f13); + } + _f0 = combine4x2_ps(_f00, _f10); + _f1 = combine4x2_ps(_f01, _f11); + _f2 = combine4x2_ps(_f02, _f12); + _f3 = combine4x2_ps(_f03, _f13); + if (pC) + { + if (broadcast_type_C == 0) + { + __m256 _c = _mm256_set1_ps(c0); + _f0 = _mm256_add_ps(_f0, _c); + _f1 = _mm256_add_ps(_f1, _c); + _f2 = _mm256_add_ps(_f2, _c); + _f3 = _mm256_add_ps(_f3, _c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = _mm256_add_ps(_f0, _mm256_set1_ps(c0)); + _f1 = _mm256_add_ps(_f1, _mm256_set1_ps(c1)); + _f2 = _mm256_add_ps(_f2, _mm256_set1_ps(c2)); + _f3 = _mm256_add_ps(_f3, _mm256_set1_ps(c3)); + } + if (broadcast_type_C == 3) + { + __m256 _c0 = _mm256_loadu_ps(pC); + __m256 _c1 = _mm256_loadu_ps(pC + N); + __m256 _c2 = _mm256_loadu_ps(pC + N * 2); + __m256 _c3 = _mm256_loadu_ps(pC + N * 3); + if (beta == 1.f) + { + _f0 = _mm256_add_ps(_f0, _c0); + _f1 = _mm256_add_ps(_f1, _c1); + _f2 = _mm256_add_ps(_f2, _c2); + _f3 = _mm256_add_ps(_f3, _c3); + } + else + { + __m256 _beta = _mm256_set1_ps(beta); +#if __FMA__ + _f0 = _mm256_fmadd_ps(_c0, _beta, _f0); + _f1 = _mm256_fmadd_ps(_c1, _beta, _f1); + _f2 = _mm256_fmadd_ps(_c2, _beta, _f2); + _f3 = _mm256_fmadd_ps(_c3, _beta, _f3); +#else + _f0 = _mm256_add_ps(_f0, _mm256_mul_ps(_c0, _beta)); + _f1 = _mm256_add_ps(_f1, _mm256_mul_ps(_c1, _beta)); + _f2 = _mm256_add_ps(_f2, _mm256_mul_ps(_c2, _beta)); + _f3 = _mm256_add_ps(_f3, _mm256_mul_ps(_c3, _beta)); +#endif + } + pC += 8; + } + if (broadcast_type_C == 4) + { + __m256 _c = _mm256_loadu_ps(pC); + if (beta != 1.f) + _c = _mm256_mul_ps(_c, _mm256_set1_ps(beta)); + _f0 = _mm256_add_ps(_f0, _c); + _f1 = _mm256_add_ps(_f1, _c); + _f2 = _mm256_add_ps(_f2, _c); + _f3 = _mm256_add_ps(_f3, _c); + pC += 8; + } + } + if (alpha != 1.f) + { + __m256 _alpha = _mm256_set1_ps(alpha); + _f0 = _mm256_mul_ps(_f0, _alpha); + _f1 = _mm256_mul_ps(_f1, _alpha); + _f2 = _mm256_mul_ps(_f2, _alpha); + _f3 = _mm256_mul_ps(_f3, _alpha); + } + { + __m128 _r0 = _mm256_castps256_ps128(_f0); + __m128 _r1 = _mm256_castps256_ps128(_f1); + __m128 _r2 = _mm256_castps256_ps128(_f2); + __m128 _r3 = _mm256_castps256_ps128(_f3); + _MM_TRANSPOSE4_PS(_r0, _r1, _r2, _r3); + _mm_storeu_ps(p0, _r0); + _mm_storeu_ps(p0 + out_hstep, _r1); + _mm_storeu_ps(p0 + out_hstep * 2, _r2); + _mm_storeu_ps(p0 + out_hstep * 3, _r3); + } + { + __m128 _r0 = _mm256_extractf128_ps(_f0, 1); + __m128 _r1 = _mm256_extractf128_ps(_f1, 1); + __m128 _r2 = _mm256_extractf128_ps(_f2, 1); + __m128 _r3 = _mm256_extractf128_ps(_f3, 1); + _MM_TRANSPOSE4_PS(_r0, _r1, _r2, _r3); + _mm_storeu_ps(p0 + out_hstep * 4, _r0); + _mm_storeu_ps(p0 + out_hstep * 5, _r1); + _mm_storeu_ps(p0 + out_hstep * 6, _r2); + _mm_storeu_ps(p0 + out_hstep * 7, _r3); + } + p0 += out_hstep * 8; + pp += 32; + } +#endif // __AVX512F__ + for (; jj + 3 < max_jj; jj += 4) + { + __m128 _f0 = _mm_loadu_ps(pp + 0); + __m128 _f1 = _mm_loadu_ps(pp + 4); + __m128 _f2 = _mm_loadu_ps(pp + 8); + __m128 _f3 = _mm_loadu_ps(pp + 12); + { + _f1 = _mm_shuffle_ps(_f1, _f1, _MM_SHUFFLE(2, 1, 0, 3)); + _f3 = _mm_shuffle_ps(_f3, _f3, _MM_SHUFFLE(2, 1, 0, 3)); + __m128 _tmp0 = _mm_unpacklo_ps(_f0, _f3); + __m128 _tmp1 = _mm_unpackhi_ps(_f0, _f3); + __m128 _tmp2 = _mm_unpacklo_ps(_f2, _f1); + __m128 _tmp3 = _mm_unpackhi_ps(_f2, _f1); + _f0 = _mm_castpd_ps(_mm_unpacklo_pd(_mm_castps_pd(_tmp0), _mm_castps_pd(_tmp2))); + _f1 = _mm_castpd_ps(_mm_unpackhi_pd(_mm_castps_pd(_tmp0), _mm_castps_pd(_tmp2))); + _f2 = _mm_castpd_ps(_mm_unpacklo_pd(_mm_castps_pd(_tmp3), _mm_castps_pd(_tmp1))); + _f3 = _mm_castpd_ps(_mm_unpackhi_pd(_mm_castps_pd(_tmp3), _mm_castps_pd(_tmp1))); + _f1 = _mm_shuffle_ps(_f1, _f1, _MM_SHUFFLE(2, 1, 0, 3)); + _f3 = _mm_shuffle_ps(_f3, _f3, _MM_SHUFFLE(2, 1, 0, 3)); + _MM_TRANSPOSE4_PS(_f0, _f1, _f2, _f3); + } + if (pC) + { + if (broadcast_type_C == 0) + { + __m128 _c = _mm_set1_ps(c0); + _f0 = _mm_add_ps(_f0, _c); + _f1 = _mm_add_ps(_f1, _c); + _f2 = _mm_add_ps(_f2, _c); + _f3 = _mm_add_ps(_f3, _c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = _mm_add_ps(_f0, _mm_set1_ps(c0)); + _f1 = _mm_add_ps(_f1, _mm_set1_ps(c1)); + _f2 = _mm_add_ps(_f2, _mm_set1_ps(c2)); + _f3 = _mm_add_ps(_f3, _mm_set1_ps(c3)); + } + if (broadcast_type_C == 3) + { + __m128 _c0 = _mm_loadu_ps(pC); + __m128 _c1 = _mm_loadu_ps(pC + N); + __m128 _c2 = _mm_loadu_ps(pC + N * 2); + __m128 _c3 = _mm_loadu_ps(pC + N * 3); + if (beta == 1.f) + { + _f0 = _mm_add_ps(_f0, _c0); + _f1 = _mm_add_ps(_f1, _c1); + _f2 = _mm_add_ps(_f2, _c2); + _f3 = _mm_add_ps(_f3, _c3); + } + else + { + __m128 _beta = _mm_set1_ps(beta); +#if __FMA__ + _f0 = _mm_fmadd_ps(_c0, _beta, _f0); + _f1 = _mm_fmadd_ps(_c1, _beta, _f1); + _f2 = _mm_fmadd_ps(_c2, _beta, _f2); + _f3 = _mm_fmadd_ps(_c3, _beta, _f3); +#else + _f0 = _mm_add_ps(_f0, _mm_mul_ps(_c0, _beta)); + _f1 = _mm_add_ps(_f1, _mm_mul_ps(_c1, _beta)); + _f2 = _mm_add_ps(_f2, _mm_mul_ps(_c2, _beta)); + _f3 = _mm_add_ps(_f3, _mm_mul_ps(_c3, _beta)); +#endif + } + pC += 4; + } + if (broadcast_type_C == 4) + { + __m128 _c = _mm_loadu_ps(pC); + if (beta != 1.f) + _c = _mm_mul_ps(_c, _mm_set1_ps(beta)); + _f0 = _mm_add_ps(_f0, _c); + _f1 = _mm_add_ps(_f1, _c); + _f2 = _mm_add_ps(_f2, _c); + _f3 = _mm_add_ps(_f3, _c); + pC += 4; + } + } + if (alpha != 1.f) + { + __m128 _alpha = _mm_set1_ps(alpha); + _f0 = _mm_mul_ps(_f0, _alpha); + _f1 = _mm_mul_ps(_f1, _alpha); + _f2 = _mm_mul_ps(_f2, _alpha); + _f3 = _mm_mul_ps(_f3, _alpha); + } + { + __m128 _r0 = _f0; + __m128 _r1 = _f1; + __m128 _r2 = _f2; + __m128 _r3 = _f3; + _MM_TRANSPOSE4_PS(_r0, _r1, _r2, _r3); + _mm_storeu_ps(p0, _r0); + _mm_storeu_ps(p0 + out_hstep, _r1); + _mm_storeu_ps(p0 + out_hstep * 2, _r2); + _mm_storeu_ps(p0 + out_hstep * 3, _r3); + } + p0 += out_hstep * 4; + pp += 16; + } +#endif // defined(__x86_64__) || defined(_M_X64) + for (; jj + 1 < max_jj; jj += 2) + { + __m128 _t0 = _mm_loadu_ps(pp); + __m128 _t1 = _mm_loadu_ps(pp + 4); + __m128 _tmp0 = _mm_shuffle_ps(_t0, _t0, _MM_SHUFFLE(3, 1, 2, 0)); + __m128 _tmp1 = _mm_shuffle_ps(_t1, _t1, _MM_SHUFFLE(0, 2, 3, 1)); + __m128 _c0v = _mm_unpacklo_ps(_tmp0, _tmp1); + __m128 _c1v = _mm_unpackhi_ps(_tmp0, _tmp1); + _c1v = _mm_shuffle_ps(_c1v, _c1v, _MM_SHUFFLE(2, 1, 0, 3)); + __m128 _f01 = _mm_unpacklo_ps(_c0v, _c1v); + __m128 _f23 = _mm_unpackhi_ps(_c0v, _c1v); + __m128 _f0 = _f01; + __m128 _f1 = _mm_movehl_ps(_f01, _f01); + __m128 _f2 = _f23; + __m128 _f3 = _mm_movehl_ps(_f23, _f23); + if (pC) + { + if (broadcast_type_C == 0) + { + __m128 _c = _mm_set1_ps(c0); + _f0 = _mm_add_ps(_f0, _c); + _f1 = _mm_add_ps(_f1, _c); + _f2 = _mm_add_ps(_f2, _c); + _f3 = _mm_add_ps(_f3, _c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = _mm_add_ps(_f0, _mm_set1_ps(c0)); + _f1 = _mm_add_ps(_f1, _mm_set1_ps(c1)); + _f2 = _mm_add_ps(_f2, _mm_set1_ps(c2)); + _f3 = _mm_add_ps(_f3, _mm_set1_ps(c3)); + } + if (broadcast_type_C == 3) + { + __m128 _c0 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pC); + __m128 _c1 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N)); + __m128 _c2 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N * 2)); + __m128 _c3 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N * 3)); + if (beta == 1.f) + { + _f0 = _mm_add_ps(_f0, _c0); + _f1 = _mm_add_ps(_f1, _c1); + _f2 = _mm_add_ps(_f2, _c2); + _f3 = _mm_add_ps(_f3, _c3); + } + else + { + __m128 _beta = _mm_set1_ps(beta); +#if __FMA__ + _f0 = _mm_fmadd_ps(_c0, _beta, _f0); + _f1 = _mm_fmadd_ps(_c1, _beta, _f1); + _f2 = _mm_fmadd_ps(_c2, _beta, _f2); + _f3 = _mm_fmadd_ps(_c3, _beta, _f3); +#else + _f0 = _mm_add_ps(_f0, _mm_mul_ps(_c0, _beta)); + _f1 = _mm_add_ps(_f1, _mm_mul_ps(_c1, _beta)); + _f2 = _mm_add_ps(_f2, _mm_mul_ps(_c2, _beta)); + _f3 = _mm_add_ps(_f3, _mm_mul_ps(_c3, _beta)); +#endif + } + pC += 2; + } + if (broadcast_type_C == 4) + { + __m128 _c = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pC); + if (beta != 1.f) + _c = _mm_mul_ps(_c, _mm_set1_ps(beta)); + _f0 = _mm_add_ps(_f0, _c); + _f1 = _mm_add_ps(_f1, _c); + _f2 = _mm_add_ps(_f2, _c); + _f3 = _mm_add_ps(_f3, _c); + pC += 2; + } + } + if (alpha != 1.f) + { + __m128 _alpha = _mm_set1_ps(alpha); + _f0 = _mm_mul_ps(_f0, _alpha); + _f1 = _mm_mul_ps(_f1, _alpha); + _f2 = _mm_mul_ps(_f2, _alpha); + _f3 = _mm_mul_ps(_f3, _alpha); + } + { + __m128 _r0 = _f0; + __m128 _r1 = _f1; + __m128 _r2 = _f2; + __m128 _r3 = _f3; + _MM_TRANSPOSE4_PS(_r0, _r1, _r2, _r3); + _mm_storeu_ps(p0, _r0); + _mm_storeu_ps(p0 + out_hstep, _r1); + } + p0 += out_hstep * 2; + pp += 8; + } + for (; jj < max_jj; jj += 1) + { + __m128 _f = _mm_loadu_ps(pp); + if (pC) + { + if (broadcast_type_C == 0) + _f = _mm_add_ps(_f, _mm_set1_ps(c0)); + if (broadcast_type_C == 1 || broadcast_type_C == 2) + _f = _mm_add_ps(_f, _mm_setr_ps(c0, c1, c2, c3)); + if (broadcast_type_C == 3 || broadcast_type_C == 4) + { + __m128 _c; + if (broadcast_type_C == 3) + _c = _mm_setr_ps(pC[0], pC[N], pC[N * 2], pC[N * 3]); + if (broadcast_type_C == 4) + _c = _mm_set1_ps(pC[0]); + if (beta == 1.f) + { + _f = _mm_add_ps(_f, _c); + } + else + { +#if __FMA__ + _f = _mm_fmadd_ps(_c, _mm_set1_ps(beta), _f); +#else + _f = _mm_add_ps(_f, _mm_mul_ps(_c, _mm_set1_ps(beta))); +#endif + } + pC++; + } + } + if (alpha != 1.f) + _f = _mm_mul_ps(_f, _mm_set1_ps(alpha)); + _mm_storeu_ps(p0, _f); + p0 += out_hstep; + pp += 4; + } + } + +#endif // __SSE2__ + + for (; ii + 1 < max_ii; ii += 2) + { + float* p0 = (float*)top_blob + j * out_hstep + i + ii; + + float c0 = 0.f; + float c1 = c0; + if (pC) + { + if (broadcast_type_C == 0) + { + c0 = pC[0]; + c1 = c0; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + c0 = pC[i + ii]; + c1 = pC[i + ii + 1]; + } + if (broadcast_type_C == 3) + pC = (const float*)C + (i + ii) * N + j; + if (broadcast_type_C == 4) + pC = (const float*)C + j; + if (broadcast_type_C == 0 && beta != 1.f) + c0 *= beta; + if ((broadcast_type_C == 1 || broadcast_type_C == 2) && beta != 1.f) + { + c0 *= beta; + c1 *= beta; + } + } + + int jj = 0; +#if __SSE2__ +#if defined(__x86_64__) || defined(_M_X64) +#if __AVX512F__ + for (; jj + 7 < max_jj; jj += 8) + { + __m128 _t0 = _mm_loadu_ps(pp); + __m128 _t1 = _mm_loadu_ps(pp + 4); + __m128 _t2 = _mm_loadu_ps(pp + 8); + __m128 _t3 = _mm_loadu_ps(pp + 12); + _t2 = _mm_shuffle_ps(_t2, _t2, _MM_SHUFFLE(2, 3, 0, 1)); + _t3 = _mm_shuffle_ps(_t3, _t3, _MM_SHUFFLE(2, 3, 0, 1)); + __m128 _tmp0x = _mm_unpacklo_ps(_t0, _t2); + __m128 _tmp1x = _mm_unpackhi_ps(_t0, _t2); + __m128 _tmp2x = _mm_unpacklo_ps(_t1, _t3); + __m128 _tmp3x = _mm_unpackhi_ps(_t1, _t3); + _t0 = _mm_castpd_ps(_mm_unpacklo_pd(_mm_castps_pd(_tmp0x), _mm_castps_pd(_tmp1x))); + _t1 = _mm_castpd_ps(_mm_unpacklo_pd(_mm_castps_pd(_tmp2x), _mm_castps_pd(_tmp3x))); + _t2 = _mm_castpd_ps(_mm_unpackhi_pd(_mm_castps_pd(_tmp0x), _mm_castps_pd(_tmp1x))); + _t3 = _mm_castpd_ps(_mm_unpackhi_pd(_mm_castps_pd(_tmp2x), _mm_castps_pd(_tmp3x))); + _t2 = _mm_shuffle_ps(_t2, _t2, _MM_SHUFFLE(2, 3, 0, 1)); + _t3 = _mm_shuffle_ps(_t3, _t3, _MM_SHUFFLE(2, 3, 0, 1)); + __m256 _f0 = combine4x2_ps(_t0, _t1); + __m256 _f1 = combine4x2_ps(_t2, _t3); + if (pC) + { + if (broadcast_type_C == 0) + { + __m256 _c = _mm256_set1_ps(c0); + _f0 = _mm256_add_ps(_f0, _c); + _f1 = _mm256_add_ps(_f1, _c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = _mm256_add_ps(_f0, _mm256_set1_ps(c0)); + _f1 = _mm256_add_ps(_f1, _mm256_set1_ps(c1)); + } + if (broadcast_type_C == 3) + { + __m256 _c0 = _mm256_loadu_ps(pC); + __m256 _c1 = _mm256_loadu_ps(pC + N); + if (beta == 1.f) + { + _f0 = _mm256_add_ps(_f0, _c0); + _f1 = _mm256_add_ps(_f1, _c1); + } + else + { + __m256 _beta = _mm256_set1_ps(beta); +#if __FMA__ + _f0 = _mm256_fmadd_ps(_c0, _beta, _f0); + _f1 = _mm256_fmadd_ps(_c1, _beta, _f1); +#else + _f0 = _mm256_add_ps(_f0, _mm256_mul_ps(_c0, _beta)); + _f1 = _mm256_add_ps(_f1, _mm256_mul_ps(_c1, _beta)); +#endif + } + pC += 8; + } + if (broadcast_type_C == 4) + { + __m256 _c = _mm256_loadu_ps(pC); + if (beta != 1.f) + _c = _mm256_mul_ps(_c, _mm256_set1_ps(beta)); + _f0 = _mm256_add_ps(_f0, _c); + _f1 = _mm256_add_ps(_f1, _c); + pC += 8; + } + } + if (alpha != 1.f) + { + __m256 _alpha = _mm256_set1_ps(alpha); + _f0 = _mm256_mul_ps(_f0, _alpha); + _f1 = _mm256_mul_ps(_f1, _alpha); + } + __m256 _tmp0 = _mm256_unpacklo_ps(_f0, _f1); + __m256 _tmp1 = _mm256_unpackhi_ps(_f0, _f1); + __m128 _r0 = _mm256_castps256_ps128(_tmp0); + __m128 _r1 = _mm256_castps256_ps128(_tmp1); + _mm_storel_pi((__m64*)p0, _r0); + _mm_storeh_pi((__m64*)(p0 + out_hstep), _r0); + _mm_storel_pi((__m64*)(p0 + out_hstep * 2), _r1); + _mm_storeh_pi((__m64*)(p0 + out_hstep * 3), _r1); + _r0 = _mm256_extractf128_ps(_tmp0, 1); + _r1 = _mm256_extractf128_ps(_tmp1, 1); + _mm_storel_pi((__m64*)(p0 + out_hstep * 4), _r0); + _mm_storeh_pi((__m64*)(p0 + out_hstep * 5), _r0); + _mm_storel_pi((__m64*)(p0 + out_hstep * 6), _r1); + _mm_storeh_pi((__m64*)(p0 + out_hstep * 7), _r1); + p0 += out_hstep * 8; + pp += 16; + } +#endif // __AVX512F__ + for (; jj + 3 < max_jj; jj += 4) + { + __m128 _t0 = _mm_loadu_ps(pp + 0); + __m128 _t1 = _mm_loadu_ps(pp + 4); + __m128 _tmp0x = _mm_unpacklo_ps(_t0, _t1); + __m128 _tmp1x = _mm_unpackhi_ps(_t0, _t1); + __m128 _f0 = _mm_castpd_ps(_mm_unpacklo_pd(_mm_castps_pd(_tmp0x), _mm_castps_pd(_tmp1x))); + __m128 _f1 = _mm_castpd_ps(_mm_unpackhi_pd(_mm_castps_pd(_tmp1x), _mm_castps_pd(_tmp0x))); + _f1 = _mm_shuffle_ps(_f1, _f1, _MM_SHUFFLE(0, 3, 2, 1)); + if (pC) + { + if (broadcast_type_C == 0) + { + __m128 _c = _mm_set1_ps(c0); + _f0 = _mm_add_ps(_f0, _c); + _f1 = _mm_add_ps(_f1, _c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = _mm_add_ps(_f0, _mm_set1_ps(c0)); + _f1 = _mm_add_ps(_f1, _mm_set1_ps(c1)); + } + if (broadcast_type_C == 3) + { + __m128 _c0 = _mm_loadu_ps(pC); + __m128 _c1 = _mm_loadu_ps(pC + N); + if (beta == 1.f) + { + _f0 = _mm_add_ps(_f0, _c0); + _f1 = _mm_add_ps(_f1, _c1); + } + else + { + __m128 _beta = _mm_set1_ps(beta); +#if __FMA__ + _f0 = _mm_fmadd_ps(_c0, _beta, _f0); + _f1 = _mm_fmadd_ps(_c1, _beta, _f1); +#else + _f0 = _mm_add_ps(_f0, _mm_mul_ps(_c0, _beta)); + _f1 = _mm_add_ps(_f1, _mm_mul_ps(_c1, _beta)); +#endif + } + pC += 4; + } + if (broadcast_type_C == 4) + { + __m128 _c = _mm_loadu_ps(pC); + if (beta != 1.f) + _c = _mm_mul_ps(_c, _mm_set1_ps(beta)); + _f0 = _mm_add_ps(_f0, _c); + _f1 = _mm_add_ps(_f1, _c); + pC += 4; + } + } + if (alpha != 1.f) + { + __m128 _alpha = _mm_set1_ps(alpha); + _f0 = _mm_mul_ps(_f0, _alpha); + _f1 = _mm_mul_ps(_f1, _alpha); + } + __m128 _tmp0 = _mm_unpacklo_ps(_f0, _f1); + __m128 _tmp1 = _mm_unpackhi_ps(_f0, _f1); + _mm_storel_pi((__m64*)p0, _tmp0); + _mm_storeh_pi((__m64*)(p0 + out_hstep), _tmp0); + _mm_storel_pi((__m64*)(p0 + out_hstep * 2), _tmp1); + _mm_storeh_pi((__m64*)(p0 + out_hstep * 3), _tmp1); + p0 += out_hstep * 4; + pp += 8; + } +#endif // defined(__x86_64__) || defined(_M_X64) + for (; jj + 1 < max_jj; jj += 2) + { + __m128 _f0 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pp + 0)); + __m128 _f1 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pp + 2)); + if (pC) + { + if (broadcast_type_C == 0) + { + __m128 _c = _mm_set1_ps(c0); + _f0 = _mm_add_ps(_f0, _c); + _f1 = _mm_add_ps(_f1, _c); + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = _mm_add_ps(_f0, _mm_set1_ps(c0)); + _f1 = _mm_add_ps(_f1, _mm_set1_ps(c1)); + } + if (broadcast_type_C == 3) + { + __m128 _c0 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pC); + __m128 _c1 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N)); + if (beta == 1.f) + { + _f0 = _mm_add_ps(_f0, _c0); + _f1 = _mm_add_ps(_f1, _c1); + } + else + { + __m128 _beta = _mm_set1_ps(beta); +#if __FMA__ + _f0 = _mm_fmadd_ps(_c0, _beta, _f0); + _f1 = _mm_fmadd_ps(_c1, _beta, _f1); +#else + _f0 = _mm_add_ps(_f0, _mm_mul_ps(_c0, _beta)); + _f1 = _mm_add_ps(_f1, _mm_mul_ps(_c1, _beta)); +#endif + } + pC += 2; + } + if (broadcast_type_C == 4) + { + __m128 _c = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pC); + if (beta != 1.f) + _c = _mm_mul_ps(_c, _mm_set1_ps(beta)); + _f0 = _mm_add_ps(_f0, _c); + _f1 = _mm_add_ps(_f1, _c); + pC += 2; + } + } + if (alpha != 1.f) + { + __m128 _alpha = _mm_set1_ps(alpha); + _f0 = _mm_mul_ps(_f0, _alpha); + _f1 = _mm_mul_ps(_f1, _alpha); + } + __m128 _tmp0 = _mm_unpacklo_ps(_f0, _f1); + _mm_storel_pi((__m64*)p0, _tmp0); + _mm_storeh_pi((__m64*)(p0 + out_hstep), _tmp0); + p0 += out_hstep * 2; + pp += 4; + } + for (; jj < max_jj; jj += 1) + { + __m128 _f = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pp); + if (pC) + { + if (broadcast_type_C == 0) + _f = _mm_add_ps(_f, _mm_set1_ps(c0)); + if (broadcast_type_C == 1 || broadcast_type_C == 2) + _f = _mm_add_ps(_f, _mm_setr_ps(c0, c1, 0.f, 0.f)); + if (broadcast_type_C == 3 || broadcast_type_C == 4) + { + __m128 _c; + if (broadcast_type_C == 3) + _c = _mm_setr_ps(pC[0], pC[N], 0.f, 0.f); + if (broadcast_type_C == 4) + _c = _mm_set1_ps(pC[0]); + if (beta == 1.f) + { + _f = _mm_add_ps(_f, _c); + } + else + { +#if __FMA__ + _f = _mm_fmadd_ps(_c, _mm_set1_ps(beta), _f); +#else + _f = _mm_add_ps(_f, _mm_mul_ps(_c, _mm_set1_ps(beta))); +#endif + } + pC++; + } + } + if (alpha != 1.f) + _f = _mm_mul_ps(_f, _mm_set1_ps(alpha)); + _mm_storel_pi((__m64*)p0, _f); + p0 += out_hstep; + pp += 2; + } +#else + for (; jj + 1 < max_jj; jj += 2) + { + float f0_0 = pp[0]; + float f0_1 = pp[1]; + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + f0_0 += c0; + f0_1 += c0; + } + if (broadcast_type_C == 3 || broadcast_type_C == 4) + { + f0_0 += pC[0] * beta; + f0_1 += pC[1] * beta; + } + } + if (alpha != 1.f) + { + f0_0 *= alpha; + f0_1 *= alpha; + } + + float f1_0 = pp[2]; + float f1_1 = pp[3]; + if (pC) + { + if (broadcast_type_C == 0) + { + f1_0 += c0; + f1_1 += c0; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + f1_0 += c1; + f1_1 += c1; + } + if (broadcast_type_C == 3) + { + f1_0 += pC[N] * beta; + f1_1 += pC[N + 1] * beta; + pC += 2; + } + if (broadcast_type_C == 4) + { + f1_0 += pC[0] * beta; + f1_1 += pC[1] * beta; + pC += 2; + } + } + if (alpha != 1.f) + { + f1_0 *= alpha; + f1_1 *= alpha; + } + + p0[0] = f0_0; + p0[1] = f1_0; + p0[out_hstep] = f0_1; + p0[out_hstep + 1] = f1_1; + + p0 += out_hstep * 2; + pp += 4; + } + for (; jj < max_jj; jj += 1) + { + float f0_0 = pp[0]; + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) f0_0 += c0; + if (broadcast_type_C == 3 || broadcast_type_C == 4) f0_0 += pC[0] * beta; + } + if (alpha != 1.f) f0_0 *= alpha; + + float f1_0 = pp[1]; + if (pC) + { + if (broadcast_type_C == 0) f1_0 += c0; + if (broadcast_type_C == 1 || broadcast_type_C == 2) f1_0 += c1; + if (broadcast_type_C == 3) + { + f1_0 += pC[N] * beta; + pC++; + } + if (broadcast_type_C == 4) + { + f1_0 += pC[0] * beta; + pC++; + } + } + if (alpha != 1.f) f1_0 *= alpha; + + p0[0] = f0_0; + p0[1] = f1_0; + + p0 += out_hstep; + pp += 2; + } +#endif // __SSE2__ + } + + for (; ii < max_ii; ii += 1) + { + float* p0 = (float*)top_blob + j * out_hstep + i + ii; + + float c0 = 0.f; + if (pC) + { + if (broadcast_type_C == 0) + c0 = pC[0]; + if (broadcast_type_C == 1 || broadcast_type_C == 2) + c0 = pC[i + ii]; + if (broadcast_type_C == 3) + pC = (const float*)C + (i + ii) * N + j; + if (broadcast_type_C == 4) + pC = (const float*)C + j; + if ((broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) && beta != 1.f) + c0 *= beta; + } + + int jj = 0; +#if __SSE2__ +#if defined(__x86_64__) || defined(_M_X64) +#if __AVX512F__ + for (; jj + 7 < max_jj; jj += 8) + { + __m256 _f = _mm256_loadu_ps(pp); + if (pC) + { + if (broadcast_type_C == 3 || broadcast_type_C == 4) + { + __m256 _c = _mm256_loadu_ps(pC); + if (beta == 1.f) + { + _f = _mm256_add_ps(_f, _c); + } + else + { +#if __FMA__ + _f = _mm256_fmadd_ps(_c, _mm256_set1_ps(beta), _f); +#else + _f = _mm256_add_ps(_f, _mm256_mul_ps(_c, _mm256_set1_ps(beta))); +#endif + } + pC += 8; + } + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f = _mm256_add_ps(_f, _mm256_set1_ps(c0)); + } + } + if (alpha != 1.f) + _f = _mm256_mul_ps(_f, _mm256_set1_ps(alpha)); + + if (out_hstep == 1) + { + _mm256_storeu_ps(p0, _f); + } + else + { + __m128 _r0 = _mm256_castps256_ps128(_f); + __m128 _r1 = _mm256_extractf128_ps(_f, 1); + _mm_store_ss(p0, _r0); + _r0 = _mm_shuffle_ps(_r0, _r0, _MM_SHUFFLE(0, 3, 2, 1)); + _mm_store_ss(p0 + out_hstep, _r0); + _r0 = _mm_shuffle_ps(_r0, _r0, _MM_SHUFFLE(0, 3, 2, 1)); + _mm_store_ss(p0 + out_hstep * 2, _r0); + _r0 = _mm_shuffle_ps(_r0, _r0, _MM_SHUFFLE(0, 3, 2, 1)); + _mm_store_ss(p0 + out_hstep * 3, _r0); + _mm_store_ss(p0 + out_hstep * 4, _r1); + _r1 = _mm_shuffle_ps(_r1, _r1, _MM_SHUFFLE(0, 3, 2, 1)); + _mm_store_ss(p0 + out_hstep * 5, _r1); + _r1 = _mm_shuffle_ps(_r1, _r1, _MM_SHUFFLE(0, 3, 2, 1)); + _mm_store_ss(p0 + out_hstep * 6, _r1); + _r1 = _mm_shuffle_ps(_r1, _r1, _MM_SHUFFLE(0, 3, 2, 1)); + _mm_store_ss(p0 + out_hstep * 7, _r1); + } + p0 += out_hstep * 8; + pp += 8; + } +#endif // __AVX512F__ + for (; jj + 3 < max_jj; jj += 4) + { + __m128 _f = _mm_loadu_ps(pp); + if (pC) + { + if (broadcast_type_C == 3 || broadcast_type_C == 4) + { + __m128 _c = _mm_loadu_ps(pC); + if (beta == 1.f) + { + _f = _mm_add_ps(_f, _c); + } + else + { +#if __FMA__ + _f = _mm_fmadd_ps(_c, _mm_set1_ps(beta), _f); +#else + _f = _mm_add_ps(_f, _mm_mul_ps(_c, _mm_set1_ps(beta))); +#endif + } + pC += 4; + } + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f = _mm_add_ps(_f, _mm_set1_ps(c0)); + } + } + if (alpha != 1.f) + _f = _mm_mul_ps(_f, _mm_set1_ps(alpha)); + + if (out_hstep == 1) + { + _mm_storeu_ps(p0, _f); + } + else + { + _mm_store_ss(p0, _f); + _f = _mm_shuffle_ps(_f, _f, _MM_SHUFFLE(0, 3, 2, 1)); + _mm_store_ss(p0 + out_hstep, _f); + _f = _mm_shuffle_ps(_f, _f, _MM_SHUFFLE(0, 3, 2, 1)); + _mm_store_ss(p0 + out_hstep * 2, _f); + _f = _mm_shuffle_ps(_f, _f, _MM_SHUFFLE(0, 3, 2, 1)); + _mm_store_ss(p0 + out_hstep * 3, _f); + } + p0 += out_hstep * 4; + pp += 4; + } +#endif // defined(__x86_64__) || defined(_M_X64) + for (; jj + 1 < max_jj; jj += 2) + { + __m128 _f = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pp); + if (pC) + { + if (broadcast_type_C == 3 || broadcast_type_C == 4) + { + __m128 _c = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pC); + if (beta == 1.f) + { + _f = _mm_add_ps(_f, _c); + } + else + { +#if __FMA__ + _f = _mm_fmadd_ps(_c, _mm_set1_ps(beta), _f); +#else + _f = _mm_add_ps(_f, _mm_mul_ps(_c, _mm_set1_ps(beta))); +#endif + } + pC += 2; + } + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f = _mm_add_ps(_f, _mm_set1_ps(c0)); + } + } + if (alpha != 1.f) + _f = _mm_mul_ps(_f, _mm_set1_ps(alpha)); + + if (out_hstep == 1) + { + _mm_storel_pi((__m64*)p0, _f); + } + else + { + _mm_store_ss(p0, _f); + _f = _mm_shuffle_ps(_f, _f, _MM_SHUFFLE(0, 3, 2, 1)); + _mm_store_ss(p0 + out_hstep, _f); + } + p0 += out_hstep * 2; + pp += 2; + } +#else + for (; jj + 1 < max_jj; jj += 2) + { + float f0_0 = pp[0]; + float f0_1 = pp[1]; + if (pC) + { + if (broadcast_type_C == 3 || broadcast_type_C == 4) + { + f0_0 += pC[0] * beta; + f0_1 += pC[1] * beta; + pC += 2; + } + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + { + f0_0 += c0; + f0_1 += c0; + } + } + if (alpha != 1.f) + { + f0_0 *= alpha; + f0_1 *= alpha; + } + p0[0] = f0_0; + p0[out_hstep] = f0_1; + p0 += out_hstep * 2; + pp += 2; + } +#endif // __SSE2__ + for (; jj < max_jj; jj += 1) + { + float f0_0 = pp[0]; + if (pC) + { + if (broadcast_type_C == 3 || broadcast_type_C == 4) + { + f0_0 += beta == 1.f ? pC[0] : pC[0] * beta; + pC++; + } + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + f0_0 += c0; + } + if (alpha != 1.f) f0_0 *= alpha; + p0[0] = f0_0; + p0 += out_hstep; + pp += 1; + } + } +} + +static void get_optimal_tile_mnk_wq_int8(int M, int N, int K, int constant_TILE_M, int constant_TILE_N, int constant_TILE_K, int& TILE_M, int& TILE_N, int& TILE_K, int nT) +{ + // resolve optimal tile size from cache size + const size_t l2_cache_size = get_cpu_level2_cache_size(); + + if (nT == 0) + nT = get_physical_big_cpu_count(); + + const int tile_size = std::max(1, (int)((float)l2_cache_size / 2 / sizeof(signed char) / std::max(1, K))); + +#if defined(__x86_64__) || defined(_M_X64) +#if __AVX512F__ + TILE_M = std::max(16, tile_size / 16 * 16); + TILE_N = std::max(8, tile_size / 8 * 8); +#else + TILE_M = std::max(8, tile_size / 8 * 8); + TILE_N = std::max(4, tile_size / 4 * 4); +#endif // __AVX512F__ +#else +#if __SSE2__ + TILE_M = std::max(4, tile_size / 4 * 4); +#else + TILE_M = std::max(2, tile_size / 2 * 2); +#endif // __SSE2__ + TILE_N = std::max(2, tile_size / 2 * 2); +#endif + TILE_K = K; + + TILE_M *= std::min(nT, get_physical_cpu_count()); + + if (M > 0) + { + const int nn_M = (M + TILE_M - 1) / TILE_M; +#if defined(__x86_64__) || defined(_M_X64) +#if __AVX512F__ + TILE_M = std::min(TILE_M, ((M + nn_M - 1) / nn_M + 15) / 16 * 16); +#else + TILE_M = std::min(TILE_M, ((M + nn_M - 1) / nn_M + 7) / 8 * 8); +#endif // __AVX512F__ +#else +#if __SSE2__ + TILE_M = std::min(TILE_M, ((M + nn_M - 1) / nn_M + 3) / 4 * 4); +#else + TILE_M = std::min(TILE_M, ((M + nn_M - 1) / nn_M + 1) / 2 * 2); +#endif // __SSE2__ +#endif + } + + if (N > 0) + { + const int nn_N = (N + TILE_N - 1) / TILE_N; +#if defined(__x86_64__) || defined(_M_X64) +#if __AVX512F__ + TILE_N = std::min(TILE_N, ((N + nn_N - 1) / nn_N + 7) / 8 * 8); +#else + TILE_N = std::min(TILE_N, ((N + nn_N - 1) / nn_N + 3) / 4 * 4); +#endif // __AVX512F__ +#else + TILE_N = std::min(TILE_N, ((N + nn_N - 1) / nn_N + 1) / 2 * 2); +#endif + } + + if (nT > 1) + { +#if defined(__x86_64__) || defined(_M_X64) +#if __AVX512F__ + TILE_M = std::min(TILE_M, (std::max(1, TILE_M / nT) + 15) / 16 * 16); +#else + TILE_M = std::min(TILE_M, (std::max(1, TILE_M / nT) + 7) / 8 * 8); +#endif // __AVX512F__ +#else +#if __SSE2__ + TILE_M = std::min(TILE_M, (std::max(1, TILE_M / nT) + 3) / 4 * 4); +#else + TILE_M = std::min(TILE_M, (std::max(1, TILE_M / nT) + 1) / 2 * 2); +#endif // __SSE2__ +#endif + } + + // always take constant TILE_M/N value when provided + if (constant_TILE_M > 0) + { +#if defined(__x86_64__) || defined(_M_X64) +#if __AVX512F__ + TILE_M = (constant_TILE_M + 15) / 16 * 16; +#else + TILE_M = (constant_TILE_M + 7) / 8 * 8; +#endif // __AVX512F__ +#else +#if __SSE2__ + TILE_M = (constant_TILE_M + 3) / 4 * 4; +#else + TILE_M = (constant_TILE_M + 1) / 2 * 2; +#endif // __SSE2__ +#endif + } + + if (constant_TILE_N > 0) + { +#if defined(__x86_64__) || defined(_M_X64) +#if __AVX512F__ + TILE_N = (constant_TILE_N + 7) / 8 * 8; +#else + TILE_N = (constant_TILE_N + 3) / 4 * 4; +#endif // __AVX512F__ +#else + TILE_N = (constant_TILE_N + 1) / 2 * 2; +#endif + } + + (void)constant_TILE_K; +} diff --git a/src/layer/x86/gemm_x86.cpp b/src/layer/x86/gemm_x86.cpp index 0d176d4778a5..dd583b0ec0ad 100644 --- a/src/layer/x86/gemm_x86.cpp +++ b/src/layer/x86/gemm_x86.cpp @@ -20,6 +20,10 @@ namespace ncnn { +#if NCNN_WEIGHT_QUANT +#include "gemm_wq_int8.h" +#endif + #if NCNN_INT8 #include "gemm_int8.h" #endif @@ -7432,11 +7436,289 @@ static int gemm_AT_BT_x86(const Mat& AT, const Mat& BT, const Mat& C, Mat& top_b return 0; } +#if NCNN_WEIGHT_QUANT +static int gemm_weight_quantize_bits_x86(int quantize_term) +{ + return (quantize_term / 100) % 10; +} + +static int gemm_weight_quantize_block_size_x86(int quantize_term) +{ + const int block_size_code = quantize_term % 10; + if (block_size_code == 0) return 32; + if (block_size_code == 1) return 64; + if (block_size_code == 2) return 128; + return 0; +} + +static int pack_B_wq_int8_x86(const Mat& B, const Mat& B_scales, Mat& packed_B, Mat& packed_B_descales, int N, int K, int block_size, const Option& opt) +{ + const int block_count = (K + block_size - 1) / block_size; + + // compact persistent stream, each output column occupies exactly K bytes + packed_B.create(N * K, (size_t)1u, 1, opt.blob_allocator); + if (packed_B.empty()) + return -100; + packed_B.cstep = (size_t)N * K; + + packed_B_descales.create(N * block_count, (size_t)4u, 1, opt.blob_allocator); + if (packed_B_descales.empty()) + return -100; + packed_B_descales.cstep = (size_t)N * block_count; + + unsigned char* packed_B_ptr = packed_B; + float* packed_B_descales_ptr = packed_B_descales; + + pack_B_tile_wq_int8(B, B_scales, packed_B_ptr, packed_B_descales_ptr, 0, N, K, block_size); + + return 0; +} + +static int gemm_BT_x86_wq_int8(const Mat& A, const Mat& packed_B, const Mat& packed_B_descales, const Mat& input_scales, const Mat& C, Mat& top_blob, int broadcast_type_C, int N, int K, int block_size, int transA, int output_transpose, float alpha, float beta, int constant_TILE_M, int constant_TILE_N, int constant_TILE_K, int nT, const Option& opt) +{ + const int M = transA ? A.w : (A.dims == 3 ? A.c : A.h) * A.elempack; + const int block_count = (K + block_size - 1) / block_size; + int TILE_M, TILE_N, TILE_K; + get_optimal_tile_mnk_wq_int8(M, N, K, constant_TILE_M, constant_TILE_N, constant_TILE_K, TILE_M, TILE_N, TILE_K, nT); + + (void)TILE_K; + const int mr = std::min(M, TILE_M); + const int nr = std::min(N, TILE_N); + const int nn_M = (M + TILE_M - 1) / TILE_M; + const int nn_N = (N + TILE_N - 1) / TILE_N; + const float* input_scale_ptr = input_scales; + int AT_hstep = K; +#if NCNN_AVX512VNNI || NCNN_AVXVNNI + bool has_w_shift = ncnn::cpu_support_x86_avx512_vnni() || ncnn::cpu_support_x86_avx_vnni(); +#if NCNN_AVXVNNIINT8 + if (ncnn::cpu_support_x86_avx_vnni_int8()) + has_w_shift = false; +#endif // NCNN_AVXVNNIINT8 + if (has_w_shift) + AT_hstep += 4 * ((K + block_size - 4) / block_size); +#endif // NCNN_AVX512VNNI || NCNN_AVXVNNI + const Mat BT = packed_B.reshape(K, N); + const Mat BT_descales = packed_B_descales.reshape(block_count, N); + + Mat topT(mr * nr, 1, nT, 4u, opt.workspace_allocator); + if (topT.empty()) + return -100; + + if (nT > nn_M) + { + Mat AT(AT_hstep, mr, nn_M, 1u, opt.workspace_allocator); + Mat AT_descales(block_count, mr, nn_M, 4u, opt.workspace_allocator); + if (AT.empty() || AT_descales.empty()) + return -100; + + #pragma omp parallel for num_threads(nT) + for (int ppi = 0; ppi < nn_M; ppi++) + { + const int i = ppi * TILE_M; + const int max_ii = std::min(M - i, TILE_M); + + Mat AT_tile = AT.channel(i / TILE_M); + Mat AT_descales_tile = AT_descales.channel(i / TILE_M); + + if (transA) + transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); + else + quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); + } + + const int nn_MN = nn_M * nn_N; + #pragma omp parallel for num_threads(nT) + for (int ppij = 0; ppij < nn_MN; ppij++) + { + const int ppi = ppij / nn_N; + const int ppj = ppij % nn_N; + + const int i = ppi * TILE_M; + const int j = ppj * TILE_N; + + const int max_ii = std::min(M - i, TILE_M); + const int max_jj = std::min(N - j, TILE_N); + + Mat AT_tile = AT.channel(i / TILE_M); + Mat AT_descales_tile = AT_descales.channel(i / TILE_M); + Mat BT_tile = BT.row_range(j, max_jj); + Mat BT_descales_tile = BT_descales.row_range(j, max_jj); + Mat topT_tile = topT.channel(get_omp_thread_num()); + + gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); + + if (output_transpose) + transpose_unpack_output_tile_wq_int8(topT_tile, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, N, alpha, beta); + else + unpack_output_tile_wq_int8(topT_tile, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, N, alpha, beta); + } + } + else + { + Mat ATX(AT_hstep, mr, nT, 1u, opt.workspace_allocator); + Mat ATX_descales(block_count, mr, nT, 4u, opt.workspace_allocator); + if (ATX.empty() || ATX_descales.empty()) + return -100; + + #pragma omp parallel for num_threads(nT) + for (int ppi = 0; ppi < nn_M; ppi++) + { + const int i = ppi * TILE_M; + const int max_ii = std::min(M - i, TILE_M); + + Mat AT_tile = ATX.channel(get_omp_thread_num()); + Mat AT_descales_tile = ATX_descales.channel(get_omp_thread_num()); + Mat topT_tile = topT.channel(get_omp_thread_num()); + + if (transA) + transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); + else + quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); + + for (int j = 0; j < N; j += TILE_N) + { + const int max_jj = std::min(N - j, TILE_N); + + Mat BT_tile = BT.row_range(j, max_jj); + Mat BT_descales_tile = BT_descales.row_range(j, max_jj); + + gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); + + if (output_transpose) + transpose_unpack_output_tile_wq_int8(topT_tile, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, N, alpha, beta); + else + unpack_output_tile_wq_int8(topT_tile, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, N, alpha, beta); + } + } + } + + return 0; +} + +int Gemm_x86::forward_weight_block_quantize_int8(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& opt) const +{ + const Mat& A = bottom_blobs[0]; + if (A.elemsize != 4u || A.elempack != 1) + { + NCNN_LOGE("Gemm unsupported input"); + return -1; + } + + if (transA && A.dims != 2) + { + NCNN_LOGE("Gemm unsupported input"); + return -1; + } + + const int K = transA ? A.h : A.w; + if (K != constantK) + { + NCNN_LOGE("Gemm weight block quantize K mismatch"); + return -1; + } + + const int M = transA ? A.w : A.dims == 3 ? A.c : A.h; + const int N = constantN; + const int block_size = gemm_weight_quantize_block_size_x86(quantize_term); + + Mat C; + int broadcast_type_C = -1; + if (constantC) + { + C = C_data; + broadcast_type_C = constant_broadcast_type_C; + } + else + { + if (bottom_blobs.size() == 2) + C = bottom_blobs[1]; + + if (!C.empty()) + { + bool matched = false; + if (C.dims == 1 && C.w == 1) + { + broadcast_type_C = 0; + matched = true; + } + if (C.dims == 1 && C.w == M) + { + broadcast_type_C = 1; + matched = true; + } + if (C.dims == 1 && C.w == N) + { + broadcast_type_C = 4; + matched = true; + } + if (C.dims == 2 && C.w == 1 && C.h == M) + { + broadcast_type_C = 2; + matched = true; + } + if (C.dims == 2 && C.w == N && C.h == M) + { + broadcast_type_C = 3; + matched = true; + } + if (C.dims == 2 && C.w == N && C.h == 1) + { + broadcast_type_C = 4; + matched = true; + } + + if (!matched || C.elemsize != 4u || C.elempack != 1) + { + NCNN_LOGE("Gemm unsupported C"); + return -1; + } + } + } + + if (!C.empty() && (C.elemsize != 4u || C.elempack != 1)) + { + NCNN_LOGE("Gemm unsupported C"); + return -1; + } + + Mat& top_blob = top_blobs[0]; + if (output_transpose) + top_blob.create(M, N, (size_t)4u, opt.blob_allocator); + else + top_blob.create(N, M, (size_t)4u, opt.blob_allocator); + if (top_blob.empty()) + return -100; + + return gemm_BT_x86_wq_int8(A, B_data_w8a8_packed, B_data_w8a8_descales, B_data_input_scales, C, top_blob, broadcast_type_C, N, K, block_size, transA, output_transpose, alpha, beta, constant_TILE_M, constant_TILE_N, constant_TILE_K, opt.num_threads, opt); +} +#endif // NCNN_WEIGHT_QUANT + int Gemm_x86::create_pipeline(const Option& opt) { if (weight_block_quantize) { - return 0; +#if NCNN_WEIGHT_QUANT + if (gemm_weight_quantize_bits_x86(quantize_term) == 8) + { + if (!B_data_w8a8_packed.empty()) + return 0; + + Mat packed_B; + Mat packed_B_descales; + int ret = pack_B_wq_int8_x86(B_data, B_data_quantize_scales, packed_B, packed_B_descales, constantN, constantK, gemm_weight_quantize_block_size_x86(quantize_term), opt); + if (ret != 0) + return ret; + + B_data_w8a8_packed = packed_B; + B_data_w8a8_descales = packed_B_descales; + B_data.release(); + B_data_quantize_scales.release(); + + return 0; + } +#endif // NCNN_WEIGHT_QUANT + + return Gemm::create_pipeline(opt); } #if NCNN_INT8 @@ -7591,10 +7873,25 @@ int Gemm_x86::create_pipeline(const Option& opt) return 0; } +int Gemm_x86::destroy_pipeline(const Option& opt) +{ +#if NCNN_WEIGHT_QUANT + B_data_w8a8_packed.release(); + B_data_w8a8_descales.release(); +#endif + + return Gemm::destroy_pipeline(opt); +} + int Gemm_x86::forward(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& opt) const { if (weight_block_quantize) { +#if NCNN_WEIGHT_QUANT + if (gemm_weight_quantize_bits_x86(quantize_term) == 8 && !B_data_w8a8_packed.empty()) + return forward_weight_block_quantize_int8(bottom_blobs, top_blobs, opt); +#endif + return Gemm::forward(bottom_blobs, top_blobs, opt); } diff --git a/src/layer/x86/gemm_x86.h b/src/layer/x86/gemm_x86.h index 30263a951f63..74c27ad565c3 100644 --- a/src/layer/x86/gemm_x86.h +++ b/src/layer/x86/gemm_x86.h @@ -14,10 +14,14 @@ class Gemm_x86 : public Gemm Gemm_x86(); virtual int create_pipeline(const Option& opt); + virtual int destroy_pipeline(const Option& opt); virtual int forward(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& opt) const; protected: +#if NCNN_WEIGHT_QUANT + int forward_weight_block_quantize_int8(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& opt) const; +#endif #if NCNN_BF16 int create_pipeline_bf16s(const Option& opt); int forward_bf16s(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& opt) const; @@ -32,6 +36,10 @@ class Gemm_x86 : public Gemm Mat AT_data; Mat BT_data; Mat CT_data; +#if NCNN_WEIGHT_QUANT + Mat B_data_w8a8_packed; + Mat B_data_w8a8_descales; +#endif }; // expose some gemm internal routines for convolution uses diff --git a/src/layer/x86/gemm_x86_avx2.cpp b/src/layer/x86/gemm_x86_avx2.cpp index cab6c7579757..3bbc3f35b249 100644 --- a/src/layer/x86/gemm_x86_avx2.cpp +++ b/src/layer/x86/gemm_x86_avx2.cpp @@ -17,6 +17,40 @@ namespace ncnn { #include "gemm_int8.h" +#if NCNN_WEIGHT_QUANT +#include "gemm_wq_int8.h" + +void pack_B_tile_wq_int8_avx2(const Mat& B, const Mat& B_scales, unsigned char* pp, float* pd, int j, int max_jj, int K, int block_size) +{ + pack_B_tile_wq_int8(B, B_scales, pp, pd, j, max_jj, K, block_size); +} + +void quantize_A_tile_wq_int8_avx2(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr) +{ + quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); +} + +void transpose_quantize_A_tile_wq_int8_avx2(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr) +{ + transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); +} + +void gemm_transB_packed_tile_wq_int8_avx2(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int block_size) +{ + gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); +} + +void unpack_output_tile_wq_int8_avx2(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta) +{ + unpack_output_tile_wq_int8(topT, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, N, alpha, beta); +} + +void transpose_unpack_output_tile_wq_int8_avx2(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta) +{ + transpose_unpack_output_tile_wq_int8(topT, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, N, alpha, beta); +} +#endif // NCNN_WEIGHT_QUANT + void pack_A_tile_int8_avx2(const Mat& A, Mat& AT, int i, int max_ii, int k, int max_kk) { pack_A_tile_int8(A, AT, i, max_ii, k, max_kk); diff --git a/src/layer/x86/gemm_x86_avx512vnni.cpp b/src/layer/x86/gemm_x86_avx512vnni.cpp index d8f46b5bba53..8b03b824dd37 100644 --- a/src/layer/x86/gemm_x86_avx512vnni.cpp +++ b/src/layer/x86/gemm_x86_avx512vnni.cpp @@ -20,6 +20,40 @@ namespace ncnn { #include "gemm_int8.h" +#if NCNN_WEIGHT_QUANT +#include "gemm_wq_int8.h" + +void pack_B_tile_wq_int8_avx512vnni(const Mat& B, const Mat& B_scales, unsigned char* pp, float* pd, int j, int max_jj, int K, int block_size) +{ + pack_B_tile_wq_int8(B, B_scales, pp, pd, j, max_jj, K, block_size); +} + +void quantize_A_tile_wq_int8_avx512vnni(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr) +{ + quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); +} + +void transpose_quantize_A_tile_wq_int8_avx512vnni(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr) +{ + transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); +} + +void gemm_transB_packed_tile_wq_int8_avx512vnni(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int block_size) +{ + gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); +} + +void unpack_output_tile_wq_int8_avx512vnni(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta) +{ + unpack_output_tile_wq_int8(topT, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, N, alpha, beta); +} + +void transpose_unpack_output_tile_wq_int8_avx512vnni(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta) +{ + transpose_unpack_output_tile_wq_int8(topT, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, N, alpha, beta); +} +#endif // NCNN_WEIGHT_QUANT + void pack_A_tile_int8_avx512vnni(const Mat& A, Mat& AT, int i, int max_ii, int k, int max_kk) { pack_A_tile_int8(A, AT, i, max_ii, k, max_kk); diff --git a/src/layer/x86/gemm_x86_avxvnni.cpp b/src/layer/x86/gemm_x86_avxvnni.cpp index 03120ff2e5b2..31f350f69f40 100644 --- a/src/layer/x86/gemm_x86_avxvnni.cpp +++ b/src/layer/x86/gemm_x86_avxvnni.cpp @@ -17,6 +17,40 @@ namespace ncnn { #include "gemm_int8.h" +#if NCNN_WEIGHT_QUANT +#include "gemm_wq_int8.h" + +void pack_B_tile_wq_int8_avxvnni(const Mat& B, const Mat& B_scales, unsigned char* pp, float* pd, int j, int max_jj, int K, int block_size) +{ + pack_B_tile_wq_int8(B, B_scales, pp, pd, j, max_jj, K, block_size); +} + +void quantize_A_tile_wq_int8_avxvnni(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr) +{ + quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); +} + +void transpose_quantize_A_tile_wq_int8_avxvnni(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr) +{ + transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); +} + +void gemm_transB_packed_tile_wq_int8_avxvnni(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int block_size) +{ + gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); +} + +void unpack_output_tile_wq_int8_avxvnni(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta) +{ + unpack_output_tile_wq_int8(topT, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, N, alpha, beta); +} + +void transpose_unpack_output_tile_wq_int8_avxvnni(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta) +{ + transpose_unpack_output_tile_wq_int8(topT, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, N, alpha, beta); +} +#endif // NCNN_WEIGHT_QUANT + void pack_A_tile_int8_avxvnni(const Mat& A, Mat& AT, int i, int max_ii, int k, int max_kk) { pack_A_tile_int8(A, AT, i, max_ii, k, max_kk); diff --git a/src/layer/x86/gemm_x86_avxvnniint8.cpp b/src/layer/x86/gemm_x86_avxvnniint8.cpp index 59e3f7e187b5..f083d0e31b1e 100644 --- a/src/layer/x86/gemm_x86_avxvnniint8.cpp +++ b/src/layer/x86/gemm_x86_avxvnniint8.cpp @@ -17,6 +17,40 @@ namespace ncnn { #include "gemm_int8.h" +#if NCNN_WEIGHT_QUANT +#include "gemm_wq_int8.h" + +void pack_B_tile_wq_int8_avxvnniint8(const Mat& B, const Mat& B_scales, unsigned char* pp, float* pd, int j, int max_jj, int K, int block_size) +{ + pack_B_tile_wq_int8(B, B_scales, pp, pd, j, max_jj, K, block_size); +} + +void quantize_A_tile_wq_int8_avxvnniint8(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr) +{ + quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); +} + +void transpose_quantize_A_tile_wq_int8_avxvnniint8(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr) +{ + transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); +} + +void gemm_transB_packed_tile_wq_int8_avxvnniint8(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int block_size) +{ + gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); +} + +void unpack_output_tile_wq_int8_avxvnniint8(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta) +{ + unpack_output_tile_wq_int8(topT, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, N, alpha, beta); +} + +void transpose_unpack_output_tile_wq_int8_avxvnniint8(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta) +{ + transpose_unpack_output_tile_wq_int8(topT, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, N, alpha, beta); +} +#endif // NCNN_WEIGHT_QUANT + void pack_A_tile_int8_avxvnniint8(const Mat& A, Mat& AT, int i, int max_ii, int k, int max_kk) { pack_A_tile_int8(A, AT, i, max_ii, k, max_kk); diff --git a/src/layer/x86/gemm_x86_xop.cpp b/src/layer/x86/gemm_x86_xop.cpp index e2f58dbca5f9..f2ee1367331f 100644 --- a/src/layer/x86/gemm_x86_xop.cpp +++ b/src/layer/x86/gemm_x86_xop.cpp @@ -17,6 +17,15 @@ namespace ncnn { #include "gemm_int8.h" +#if NCNN_WEIGHT_QUANT +#include "gemm_wq_int8.h" + +void gemm_transB_packed_tile_wq_int8_xop(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int block_size) +{ + gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); +} +#endif // NCNN_WEIGHT_QUANT + void gemm_transB_packed_tile_int8_xop(const Mat& AT_tile, const Mat& BT_tile, Mat& topT_tile, int i, int max_ii, int j, int max_jj, int k, int max_kk) { gemm_transB_packed_tile_int8(AT_tile, BT_tile, topT_tile, i, max_ii, j, max_jj, k, max_kk); diff --git a/src/layer/x86/multiheadattention_x86.cpp b/src/layer/x86/multiheadattention_x86.cpp index b1cc65423eb0..5e0d4583760e 100644 --- a/src/layer/x86/multiheadattention_x86.cpp +++ b/src/layer/x86/multiheadattention_x86.cpp @@ -29,10 +29,355 @@ MultiHeadAttention_x86::MultiHeadAttention_x86() o_gemm = 0; } +#if NCNN_WEIGHT_QUANT +int MultiHeadAttention_x86::create_pipeline_wq_int8(const Option& _opt) +{ + if (q_gemm) + return 0; + + Option opt = _opt; + Option opt_wq = opt; + opt_wq.use_packing_layout = false; + opt_wq.use_fp16_packed = false; + opt_wq.use_fp16_storage = false; + opt_wq.use_fp16_arithmetic = false; + opt_wq.use_bf16_packed = false; + opt_wq.use_bf16_storage = false; + + { + qk_softmax = ncnn::create_layer_cpu(ncnn::LayerType::Softmax); + if (!qk_softmax) + return -100; + ncnn::ParamDict pd; + pd.set(0, -1); + pd.set(1, 1); + int ret = qk_softmax->load_param(pd); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = qk_softmax->load_model(ModelBinFromMatArray(0)); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = qk_softmax->create_pipeline(opt); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + } + + const int qdim = weight_data_size / embed_dim; + + { + q_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!q_gemm) + { + destroy_pipeline(_opt); + return -100; + } + ncnn::ParamDict pd; + pd.set(0, scale); + pd.set(1, 1.f); + pd.set(2, 0); // transA + pd.set(3, 1); // transB + pd.set(4, 0); // constantA + pd.set(5, 1); // constantB + pd.set(6, 1); // constantC + pd.set(7, 0); // M + pd.set(8, embed_dim); // N + pd.set(9, qdim); // K + pd.set(10, 4); // constant_broadcast_type_C + pd.set(11, 0); // output_N1M + pd.set(12, 0); // output_elempack + pd.set(14, 1); // output_transpose + pd.set(18, quantize_term); + int ret = q_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + Mat weights[4]; + weights[0] = q_weight_data; + weights[1] = q_bias_data; + weights[2] = q_weight_data_quantize_scales; + weights[3] = q_weight_data_input_scales; + ret = q_gemm->load_model(ModelBinFromMatArray(weights)); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = q_gemm->create_pipeline(opt_wq); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + } + + { + k_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!k_gemm) + { + destroy_pipeline(_opt); + return -100; + } + ncnn::ParamDict pd; + pd.set(2, 0); // transA + pd.set(3, 1); // transB + pd.set(4, 0); // constantA + pd.set(5, 1); // constantB + pd.set(6, 1); // constantC + pd.set(7, 0); // M + pd.set(8, embed_dim); // N + pd.set(9, kdim); // K + pd.set(10, 4); // constant_broadcast_type_C + pd.set(11, 0); // output_N1M + pd.set(12, 0); // output_elempack + pd.set(14, 1); // output_transpose + pd.set(18, quantize_term); + int ret = k_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + Mat weights[4]; + weights[0] = k_weight_data; + weights[1] = k_bias_data; + weights[2] = k_weight_data_quantize_scales; + weights[3] = k_weight_data_input_scales; + ret = k_gemm->load_model(ModelBinFromMatArray(weights)); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = k_gemm->create_pipeline(opt_wq); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + } + + { + v_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!v_gemm) + { + destroy_pipeline(_opt); + return -100; + } + ncnn::ParamDict pd; + pd.set(2, 0); // transA + pd.set(3, 1); // transB + pd.set(4, 0); // constantA + pd.set(5, 1); // constantB + pd.set(6, 1); // constantC + pd.set(7, 0); // M + pd.set(8, embed_dim); // N + pd.set(9, vdim); // K + pd.set(10, 4); // constant_broadcast_type_C + pd.set(11, 0); // output_N1M + pd.set(12, 0); // output_elempack + pd.set(14, 1); // output_transpose + pd.set(18, quantize_term); + int ret = v_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + Mat weights[4]; + weights[0] = v_weight_data; + weights[1] = v_bias_data; + weights[2] = v_weight_data_quantize_scales; + weights[3] = v_weight_data_input_scales; + ret = v_gemm->load_model(ModelBinFromMatArray(weights)); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = v_gemm->create_pipeline(opt_wq); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + } + + { + o_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!o_gemm) + { + destroy_pipeline(_opt); + return -100; + } + ncnn::ParamDict pd; + pd.set(2, 1); // transA + pd.set(3, 1); // transB + pd.set(4, 0); // constantA + pd.set(5, 1); // constantB + pd.set(6, 1); // constantC + pd.set(7, 0); // M = outch + pd.set(8, qdim); // N = size + pd.set(9, embed_dim); // K = maxk*inch + pd.set(10, 4); // constant_broadcast_type_C + pd.set(11, 0); // output_N1M + pd.set(18, quantize_term); + int ret = o_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + Mat weights[4]; + weights[0] = out_weight_data; + weights[1] = out_bias_data; + weights[2] = out_weight_data_quantize_scales; + weights[3] = out_weight_data_input_scales; + ret = o_gemm->load_model(ModelBinFromMatArray(weights)); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = o_gemm->create_pipeline(opt_wq); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + } + + { + qk_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!qk_gemm) + { + destroy_pipeline(_opt); + return -100; + } + ncnn::ParamDict pd; + pd.set(2, 1); // transA + pd.set(3, 0); // transB + pd.set(4, 0); // constantA + pd.set(5, 0); // constantB + pd.set(6, attn_mask ? 0 : 1); // constantC + pd.set(7, 0); // M + pd.set(8, 0); // N + pd.set(9, 0); // K + pd.set(10, attn_mask ? 3 : -1); // constant_broadcast_type_C + pd.set(11, 0); // output_N1M + pd.set(12, 1); // output_elempack + pd.set(13, 1); // output_elemtype = fp32 + int ret = qk_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = qk_gemm->load_model(ModelBinFromMatArray(0)); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + Option opt1 = opt; + opt1.use_bf16_packed = false; + opt1.use_bf16_storage = false; + opt1.num_threads = 1; + ret = qk_gemm->create_pipeline(opt1); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + } + + { + qkv_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!qkv_gemm) + { + destroy_pipeline(_opt); + return -100; + } + ncnn::ParamDict pd; + pd.set(2, 0); // transA + pd.set(3, 1); // transB + pd.set(4, 0); // constantA + pd.set(5, 0); // constantB + pd.set(6, 1); // constantC + pd.set(7, 0); // M + pd.set(8, 0); // N + pd.set(9, 0); // K + pd.set(10, -1); // constant_broadcast_type_C + pd.set(11, 0); // output_N1M + pd.set(12, 1); // output_elempack + pd.set(13, 1); // output_elemtype = fp32 + pd.set(14, 1); // output_transpose + int ret = qkv_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + ret = qkv_gemm->load_model(ModelBinFromMatArray(0)); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + Option opt1 = opt; + opt1.use_bf16_packed = false; + opt1.use_bf16_storage = false; + opt1.num_threads = 1; + ret = qkv_gemm->create_pipeline(opt1); + if (ret != 0) + { + destroy_pipeline(_opt); + return ret; + } + } + + q_weight_data.release(); + q_bias_data.release(); + k_weight_data.release(); + k_bias_data.release(); + v_weight_data.release(); + v_bias_data.release(); + out_weight_data.release(); + out_bias_data.release(); + q_weight_data_quantize_scales.release(); + k_weight_data_quantize_scales.release(); + v_weight_data_quantize_scales.release(); + out_weight_data_quantize_scales.release(); + q_weight_data_input_scales.release(); + k_weight_data_input_scales.release(); + v_weight_data_input_scales.release(); + out_weight_data_input_scales.release(); + + return 0; +} +#endif // NCNN_WEIGHT_QUANT + int MultiHeadAttention_x86::create_pipeline(const Option& _opt) { +#if NCNN_WEIGHT_QUANT if (weight_block_quantize) - return 0; + { + if (quantize_term / 100 != 8) + return MultiHeadAttention::create_pipeline(_opt); + + return create_pipeline_wq_int8(_opt); + } +#endif Option opt = _opt; if (int8_scale_term) @@ -45,18 +390,40 @@ int MultiHeadAttention_x86::create_pipeline(const Option& _opt) { qk_softmax = ncnn::create_layer_cpu(ncnn::LayerType::Softmax); + if (!qk_softmax) + return -100; ncnn::ParamDict pd; pd.set(0, -1); pd.set(1, 1); - qk_softmax->load_param(pd); - qk_softmax->load_model(ModelBinFromMatArray(0)); - qk_softmax->create_pipeline(opt); + int ret = qk_softmax->load_param(pd); + if (ret != 0) + { + destroy_pipeline(opt); + return ret; + } + ret = qk_softmax->load_model(ModelBinFromMatArray(0)); + if (ret != 0) + { + destroy_pipeline(opt); + return ret; + } + ret = qk_softmax->create_pipeline(opt); + if (ret != 0) + { + destroy_pipeline(opt); + return ret; + } } const int qdim = weight_data_size / embed_dim; { q_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!q_gemm) + { + destroy_pipeline(opt); + return -100; + } ncnn::ParamDict pd; pd.set(0, scale); pd.set(1, 1.f); @@ -75,25 +442,39 @@ int MultiHeadAttention_x86::create_pipeline(const Option& _opt) #if NCNN_INT8 pd.set(18, int8_scale_term); #endif - q_gemm->load_param(pd); + int ret = q_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(opt); + return ret; + } Mat weights[3]; weights[0] = q_weight_data; weights[1] = q_bias_data; #if NCNN_INT8 weights[2] = q_weight_data_int8_scales; #endif - q_gemm->load_model(ModelBinFromMatArray(weights)); - q_gemm->create_pipeline(opt); - - if (opt.lightmode) + ret = q_gemm->load_model(ModelBinFromMatArray(weights)); + if (ret != 0) { - q_weight_data.release(); - q_bias_data.release(); + destroy_pipeline(opt); + return ret; + } + ret = q_gemm->create_pipeline(opt); + if (ret != 0) + { + destroy_pipeline(opt); + return ret; } } { k_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!k_gemm) + { + destroy_pipeline(opt); + return -100; + } ncnn::ParamDict pd; pd.set(2, 0); // transA pd.set(3, 1); // transB @@ -110,25 +491,39 @@ int MultiHeadAttention_x86::create_pipeline(const Option& _opt) #if NCNN_INT8 pd.set(18, int8_scale_term); #endif - k_gemm->load_param(pd); + int ret = k_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(opt); + return ret; + } Mat weights[3]; weights[0] = k_weight_data; weights[1] = k_bias_data; #if NCNN_INT8 weights[2] = k_weight_data_int8_scales; #endif - k_gemm->load_model(ModelBinFromMatArray(weights)); - k_gemm->create_pipeline(opt); - - if (opt.lightmode) + ret = k_gemm->load_model(ModelBinFromMatArray(weights)); + if (ret != 0) { - k_weight_data.release(); - k_bias_data.release(); + destroy_pipeline(opt); + return ret; + } + ret = k_gemm->create_pipeline(opt); + if (ret != 0) + { + destroy_pipeline(opt); + return ret; } } { v_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!v_gemm) + { + destroy_pipeline(opt); + return -100; + } ncnn::ParamDict pd; pd.set(2, 0); // transA pd.set(3, 1); // transB @@ -145,25 +540,39 @@ int MultiHeadAttention_x86::create_pipeline(const Option& _opt) #if NCNN_INT8 pd.set(18, int8_scale_term); #endif - v_gemm->load_param(pd); + int ret = v_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(opt); + return ret; + } Mat weights[3]; weights[0] = v_weight_data; weights[1] = v_bias_data; #if NCNN_INT8 weights[2] = v_weight_data_int8_scales; #endif - v_gemm->load_model(ModelBinFromMatArray(weights)); - v_gemm->create_pipeline(opt); - - if (opt.lightmode) + ret = v_gemm->load_model(ModelBinFromMatArray(weights)); + if (ret != 0) { - v_weight_data.release(); - v_bias_data.release(); + destroy_pipeline(opt); + return ret; + } + ret = v_gemm->create_pipeline(opt); + if (ret != 0) + { + destroy_pipeline(opt); + return ret; } } { o_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!o_gemm) + { + destroy_pipeline(opt); + return -100; + } ncnn::ParamDict pd; pd.set(2, 1); // transA pd.set(3, 1); // transB @@ -178,30 +587,49 @@ int MultiHeadAttention_x86::create_pipeline(const Option& _opt) #if NCNN_INT8 pd.set(18, int8_scale_term); #endif - o_gemm->load_param(pd); + int ret = o_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(opt); + return ret; + } Mat weights[3]; weights[0] = out_weight_data; weights[1] = out_bias_data; #if NCNN_INT8 Mat out_weight_data_int8_scales(1); + if (out_weight_data_int8_scales.empty()) + { + destroy_pipeline(opt); + return -100; + } out_weight_data_int8_scales[0] = out_weight_data_int8_scale; weights[2] = out_weight_data_int8_scales; #endif - o_gemm->load_model(ModelBinFromMatArray(weights)); + ret = o_gemm->load_model(ModelBinFromMatArray(weights)); + if (ret != 0) + { + destroy_pipeline(opt); + return ret; + } Option opt_fp32 = opt; opt_fp32.use_bf16_packed = false; opt_fp32.use_bf16_storage = false; - o_gemm->create_pipeline(opt_fp32); - - if (opt.lightmode) + ret = o_gemm->create_pipeline(opt_fp32); + if (ret != 0) { - out_weight_data.release(); - out_bias_data.release(); + destroy_pipeline(opt); + return ret; } } { qk_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!qk_gemm) + { + destroy_pipeline(opt); + return -100; + } ncnn::ParamDict pd; pd.set(2, 1); // transA pd.set(3, 0); // transB @@ -218,17 +646,37 @@ int MultiHeadAttention_x86::create_pipeline(const Option& _opt) #if NCNN_INT8 pd.set(18, int8_scale_term); #endif - qk_gemm->load_param(pd); - qk_gemm->load_model(ModelBinFromMatArray(0)); + int ret = qk_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(opt); + return ret; + } + ret = qk_gemm->load_model(ModelBinFromMatArray(0)); + if (ret != 0) + { + destroy_pipeline(opt); + return ret; + } Option opt1 = opt; opt1.use_bf16_packed = false; opt1.use_bf16_storage = false; opt1.num_threads = 1; - qk_gemm->create_pipeline(opt1); + ret = qk_gemm->create_pipeline(opt1); + if (ret != 0) + { + destroy_pipeline(opt); + return ret; + } } { qkv_gemm = ncnn::create_layer_cpu(ncnn::LayerType::Gemm); + if (!qkv_gemm) + { + destroy_pipeline(opt); + return -100; + } ncnn::ParamDict pd; pd.set(2, 0); // transA pd.set(3, 1); // transB @@ -246,13 +694,40 @@ int MultiHeadAttention_x86::create_pipeline(const Option& _opt) #if NCNN_INT8 pd.set(18, int8_scale_term); #endif - qkv_gemm->load_param(pd); - qkv_gemm->load_model(ModelBinFromMatArray(0)); + int ret = qkv_gemm->load_param(pd); + if (ret != 0) + { + destroy_pipeline(opt); + return ret; + } + ret = qkv_gemm->load_model(ModelBinFromMatArray(0)); + if (ret != 0) + { + destroy_pipeline(opt); + return ret; + } Option opt1 = opt; opt1.use_bf16_packed = false; opt1.use_bf16_storage = false; opt1.num_threads = 1; - qkv_gemm->create_pipeline(opt1); + ret = qkv_gemm->create_pipeline(opt1); + if (ret != 0) + { + destroy_pipeline(opt); + return ret; + } + } + + if (opt.lightmode) + { + q_weight_data.release(); + q_bias_data.release(); + k_weight_data.release(); + k_bias_data.release(); + v_weight_data.release(); + v_bias_data.release(); + out_weight_data.release(); + out_bias_data.release(); } return 0; @@ -260,14 +735,24 @@ int MultiHeadAttention_x86::create_pipeline(const Option& _opt) int MultiHeadAttention_x86::destroy_pipeline(const Option& _opt) { - if (weight_block_quantize) - return 0; + if (weight_block_quantize && quantize_term / 100 != 8) + return MultiHeadAttention::destroy_pipeline(_opt); Option opt = _opt; - if (int8_scale_term) + if (int8_scale_term && !weight_block_quantize) { opt.use_packing_layout = false; // TODO enable packing } + Option opt_wq = opt; + if (weight_block_quantize) + { + opt_wq.use_packing_layout = false; + opt_wq.use_fp16_packed = false; + opt_wq.use_fp16_storage = false; + opt_wq.use_fp16_arithmetic = false; + opt_wq.use_bf16_packed = false; + opt_wq.use_bf16_storage = false; + } if (qk_softmax) { @@ -278,28 +763,28 @@ int MultiHeadAttention_x86::destroy_pipeline(const Option& _opt) if (q_gemm) { - q_gemm->destroy_pipeline(opt); + q_gemm->destroy_pipeline(opt_wq); delete q_gemm; q_gemm = 0; } if (k_gemm) { - k_gemm->destroy_pipeline(opt); + k_gemm->destroy_pipeline(opt_wq); delete k_gemm; k_gemm = 0; } if (v_gemm) { - v_gemm->destroy_pipeline(opt); + v_gemm->destroy_pipeline(opt_wq); delete v_gemm; v_gemm = 0; } if (o_gemm) { - o_gemm->destroy_pipeline(opt); + o_gemm->destroy_pipeline(opt_wq); delete o_gemm; o_gemm = 0; } @@ -322,7 +807,7 @@ int MultiHeadAttention_x86::destroy_pipeline(const Option& _opt) int MultiHeadAttention_x86::forward(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& _opt) const { - if (weight_block_quantize) + if (weight_block_quantize && quantize_term / 100 != 8) return MultiHeadAttention::forward(bottom_blobs, top_blobs, _opt); int q_blob_i = 0; @@ -341,10 +826,20 @@ int MultiHeadAttention_x86::forward(const std::vector& bottom_blobs, std::v const Mat& cached_xv_blob = kv_cache ? bottom_blobs[cached_xv_i] : Mat(); Option opt = _opt; - if (int8_scale_term) + if (int8_scale_term && !weight_block_quantize) { opt.use_packing_layout = false; // TODO enable packing } + Option opt_wq = opt; + if (weight_block_quantize) + { + opt_wq.use_packing_layout = false; + opt_wq.use_fp16_packed = false; + opt_wq.use_fp16_storage = false; + opt_wq.use_fp16_arithmetic = false; + opt_wq.use_bf16_packed = false; + opt_wq.use_bf16_storage = false; + } Mat attn_mask_blob_unpacked; if (attn_mask && attn_mask_blob.elempack != 1) @@ -389,7 +884,7 @@ int MultiHeadAttention_x86::forward(const std::vector& bottom_blobs, std::v const int dst_seqlen = past_seqlen > 0 ? (q_blob_i == k_blob_i ? (past_seqlen + cur_seqlen) : past_seqlen) : cur_seqlen; Mat q_affine; - int retq = q_gemm->forward(q_blob, q_affine, opt); + int retq = q_gemm->forward(q_blob, q_affine, opt_wq); if (retq != 0) return retq; @@ -399,7 +894,7 @@ int MultiHeadAttention_x86::forward(const std::vector& bottom_blobs, std::v if (q_blob_i == k_blob_i) { Mat k_affine_q; - int retk = k_gemm->forward(q_blob, k_affine_q, opt); + int retk = k_gemm->forward(q_blob, k_affine_q, opt_wq); if (retk != 0) return retk; @@ -427,7 +922,7 @@ int MultiHeadAttention_x86::forward(const std::vector& bottom_blobs, std::v } else { - int retk = k_gemm->forward(k_blob, k_affine, opt); + int retk = k_gemm->forward(k_blob, k_affine, opt_wq); if (retk != 0) return retk; } @@ -478,7 +973,7 @@ int MultiHeadAttention_x86::forward(const std::vector& bottom_blobs, std::v if (q_blob_i == v_blob_i) { Mat v_affine_q; - int retk = v_gemm->forward(v_blob, v_affine_q, opt); + int retk = v_gemm->forward(v_blob, v_affine_q, opt_wq); if (retk != 0) return retk; @@ -506,7 +1001,7 @@ int MultiHeadAttention_x86::forward(const std::vector& bottom_blobs, std::v } else { - int retv = v_gemm->forward(v_blob, v_affine, opt); + int retv = v_gemm->forward(v_blob, v_affine, opt_wq); if (retv != 0) return retv; } @@ -553,7 +1048,7 @@ int MultiHeadAttention_x86::forward(const std::vector& bottom_blobs, std::v v_affine.release(); } - int reto = o_gemm->forward(qkv_cross, top_blobs[0], opt); + int reto = o_gemm->forward(qkv_cross, top_blobs[0], opt_wq); if (reto != 0) return reto; diff --git a/src/layer/x86/multiheadattention_x86.h b/src/layer/x86/multiheadattention_x86.h index 66d88910c108..fb0fb6aa8abf 100644 --- a/src/layer/x86/multiheadattention_x86.h +++ b/src/layer/x86/multiheadattention_x86.h @@ -18,6 +18,11 @@ class MultiHeadAttention_x86 : public MultiHeadAttention virtual int forward(const std::vector& bottom_blobs, std::vector& top_blobs, const Option& opt) const; +protected: +#if NCNN_WEIGHT_QUANT + int create_pipeline_wq_int8(const Option& opt); +#endif + public: Layer* q_gemm; Layer* k_gemm; From 6d3f067b85cf77d3ce25e2ab68ae6819d83ef7f7 Mon Sep 17 00:00:00 2001 From: nihui <171016+nihui@users.noreply.github.com> Date: Thu, 16 Jul 2026 15:44:57 +0000 Subject: [PATCH 2/4] apply code-format changes --- src/layer/arm/gemm_wq_int8.h | 310 +++++++++++++------ src/layer/loongarch/gemm_loongarch.cpp | 6 +- src/layer/loongarch/gemm_wq_int8.h | 8 +- src/layer/mips/gemm_mips.cpp | 36 ++- src/layer/mips/gemm_wq_int8.h | 306 +++++++++++++----- src/layer/riscv/gemm_wq_int8.h | 13 +- src/layer/riscv/multiheadattention_riscv.cpp | 1 - src/layer/x86/gemm_wq_int8.h | 125 +++++--- 8 files changed, 560 insertions(+), 245 deletions(-) diff --git a/src/layer/arm/gemm_wq_int8.h b/src/layer/arm/gemm_wq_int8.h index 3aa907278537..09fb487e41fd 100644 --- a/src/layer/arm/gemm_wq_int8.h +++ b/src/layer/arm/gemm_wq_int8.h @@ -111,7 +111,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales _v10 = vmulq_f32(_v10, _s0); _v11 = vmulq_f32(_v11, _s1); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+w"(_v00), "+w"(_v01), "+w"(_v10), "+w"(_v11)); + asm volatile("" + : "+w"(_v00), "+w"(_v01), "+w"(_v10), "+w"(_v11)); #endif } const float32x4_t _scale0 = vdupq_n_f32(scales[r]); @@ -134,7 +135,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales { _v = vmulq_f32(_v, vld1q_f32(input_scale_ptr + k0 + kk)); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+w"(_v)); + asm volatile("" + : "+w"(_v)); #endif } const float32x4_t _scale = vdupq_n_f32(scales[r]); @@ -156,7 +158,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales v0 *= input_scale_ptr[k0 + kk]; v1 *= input_scale_ptr[k0 + kk + 1]; #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+w"(v0), "+w"(v1)); + asm volatile("" + : "+w"(v0), "+w"(v1)); #endif } *pp++ = float2int8(v0 * scales[r]); @@ -173,7 +176,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales { v *= input_scale_ptr[k0 + kk]; #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+w"(v)); + asm volatile("" + : "+w"(v)); #endif } *pp++ = float2int8(v * scales[r]); @@ -251,7 +255,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales _v10 = vmulq_f32(_v10, _s0); _v11 = vmulq_f32(_v11, _s1); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+w"(_v00), "+w"(_v01), "+w"(_v10), "+w"(_v11)); + asm volatile("" + : "+w"(_v00), "+w"(_v01), "+w"(_v10), "+w"(_v11)); #endif } const float32x4_t _scale0 = vdupq_n_f32(scales[r]); @@ -274,7 +279,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales { _v = vmulq_f32(_v, vld1q_f32(input_scale_ptr + k0 + kk)); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+w"(_v)); + asm volatile("" + : "+w"(_v)); #endif } const int8x8_t _q = float2int8(vmulq_n_f32(_v, scales[r]), vmulq_n_f32(_v, scales[r])); @@ -295,7 +301,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales v0 *= input_scale_ptr[k0 + kk]; v1 *= input_scale_ptr[k0 + kk + 1]; #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+w"(v0), "+w"(v1)); + asm volatile("" + : "+w"(v0), "+w"(v1)); #endif } *pp++ = float2int8(v0 * scales[r]); @@ -312,7 +319,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales { v *= input_scale_ptr[k0 + kk]; #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+w"(v)); + asm volatile("" + : "+w"(v)); #endif } *pp++ = float2int8(v * scales[r]); @@ -382,7 +390,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales _v10 = vmulq_f32(_v10, _s0); _v11 = vmulq_f32(_v11, _s1); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+w"(_v00), "+w"(_v01), "+w"(_v10), "+w"(_v11)); + asm volatile("" + : "+w"(_v00), "+w"(_v01), "+w"(_v10), "+w"(_v11)); #endif } const int8x8_t _q0 = float2int8(vmulq_n_f32(_v00, scales[0]), vmulq_n_f32(_v01, scales[0])); @@ -403,7 +412,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales { _v = vmulq_f32(_v, vld1q_f32(input_scale_ptr + k0 + kk)); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+w"(_v)); + asm volatile("" + : "+w"(_v)); #endif } const int8x8_t _q = float2int8(vmulq_n_f32(_v, scales[r]), vmulq_n_f32(_v, scales[r])); @@ -485,7 +495,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales { v0 *= input_scale_ptr[k0 + kk]; v1 *= input_scale_ptr[k0 + kk]; - asm volatile("" : "+w"(v0), "+w"(v1)); + asm volatile("" + : "+w"(v0), "+w"(v1)); } *pp++ = float2int8(v0 * scale0); *pp++ = float2int8(v1 * scale1); @@ -556,7 +567,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales _v0 = vmulq_f32(_v0, vld1q_f32(input_scale_ptr + k0 + kk)); _v1 = vmulq_f32(_v1, vld1q_f32(input_scale_ptr + k0 + kk + 4)); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+w"(_v0), "+w"(_v1)); + asm volatile("" + : "+w"(_v0), "+w"(_v1)); #else volatile float32x4_t _v0_ordered = _v0; volatile float32x4_t _v1_ordered = _v1; @@ -573,7 +585,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales { _v = vmulq_f32(_v, vld1q_f32(input_scale_ptr + k0 + kk)); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+w"(_v)); + asm volatile("" + : "+w"(_v)); #else volatile float32x4_t _v_ordered = _v; _v = _v_ordered; @@ -591,7 +604,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales { v *= input_scale_ptr[k]; #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+w"(v)); + asm volatile("" + : "+w"(v)); #else volatile float v_ordered = v; v = v_ordered; @@ -690,7 +704,8 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A _v0 = vmulq_n_f32(_v0, input_scale_ptr[k]); _v1 = vmulq_n_f32(_v1, input_scale_ptr[k]); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+w"(_v0), "+w"(_v1)); + asm volatile("" + : "+w"(_v0), "+w"(_v1)); #endif } _q[t] = float2int8(vmulq_f32(_v0, _scale0), vmulq_f32(_v1, _scale1)); @@ -728,7 +743,8 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A _v0 = vmulq_n_f32(_v0, input_scale_ptr[k]); _v1 = vmulq_n_f32(_v1, input_scale_ptr[k]); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+w"(_v0), "+w"(_v1)); + asm volatile("" + : "+w"(_v0), "+w"(_v1)); #endif } _q.val[t] = float2int8(vmulq_f32(_v0, _scale0), vmulq_f32(_v1, _scale1)); @@ -750,7 +766,8 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A _v0 = vmulq_n_f32(_v0, input_scale_ptr[k]); _v1 = vmulq_n_f32(_v1, input_scale_ptr[k]); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+w"(_v0), "+w"(_v1)); + asm volatile("" + : "+w"(_v0), "+w"(_v1)); #endif } _q.val[t] = float2int8(vmulq_f32(_v0, _scale0), vmulq_f32(_v1, _scale1)); @@ -768,7 +785,8 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A _v0 = vmulq_n_f32(_v0, input_scale_ptr[k]); _v1 = vmulq_n_f32(_v1, input_scale_ptr[k]); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+w"(_v0), "+w"(_v1)); + asm volatile("" + : "+w"(_v0), "+w"(_v1)); #endif } vst1_s8(pp, float2int8(vmulq_f32(_v0, _scale0), vmulq_f32(_v1, _scale1))); @@ -827,7 +845,8 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A { _v = vmulq_n_f32(_v, input_scale_ptr[k]); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+w"(_v)); + asm volatile("" + : "+w"(_v)); #else volatile float32x4_t _v_ordered = _v; _v = _v_ordered; @@ -860,7 +879,8 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A { _v = vmulq_n_f32(_v, input_scale_ptr[k]); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+w"(_v)); + asm volatile("" + : "+w"(_v)); #else volatile float32x4_t _v_ordered = _v; _v = _v_ordered; @@ -886,7 +906,8 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A { _v = vmulq_n_f32(_v, input_scale_ptr[k]); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+w"(_v)); + asm volatile("" + : "+w"(_v)); #else volatile float32x4_t _v_ordered = _v; _v = _v_ordered; @@ -908,7 +929,8 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A { _v = vmulq_n_f32(_v, input_scale_ptr[k]); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+w"(_v)); + asm volatile("" + : "+w"(_v)); #else volatile float32x4_t _v_ordered = _v; _v = _v_ordered; @@ -1071,7 +1093,8 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A { v0 *= input_scale_ptr[k0 + kk]; v1 *= input_scale_ptr[k0 + kk]; - asm volatile("" : "+w"(v0), "+w"(v1)); + asm volatile("" + : "+w"(v0), "+w"(v1)); } *pp++ = float2int8(v0 * scale0); *pp++ = float2int8(v1 * scale1); @@ -1121,7 +1144,8 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A { v *= input_scale_ptr[k]; #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+w"(v)); + asm volatile("" + : "+w"(v)); #else volatile float v_ordered = v; v = v_ordered; @@ -1845,14 +1869,22 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de const int8x16_t _a67 = vld1q_s8(pA + 48); const int8x16_t _b0 = vcombine_s8(vreinterpret_s8_s32(vld1_dup_s32((const int*)pB)), vdup_n_s8(0)); const int8x16_t _b1 = vcombine_s8(vreinterpret_s8_s32(vld1_dup_s32((const int*)(pB + 4))), vdup_n_s8(0)); - _sum0 = vdotq_laneq_s32(_sum0, _b0, _a01, 0); _sum0 = vdotq_laneq_s32(_sum0, _b1, _a01, 1); - _sum1 = vdotq_laneq_s32(_sum1, _b0, _a01, 2); _sum1 = vdotq_laneq_s32(_sum1, _b1, _a01, 3); - _sum2 = vdotq_laneq_s32(_sum2, _b0, _a23, 0); _sum2 = vdotq_laneq_s32(_sum2, _b1, _a23, 1); - _sum3 = vdotq_laneq_s32(_sum3, _b0, _a23, 2); _sum3 = vdotq_laneq_s32(_sum3, _b1, _a23, 3); - _sum4 = vdotq_laneq_s32(_sum4, _b0, _a45, 0); _sum4 = vdotq_laneq_s32(_sum4, _b1, _a45, 1); - _sum5 = vdotq_laneq_s32(_sum5, _b0, _a45, 2); _sum5 = vdotq_laneq_s32(_sum5, _b1, _a45, 3); - _sum6 = vdotq_laneq_s32(_sum6, _b0, _a67, 0); _sum6 = vdotq_laneq_s32(_sum6, _b1, _a67, 1); - _sum7 = vdotq_laneq_s32(_sum7, _b0, _a67, 2); _sum7 = vdotq_laneq_s32(_sum7, _b1, _a67, 3); + _sum0 = vdotq_laneq_s32(_sum0, _b0, _a01, 0); + _sum0 = vdotq_laneq_s32(_sum0, _b1, _a01, 1); + _sum1 = vdotq_laneq_s32(_sum1, _b0, _a01, 2); + _sum1 = vdotq_laneq_s32(_sum1, _b1, _a01, 3); + _sum2 = vdotq_laneq_s32(_sum2, _b0, _a23, 0); + _sum2 = vdotq_laneq_s32(_sum2, _b1, _a23, 1); + _sum3 = vdotq_laneq_s32(_sum3, _b0, _a23, 2); + _sum3 = vdotq_laneq_s32(_sum3, _b1, _a23, 3); + _sum4 = vdotq_laneq_s32(_sum4, _b0, _a45, 0); + _sum4 = vdotq_laneq_s32(_sum4, _b1, _a45, 1); + _sum5 = vdotq_laneq_s32(_sum5, _b0, _a45, 2); + _sum5 = vdotq_laneq_s32(_sum5, _b1, _a45, 3); + _sum6 = vdotq_laneq_s32(_sum6, _b0, _a67, 0); + _sum6 = vdotq_laneq_s32(_sum6, _b1, _a67, 1); + _sum7 = vdotq_laneq_s32(_sum7, _b0, _a67, 2); + _sum7 = vdotq_laneq_s32(_sum7, _b1, _a67, 3); pA += 64; pB += 8; } @@ -1863,10 +1895,14 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de const int8x16_t _b = vcombine_s8(vreinterpret_s8_s32(vld1_dup_s32((const int*)pB)), vdup_n_s8(0)); const int8x16_t _a0 = vld1q_s8(pA); const int8x16_t _a1 = vld1q_s8(pA + 16); - _sum0 = vdotq_laneq_s32(_sum0, _b, _a0, 0); _sum1 = vdotq_laneq_s32(_sum1, _b, _a0, 1); - _sum2 = vdotq_laneq_s32(_sum2, _b, _a0, 2); _sum3 = vdotq_laneq_s32(_sum3, _b, _a0, 3); - _sum4 = vdotq_laneq_s32(_sum4, _b, _a1, 0); _sum5 = vdotq_laneq_s32(_sum5, _b, _a1, 1); - _sum6 = vdotq_laneq_s32(_sum6, _b, _a1, 2); _sum7 = vdotq_laneq_s32(_sum7, _b, _a1, 3); + _sum0 = vdotq_laneq_s32(_sum0, _b, _a0, 0); + _sum1 = vdotq_laneq_s32(_sum1, _b, _a0, 1); + _sum2 = vdotq_laneq_s32(_sum2, _b, _a0, 2); + _sum3 = vdotq_laneq_s32(_sum3, _b, _a0, 3); + _sum4 = vdotq_laneq_s32(_sum4, _b, _a1, 0); + _sum5 = vdotq_laneq_s32(_sum5, _b, _a1, 1); + _sum6 = vdotq_laneq_s32(_sum6, _b, _a1, 2); + _sum7 = vdotq_laneq_s32(_sum7, _b, _a1, 3); pA += 32; pB += 4; } @@ -1892,14 +1928,22 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de { const int8x8_t _b = vset_lane_s8(pB[0], vdup_n_s8(0), 0); const int8x8_t _a = vld1_s8(pA); - const int16x8_t _p0 = vmull_s8(_b, vdup_lane_s8(_a, 0)); const int16x8_t _p1 = vmull_s8(_b, vdup_lane_s8(_a, 1)); - const int16x8_t _p2 = vmull_s8(_b, vdup_lane_s8(_a, 2)); const int16x8_t _p3 = vmull_s8(_b, vdup_lane_s8(_a, 3)); - const int16x8_t _p4 = vmull_s8(_b, vdup_lane_s8(_a, 4)); const int16x8_t _p5 = vmull_s8(_b, vdup_lane_s8(_a, 5)); - const int16x8_t _p6 = vmull_s8(_b, vdup_lane_s8(_a, 6)); const int16x8_t _p7 = vmull_s8(_b, vdup_lane_s8(_a, 7)); - _sum0 = vaddq_s32(_sum0, vmovl_s16(vget_low_s16(_p0))); _sum1 = vaddq_s32(_sum1, vmovl_s16(vget_low_s16(_p1))); - _sum2 = vaddq_s32(_sum2, vmovl_s16(vget_low_s16(_p2))); _sum3 = vaddq_s32(_sum3, vmovl_s16(vget_low_s16(_p3))); - _sum4 = vaddq_s32(_sum4, vmovl_s16(vget_low_s16(_p4))); _sum5 = vaddq_s32(_sum5, vmovl_s16(vget_low_s16(_p5))); - _sum6 = vaddq_s32(_sum6, vmovl_s16(vget_low_s16(_p6))); _sum7 = vaddq_s32(_sum7, vmovl_s16(vget_low_s16(_p7))); + const int16x8_t _p0 = vmull_s8(_b, vdup_lane_s8(_a, 0)); + const int16x8_t _p1 = vmull_s8(_b, vdup_lane_s8(_a, 1)); + const int16x8_t _p2 = vmull_s8(_b, vdup_lane_s8(_a, 2)); + const int16x8_t _p3 = vmull_s8(_b, vdup_lane_s8(_a, 3)); + const int16x8_t _p4 = vmull_s8(_b, vdup_lane_s8(_a, 4)); + const int16x8_t _p5 = vmull_s8(_b, vdup_lane_s8(_a, 5)); + const int16x8_t _p6 = vmull_s8(_b, vdup_lane_s8(_a, 6)); + const int16x8_t _p7 = vmull_s8(_b, vdup_lane_s8(_a, 7)); + _sum0 = vaddq_s32(_sum0, vmovl_s16(vget_low_s16(_p0))); + _sum1 = vaddq_s32(_sum1, vmovl_s16(vget_low_s16(_p1))); + _sum2 = vaddq_s32(_sum2, vmovl_s16(vget_low_s16(_p2))); + _sum3 = vaddq_s32(_sum3, vmovl_s16(vget_low_s16(_p3))); + _sum4 = vaddq_s32(_sum4, vmovl_s16(vget_low_s16(_p4))); + _sum5 = vaddq_s32(_sum5, vmovl_s16(vget_low_s16(_p5))); + _sum6 = vaddq_s32(_sum6, vmovl_s16(vget_low_s16(_p6))); + _sum7 = vaddq_s32(_sum7, vmovl_s16(vget_low_s16(_p7))); pA += 8; pB++; } @@ -1919,14 +1963,22 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB_descales++; } - vst1q_lane_f32(outptr, _fsum0, 0); outptr++; - vst1q_lane_f32(outptr, _fsum1, 0); outptr++; - vst1q_lane_f32(outptr, _fsum2, 0); outptr++; - vst1q_lane_f32(outptr, _fsum3, 0); outptr++; - vst1q_lane_f32(outptr, _fsum4, 0); outptr++; - vst1q_lane_f32(outptr, _fsum5, 0); outptr++; - vst1q_lane_f32(outptr, _fsum6, 0); outptr++; - vst1q_lane_f32(outptr, _fsum7, 0); outptr++; + vst1q_lane_f32(outptr, _fsum0, 0); + outptr++; + vst1q_lane_f32(outptr, _fsum1, 0); + outptr++; + vst1q_lane_f32(outptr, _fsum2, 0); + outptr++; + vst1q_lane_f32(outptr, _fsum3, 0); + outptr++; + vst1q_lane_f32(outptr, _fsum4, 0); + outptr++; + vst1q_lane_f32(outptr, _fsum5, 0); + outptr++; + vst1q_lane_f32(outptr, _fsum6, 0); + outptr++; + vst1q_lane_f32(outptr, _fsum7, 0); + outptr++; } } #endif // __aarch64__ @@ -3652,61 +3704,111 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b if (broadcast_type_C == 4) { const float32x2_t _c = beta == 1.f ? vld1_f32(pC) : vmul_n_f32(vld1_f32(pC), beta); - _out0 = vadd_f32(_out0, _c); _out1 = vadd_f32(_out1, _c); - _out2 = vadd_f32(_out2, _c); _out3 = vadd_f32(_out3, _c); - _out4 = vadd_f32(_out4, _c); _out5 = vadd_f32(_out5, _c); - _out6 = vadd_f32(_out6, _c); _out7 = vadd_f32(_out7, _c); + _out0 = vadd_f32(_out0, _c); + _out1 = vadd_f32(_out1, _c); + _out2 = vadd_f32(_out2, _c); + _out3 = vadd_f32(_out3, _c); + _out4 = vadd_f32(_out4, _c); + _out5 = vadd_f32(_out5, _c); + _out6 = vadd_f32(_out6, _c); + _out7 = vadd_f32(_out7, _c); pC += 2; } } if (alpha != 1.f) { - _out0 = vmul_n_f32(_out0, alpha); _out1 = vmul_n_f32(_out1, alpha); - _out2 = vmul_n_f32(_out2, alpha); _out3 = vmul_n_f32(_out3, alpha); - _out4 = vmul_n_f32(_out4, alpha); _out5 = vmul_n_f32(_out5, alpha); - _out6 = vmul_n_f32(_out6, alpha); _out7 = vmul_n_f32(_out7, alpha); - } - vst1_f32(outptr0, _out0); vst1_f32(outptr1, _out1); - vst1_f32(outptr2, _out2); vst1_f32(outptr3, _out3); - vst1_f32(outptr4, _out4); vst1_f32(outptr5, _out5); - vst1_f32(outptr6, _out6); vst1_f32(outptr7, _out7); - outptr0 += 2; outptr1 += 2; outptr2 += 2; outptr3 += 2; - outptr4 += 2; outptr5 += 2; outptr6 += 2; outptr7 += 2; + _out0 = vmul_n_f32(_out0, alpha); + _out1 = vmul_n_f32(_out1, alpha); + _out2 = vmul_n_f32(_out2, alpha); + _out3 = vmul_n_f32(_out3, alpha); + _out4 = vmul_n_f32(_out4, alpha); + _out5 = vmul_n_f32(_out5, alpha); + _out6 = vmul_n_f32(_out6, alpha); + _out7 = vmul_n_f32(_out7, alpha); + } + vst1_f32(outptr0, _out0); + vst1_f32(outptr1, _out1); + vst1_f32(outptr2, _out2); + vst1_f32(outptr3, _out3); + vst1_f32(outptr4, _out4); + vst1_f32(outptr5, _out5); + vst1_f32(outptr6, _out6); + vst1_f32(outptr7, _out7); + outptr0 += 2; + outptr1 += 2; + outptr2 += 2; + outptr3 += 2; + outptr4 += 2; + outptr5 += 2; + outptr6 += 2; + outptr7 += 2; pp += 16; } for (; jj < max_jj; jj++) { - float f0 = pp[0]; float f1 = pp[1]; float f2 = pp[2]; float f3 = pp[3]; - float f4 = pp[4]; float f5 = pp[5]; float f6 = pp[6]; float f7 = pp[7]; + float f0 = pp[0]; + float f1 = pp[1]; + float f2 = pp[2]; + float f3 = pp[3]; + float f4 = pp[4]; + float f5 = pp[5]; + float f6 = pp[6]; + float f7 = pp[7]; if (pC) { if (broadcast_type_C <= 2) { - f0 += c0; f1 += c1; f2 += c2; f3 += c3; - f4 += c4; f5 += c5; f6 += c6; f7 += c7; + f0 += c0; + f1 += c1; + f2 += c2; + f3 += c3; + f4 += c4; + f5 += c5; + f6 += c6; + f7 += c7; } if (broadcast_type_C == 3) { - f0 += pC[0] * beta; f1 += pC[c_hstep] * beta; - f2 += pC[c_hstep * 2] * beta; f3 += pC[c_hstep * 3] * beta; - f4 += pC[c_hstep * 4] * beta; f5 += pC[c_hstep * 5] * beta; - f6 += pC[c_hstep * 6] * beta; f7 += pC[c_hstep * 7] * beta; + f0 += pC[0] * beta; + f1 += pC[c_hstep] * beta; + f2 += pC[c_hstep * 2] * beta; + f3 += pC[c_hstep * 3] * beta; + f4 += pC[c_hstep * 4] * beta; + f5 += pC[c_hstep * 5] * beta; + f6 += pC[c_hstep * 6] * beta; + f7 += pC[c_hstep * 7] * beta; pC++; } if (broadcast_type_C == 4) { const float c = pC[0] * beta; - f0 += c; f1 += c; f2 += c; f3 += c; - f4 += c; f5 += c; f6 += c; f7 += c; + f0 += c; + f1 += c; + f2 += c; + f3 += c; + f4 += c; + f5 += c; + f6 += c; + f7 += c; pC++; } } - outptr0[0] = f0 * alpha; outptr1[0] = f1 * alpha; - outptr2[0] = f2 * alpha; outptr3[0] = f3 * alpha; - outptr4[0] = f4 * alpha; outptr5[0] = f5 * alpha; - outptr6[0] = f6 * alpha; outptr7[0] = f7 * alpha; - outptr0++; outptr1++; outptr2++; outptr3++; - outptr4++; outptr5++; outptr6++; outptr7++; + outptr0[0] = f0 * alpha; + outptr1[0] = f1 * alpha; + outptr2[0] = f2 * alpha; + outptr3[0] = f3 * alpha; + outptr4[0] = f4 * alpha; + outptr5[0] = f5 * alpha; + outptr6[0] = f6 * alpha; + outptr7[0] = f7 * alpha; + outptr0++; + outptr1++; + outptr2++; + outptr3++; + outptr4++; + outptr5++; + outptr6++; + outptr7++; pp += 8; } } @@ -4826,10 +4928,14 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, { if (broadcast_type_C <= 2) { - _out0 = vadd_f32(_out0, vdup_n_f32(c0)); _out1 = vadd_f32(_out1, vdup_n_f32(c1)); - _out2 = vadd_f32(_out2, vdup_n_f32(c2)); _out3 = vadd_f32(_out3, vdup_n_f32(c3)); - _out4 = vadd_f32(_out4, vdup_n_f32(c4)); _out5 = vadd_f32(_out5, vdup_n_f32(c5)); - _out6 = vadd_f32(_out6, vdup_n_f32(c6)); _out7 = vadd_f32(_out7, vdup_n_f32(c7)); + _out0 = vadd_f32(_out0, vdup_n_f32(c0)); + _out1 = vadd_f32(_out1, vdup_n_f32(c1)); + _out2 = vadd_f32(_out2, vdup_n_f32(c2)); + _out3 = vadd_f32(_out3, vdup_n_f32(c3)); + _out4 = vadd_f32(_out4, vdup_n_f32(c4)); + _out5 = vadd_f32(_out5, vdup_n_f32(c5)); + _out6 = vadd_f32(_out6, vdup_n_f32(c6)); + _out7 = vadd_f32(_out7, vdup_n_f32(c7)); } if (broadcast_type_C == 3) { @@ -4846,24 +4952,34 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, if (broadcast_type_C == 4) { const float32x2_t _c = beta == 1.f ? vld1_f32(pC) : vmul_n_f32(vld1_f32(pC), beta); - _out0 = vadd_f32(_out0, _c); _out1 = vadd_f32(_out1, _c); - _out2 = vadd_f32(_out2, _c); _out3 = vadd_f32(_out3, _c); - _out4 = vadd_f32(_out4, _c); _out5 = vadd_f32(_out5, _c); - _out6 = vadd_f32(_out6, _c); _out7 = vadd_f32(_out7, _c); + _out0 = vadd_f32(_out0, _c); + _out1 = vadd_f32(_out1, _c); + _out2 = vadd_f32(_out2, _c); + _out3 = vadd_f32(_out3, _c); + _out4 = vadd_f32(_out4, _c); + _out5 = vadd_f32(_out5, _c); + _out6 = vadd_f32(_out6, _c); + _out7 = vadd_f32(_out7, _c); pC += 2; } } if (alpha != 1.f) { - _out0 = vmul_n_f32(_out0, alpha); _out1 = vmul_n_f32(_out1, alpha); - _out2 = vmul_n_f32(_out2, alpha); _out3 = vmul_n_f32(_out3, alpha); - _out4 = vmul_n_f32(_out4, alpha); _out5 = vmul_n_f32(_out5, alpha); - _out6 = vmul_n_f32(_out6, alpha); _out7 = vmul_n_f32(_out7, alpha); + _out0 = vmul_n_f32(_out0, alpha); + _out1 = vmul_n_f32(_out1, alpha); + _out2 = vmul_n_f32(_out2, alpha); + _out3 = vmul_n_f32(_out3, alpha); + _out4 = vmul_n_f32(_out4, alpha); + _out5 = vmul_n_f32(_out5, alpha); + _out6 = vmul_n_f32(_out6, alpha); + _out7 = vmul_n_f32(_out7, alpha); } float32x4x2_t _t0 = vuzpq_f32(vcombine_f32(_out0, _out1), vcombine_f32(_out2, _out3)); float32x4x2_t _t1 = vuzpq_f32(vcombine_f32(_out4, _out5), vcombine_f32(_out6, _out7)); - vst1q_f32(outptr0, _t0.val[0]); vst1q_f32(outptr0 + 4, _t1.val[0]); - vst1q_f32(outptr1, _t0.val[1]); vst1q_f32(outptr1 + 4, _t1.val[1]); + vst1q_f32(outptr0, _t0.val[0]); + vst1q_f32(outptr0 + 4, _t1.val[0]); + vst1q_f32(outptr1, _t0.val[1]); + vst1q_f32(outptr1 + 4, _t1.val[1]); outptr0 += out_hstep * 2; outptr1 += out_hstep * 2; pp += 16; diff --git a/src/layer/loongarch/gemm_loongarch.cpp b/src/layer/loongarch/gemm_loongarch.cpp index f0dd51074e8a..aed643974e11 100644 --- a/src/layer/loongarch/gemm_loongarch.cpp +++ b/src/layer/loongarch/gemm_loongarch.cpp @@ -7461,8 +7461,10 @@ static int gemm_BT_loongarch_wq_int8(const Mat& A, const Mat& packed_B, const Ma Mat BT_tile = BT.row_range(j, max_jj); Mat BT_descales_tile = BT_descales.row_range(j, max_jj); gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); - if (output_transpose) transpose_unpack_output_tile_wq_int8(topT_tile, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, M, alpha, beta); - else unpack_output_tile_wq_int8(topT_tile, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, N, alpha, beta); + if (output_transpose) + transpose_unpack_output_tile_wq_int8(topT_tile, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, M, alpha, beta); + else + unpack_output_tile_wq_int8(topT_tile, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, N, alpha, beta); } } } diff --git a/src/layer/loongarch/gemm_wq_int8.h b/src/layer/loongarch/gemm_wq_int8.h index 72825207f51f..a2ca710c2472 100644 --- a/src/layer/loongarch/gemm_wq_int8.h +++ b/src/layer/loongarch/gemm_wq_int8.h @@ -73,7 +73,7 @@ static int pack_B_wq_int8(const Mat& B, const Mat& B_scales, Mat& BT, Mat& BT_de nr = 1; } #if __loongarch_sx - } + } #endif #if __loongarch_asx } @@ -587,7 +587,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales { v *= input_scale_ptr[k]; // preserve multiplication order for consistent rounding - asm volatile("" : "+f"(v)); + asm volatile("" + : "+f"(v)); } outptr0[k] = float2int8(v * scale); } @@ -963,7 +964,8 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A { v *= input_scale_ptr[k]; // preserve multiplication order for consistent rounding - asm volatile("" : "+f"(v)); + asm volatile("" + : "+f"(v)); } outptr0[k] = float2int8(v * scale); } diff --git a/src/layer/mips/gemm_mips.cpp b/src/layer/mips/gemm_mips.cpp index d0a696e4c398..84a72f800a5e 100644 --- a/src/layer/mips/gemm_mips.cpp +++ b/src/layer/mips/gemm_mips.cpp @@ -4635,12 +4635,36 @@ int Gemm_mips::forward_weight_block_quantize_int8(const std::vector& bottom if (!C.empty()) { bool matched = false; - if (C.dims == 1 && C.w == 1) { broadcast_type_C = 0; matched = true; } - if (C.dims == 1 && C.w == M) { broadcast_type_C = 1; matched = true; } - if (C.dims == 1 && C.w == N) { broadcast_type_C = 4; matched = true; } - if (C.dims == 2 && C.w == 1 && C.h == M) { broadcast_type_C = 2; matched = true; } - if (C.dims == 2 && C.w == N && C.h == M) { broadcast_type_C = 3; matched = true; } - if (C.dims == 2 && C.w == N && C.h == 1) { broadcast_type_C = 4; matched = true; } + if (C.dims == 1 && C.w == 1) + { + broadcast_type_C = 0; + matched = true; + } + if (C.dims == 1 && C.w == M) + { + broadcast_type_C = 1; + matched = true; + } + if (C.dims == 1 && C.w == N) + { + broadcast_type_C = 4; + matched = true; + } + if (C.dims == 2 && C.w == 1 && C.h == M) + { + broadcast_type_C = 2; + matched = true; + } + if (C.dims == 2 && C.w == N && C.h == M) + { + broadcast_type_C = 3; + matched = true; + } + if (C.dims == 2 && C.w == N && C.h == 1) + { + broadcast_type_C = 4; + matched = true; + } if (!matched || C.elemsize != 4u || C.elempack != 1) { diff --git a/src/layer/mips/gemm_wq_int8.h b/src/layer/mips/gemm_wq_int8.h index 547c1521c825..a93737caf9ea 100644 --- a/src/layer/mips/gemm_wq_int8.h +++ b/src/layer/mips/gemm_wq_int8.h @@ -540,10 +540,14 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales v31 *= s1; v32 *= s2; v33 *= s3; - asm volatile("" : "+f"(v00), "+f"(v01), "+f"(v02), "+f"(v03)); - asm volatile("" : "+f"(v10), "+f"(v11), "+f"(v12), "+f"(v13)); - asm volatile("" : "+f"(v20), "+f"(v21), "+f"(v22), "+f"(v23)); - asm volatile("" : "+f"(v30), "+f"(v31), "+f"(v32), "+f"(v33)); + asm volatile("" + : "+f"(v00), "+f"(v01), "+f"(v02), "+f"(v03)); + asm volatile("" + : "+f"(v10), "+f"(v11), "+f"(v12), "+f"(v13)); + asm volatile("" + : "+f"(v20), "+f"(v21), "+f"(v22), "+f"(v23)); + asm volatile("" + : "+f"(v30), "+f"(v31), "+f"(v32), "+f"(v33)); } pp[0] = float2int8(v00 * scale0); pp[1] = float2int8(v01 * scale0); @@ -585,8 +589,10 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales v21 *= s1; v30 *= s0; v31 *= s1; - asm volatile("" : "+f"(v00), "+f"(v01), "+f"(v10), "+f"(v11)); - asm volatile("" : "+f"(v20), "+f"(v21), "+f"(v30), "+f"(v31)); + asm volatile("" + : "+f"(v00), "+f"(v01), "+f"(v10), "+f"(v11)); + asm volatile("" + : "+f"(v20), "+f"(v21), "+f"(v30), "+f"(v31)); } pp[0] = float2int8(v00 * scale0); pp[1] = float2int8(v01 * scale0); @@ -613,7 +619,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales v1 *= s; v2 *= s; v3 *= s; - asm volatile("" : "+f"(v0), "+f"(v1), "+f"(v2), "+f"(v3)); + asm volatile("" + : "+f"(v0), "+f"(v1), "+f"(v2), "+f"(v3)); } pp[0] = float2int8(v0 * scale0); pp[1] = float2int8(v1 * scale1); @@ -672,7 +679,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales { v0 *= input_scale_ptr[k]; v1 *= input_scale_ptr[k]; - asm volatile("" : "+f"(v0), "+f"(v1)); + asm volatile("" + : "+f"(v0), "+f"(v1)); } pp[r] = float2int8(v0 * scale0); pp[4 + r] = float2int8(v1 * scale1); @@ -690,7 +698,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales { v0 *= input_scale_ptr[k]; v1 *= input_scale_ptr[k]; - asm volatile("" : "+f"(v0), "+f"(v1)); + asm volatile("" + : "+f"(v0), "+f"(v1)); } pp[r] = float2int8(v0 * scale0); pp[2 + r] = float2int8(v1 * scale1); @@ -707,7 +716,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales { v0 *= input_scale_ptr[k]; v1 *= input_scale_ptr[k]; - asm volatile("" : "+f"(v0), "+f"(v1)); + asm volatile("" + : "+f"(v0), "+f"(v1)); } pp[0] = float2int8(v0 * scale0); pp[1] = float2int8(v1 * scale1); @@ -748,7 +758,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales if (input_scale_ptr) { v0 *= input_scale_ptr[k]; - asm volatile("" : "+f"(v0)); + asm volatile("" + : "+f"(v0)); } pp[r] = float2int8(v0 * scale0); } @@ -763,7 +774,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales if (input_scale_ptr) { v0 *= input_scale_ptr[k]; - asm volatile("" : "+f"(v0)); + asm volatile("" + : "+f"(v0)); } pp[r] = float2int8(v0 * scale0); } @@ -777,7 +789,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales if (input_scale_ptr) { v0 *= input_scale_ptr[k]; - asm volatile("" : "+f"(v0)); + asm volatile("" + : "+f"(v0)); } *pp++ = float2int8(v0 * scale0); } @@ -911,8 +924,12 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A const float s1 = input_scale_ptr ? input_scale_ptr[k0 + kk + 1] : 1.f; v4f32 _p0 = __msa_fmul_w((v4f32)__msa_ld_w(p0, 0), __msa_fill_w_f32(s0)); v4f32 _p1 = __msa_fmul_w((v4f32)__msa_ld_w(p1, 0), __msa_fill_w_f32(s1)); - v16i8 _q0 = float2int8(__msa_fmul_w(_p0, (v4f32){scale0, scale1, scale2, scale3})); - v16i8 _q1 = float2int8(__msa_fmul_w(_p1, (v4f32){scale0, scale1, scale2, scale3})); + v16i8 _q0 = float2int8(__msa_fmul_w(_p0, (v4f32) { + scale0, scale1, scale2, scale3 + })); + v16i8 _q1 = float2int8(__msa_fmul_w(_p1, (v4f32) { + scale0, scale1, scale2, scale3 + })); pp[0] = __msa_copy_s_b(_q0, 0); pp[1] = __msa_copy_s_b(_q1, 0); pp[2] = __msa_copy_s_b(_q0, 1); @@ -923,8 +940,12 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A pp[7] = __msa_copy_s_b(_q1, 3); _p0 = __msa_fmul_w((v4f32)__msa_ld_w(p0 + 4, 0), __msa_fill_w_f32(s0)); _p1 = __msa_fmul_w((v4f32)__msa_ld_w(p1 + 4, 0), __msa_fill_w_f32(s1)); - _q0 = float2int8(__msa_fmul_w(_p0, (v4f32){scale4, scale5, scale6, scale7})); - _q1 = float2int8(__msa_fmul_w(_p1, (v4f32){scale4, scale5, scale6, scale7})); + _q0 = float2int8(__msa_fmul_w(_p0, (v4f32) { + scale4, scale5, scale6, scale7 + })); + _q1 = float2int8(__msa_fmul_w(_p1, (v4f32) { + scale4, scale5, scale6, scale7 + })); pp[8] = __msa_copy_s_b(_q0, 0); pp[9] = __msa_copy_s_b(_q1, 0); pp[10] = __msa_copy_s_b(_q0, 1); @@ -943,8 +964,12 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A const float s = input_scale_ptr ? input_scale_ptr[k] : 1.f; v4f32 _p0 = __msa_fmul_w((v4f32)__msa_ld_w(p0, 0), __msa_fill_w_f32(s)); v4f32 _p1 = __msa_fmul_w((v4f32)__msa_ld_w(p0 + 4, 0), __msa_fill_w_f32(s)); - const v16i8 _q0 = float2int8(__msa_fmul_w(_p0, (v4f32){scale0, scale1, scale2, scale3})); - const v16i8 _q1 = float2int8(__msa_fmul_w(_p1, (v4f32){scale4, scale5, scale6, scale7})); + const v16i8 _q0 = float2int8(__msa_fmul_w(_p0, (v4f32) { + scale0, scale1, scale2, scale3 + })); + const v16i8 _q1 = float2int8(__msa_fmul_w(_p1, (v4f32) { + scale4, scale5, scale6, scale7 + })); ((int*)pp)[0] = __msa_copy_s_w((v4i32)_q0, 0); ((int*)pp)[1] = __msa_copy_s_w((v4i32)_q1, 0); pp += 8; @@ -1045,8 +1070,10 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A v11 *= s1; v21 *= s1; v31 *= s1; - asm volatile("" : "+f"(v00), "+f"(v01), "+f"(v10), "+f"(v11)); - asm volatile("" : "+f"(v20), "+f"(v21), "+f"(v30), "+f"(v31)); + asm volatile("" + : "+f"(v00), "+f"(v01), "+f"(v10), "+f"(v11)); + asm volatile("" + : "+f"(v20), "+f"(v21), "+f"(v30), "+f"(v31)); } pp[0] = float2int8(v00 * scale0); pp[1] = float2int8(v01 * scale0); @@ -1074,7 +1101,8 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A v1 *= s; v2 *= s; v3 *= s; - asm volatile("" : "+f"(v0), "+f"(v1), "+f"(v2), "+f"(v3)); + asm volatile("" + : "+f"(v0), "+f"(v1), "+f"(v2), "+f"(v3)); } pp[0] = float2int8(v0 * scale0); pp[1] = float2int8(v1 * scale1); @@ -1132,7 +1160,8 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A const float s = input_scale_ptr[k]; v0 *= s; v1 *= s; - asm volatile("" : "+f"(v0), "+f"(v1)); + asm volatile("" + : "+f"(v0), "+f"(v1)); } pp[r] = float2int8(v0 * scale0); pp[4 + r] = float2int8(v1 * scale1); @@ -1152,7 +1181,8 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A const float s = input_scale_ptr[k]; v0 *= s; v1 *= s; - asm volatile("" : "+f"(v0), "+f"(v1)); + asm volatile("" + : "+f"(v0), "+f"(v1)); } pp[r] = float2int8(v0 * scale0); pp[2 + r] = float2int8(v1 * scale1); @@ -1171,7 +1201,8 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A const float s = input_scale_ptr[k]; v0 *= s; v1 *= s; - asm volatile("" : "+f"(v0), "+f"(v1)); + asm volatile("" + : "+f"(v0), "+f"(v1)); } pp[0] = float2int8(v0 * scale0); pp[1] = float2int8(v1 * scale1); @@ -1211,7 +1242,8 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A if (input_scale_ptr) { v0 *= input_scale_ptr[k]; - asm volatile("" : "+f"(v0)); + asm volatile("" + : "+f"(v0)); } pp[r] = float2int8(v0 * scale0); } @@ -1226,7 +1258,8 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A if (input_scale_ptr) { v0 *= input_scale_ptr[k]; - asm volatile("" : "+f"(v0)); + asm volatile("" + : "+f"(v0)); } pp[r] = float2int8(v0 * scale0); } @@ -1240,7 +1273,8 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A if (input_scale_ptr) { v0 *= input_scale_ptr[k]; - asm volatile("" : "+f"(v0)); + asm volatile("" + : "+f"(v0)); } *pp++ = float2int8(v0 * scale0); } @@ -2048,11 +2082,21 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB0 += 4; pB1 += 4; } - const v4f32 _descaleA = (v4f32){pAD[0], pAD[1], pAD[0], pAD[1]}; - const v4f32 _descaleB0 = (v4f32){pBD0[0], pBD0[0], pBD0[1], pBD0[1]}; - const v4f32 _descaleB1 = (v4f32){pBD0[2], pBD0[2], pBD0[3], pBD0[3]}; - const v4f32 _descaleB2 = (v4f32){pBD1[0], pBD1[0], pBD1[1], pBD1[1]}; - const v4f32 _descaleB3 = (v4f32){pBD1[2], pBD1[2], pBD1[3], pBD1[3]}; + const v4f32 _descaleA = (v4f32) { + pAD[0], pAD[1], pAD[0], pAD[1] + }; + const v4f32 _descaleB0 = (v4f32) { + pBD0[0], pBD0[0], pBD0[1], pBD0[1] + }; + const v4f32 _descaleB1 = (v4f32) { + pBD0[2], pBD0[2], pBD0[3], pBD0[3] + }; + const v4f32 _descaleB2 = (v4f32) { + pBD1[0], pBD1[0], pBD1[1], pBD1[1] + }; + const v4f32 _descaleB3 = (v4f32) { + pBD1[2], pBD1[2], pBD1[3], pBD1[3] + }; _fsum0 = __msa_fadd_w(_fsum0, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum0), __msa_fmul_w(_descaleA, _descaleB0))); _fsum1 = __msa_fadd_w(_fsum1, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum1), __msa_fmul_w(_descaleA, _descaleB1))); _fsum2 = __msa_fadd_w(_fsum2, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum2), __msa_fmul_w(_descaleA, _descaleB2))); @@ -2126,9 +2170,15 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pA += 2; pB += 4; } - const v4f32 _descaleA = (v4f32){pAD[0], pAD[1], pAD[0], pAD[1]}; - const v4f32 _descaleB0 = (v4f32){pBD[0], pBD[0], pBD[1], pBD[1]}; - const v4f32 _descaleB1 = (v4f32){pBD[2], pBD[2], pBD[3], pBD[3]}; + const v4f32 _descaleA = (v4f32) { + pAD[0], pAD[1], pAD[0], pAD[1] + }; + const v4f32 _descaleB0 = (v4f32) { + pBD[0], pBD[0], pBD[1], pBD[1] + }; + const v4f32 _descaleB1 = (v4f32) { + pBD[2], pBD[2], pBD[3], pBD[3] + }; _fsum0 = __msa_fadd_w(_fsum0, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum0), __msa_fmul_w(_descaleA, _descaleB0))); _fsum1 = __msa_fadd_w(_fsum1, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum1), __msa_fmul_w(_descaleA, _descaleB1))); pAD += 2; @@ -3005,10 +3055,18 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b } if (broadcast_type_C == 3) { - v4f32 _c0 = (v4f32){pC0[0], pC1[0], pC2[0], pC3[0]}; - v4f32 _c4 = (v4f32){pC4[0], pC5[0], pC6[0], pC7[0]}; - v4f32 _c1 = (v4f32){pC0[1], pC1[1], pC2[1], pC3[1]}; - v4f32 _c5 = (v4f32){pC4[1], pC5[1], pC6[1], pC7[1]}; + v4f32 _c0 = (v4f32) { + pC0[0], pC1[0], pC2[0], pC3[0] + }; + v4f32 _c4 = (v4f32) { + pC4[0], pC5[0], pC6[0], pC7[0] + }; + v4f32 _c1 = (v4f32) { + pC0[1], pC1[1], pC2[1], pC3[1] + }; + v4f32 _c5 = (v4f32) { + pC4[1], pC5[1], pC6[1], pC7[1] + }; if (beta != 1.f) { const v4f32 _beta = __msa_fill_w_f32(beta); @@ -3099,8 +3157,12 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b } if (broadcast_type_C == 3) { - v4f32 _c0 = (v4f32){pC0[0], pC1[0], pC2[0], pC3[0]}; - v4f32 _c4 = (v4f32){pC4[0], pC5[0], pC6[0], pC7[0]}; + v4f32 _c0 = (v4f32) { + pC0[0], pC1[0], pC2[0], pC3[0] + }; + v4f32 _c4 = (v4f32) { + pC4[0], pC5[0], pC6[0], pC7[0] + }; if (beta != 1.f) { const v4f32 _beta = __msa_fill_w_f32(beta); @@ -3395,8 +3457,12 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b } if (broadcast_type_C == 3) { - v4f32 _c0 = (v4f32){pC0[0], pC1[0], pC2[0], pC3[0]}; - v4f32 _c1 = (v4f32){pC0[1], pC1[1], pC2[1], pC3[1]}; + v4f32 _c0 = (v4f32) { + pC0[0], pC1[0], pC2[0], pC3[0] + }; + v4f32 _c1 = (v4f32) { + pC0[1], pC1[1], pC2[1], pC3[1] + }; if (beta != 1.f) { v4f32 _beta = __msa_fill_w_f32(beta); @@ -3460,7 +3526,9 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b } if (broadcast_type_C == 3) { - v4f32 _c0 = (v4f32){pC0[0], pC1[0], pC2[0], pC3[0]}; + v4f32 _c0 = (v4f32) { + pC0[0], pC1[0], pC2[0], pC3[0] + }; if (beta != 1.f) _c0 = __msa_fmul_w(_c0, __msa_fill_w_f32(beta)); _f0 = __msa_fadd_w(_f0, _c0); @@ -3675,10 +3743,14 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b if (pC) { if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) - _f = __msa_fadd_w(_f, (v4f32){c0, c1, c0, c1}); + _f = __msa_fadd_w(_f, (v4f32) { + c0, c1, c0, c1 + }); if (broadcast_type_C == 3) { - v4f32 _c = (v4f32){pC0[0], pC1[0], pC0[1], pC1[1]}; + v4f32 _c = (v4f32) { + pC0[0], pC1[0], pC0[1], pC1[1] + }; if (beta != 1.f) _c = __msa_fmul_w(_c, __msa_fill_w_f32(beta)); _f = __msa_fadd_w(_f, _c); @@ -3692,7 +3764,9 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b cc0 *= beta; cc1 *= beta; } - _f = __msa_fadd_w(_f, (v4f32){cc0, cc0, cc1, cc1}); + _f = __msa_fadd_w(_f, (v4f32) { + cc0, cc0, cc1, cc1 + }); } } @@ -4135,14 +4209,30 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, } if (broadcast_type_C == 3) { - v4f32 _cl0 = (v4f32){pC0[0], pC1[0], pC2[0], pC3[0]}; - v4f32 _ch0 = (v4f32){pC4[0], pC5[0], pC6[0], pC7[0]}; - v4f32 _cl1 = (v4f32){pC0[1], pC1[1], pC2[1], pC3[1]}; - v4f32 _ch1 = (v4f32){pC4[1], pC5[1], pC6[1], pC7[1]}; - v4f32 _cl2 = (v4f32){pC0[2], pC1[2], pC2[2], pC3[2]}; - v4f32 _ch2 = (v4f32){pC4[2], pC5[2], pC6[2], pC7[2]}; - v4f32 _cl3 = (v4f32){pC0[3], pC1[3], pC2[3], pC3[3]}; - v4f32 _ch3 = (v4f32){pC4[3], pC5[3], pC6[3], pC7[3]}; + v4f32 _cl0 = (v4f32) { + pC0[0], pC1[0], pC2[0], pC3[0] + }; + v4f32 _ch0 = (v4f32) { + pC4[0], pC5[0], pC6[0], pC7[0] + }; + v4f32 _cl1 = (v4f32) { + pC0[1], pC1[1], pC2[1], pC3[1] + }; + v4f32 _ch1 = (v4f32) { + pC4[1], pC5[1], pC6[1], pC7[1] + }; + v4f32 _cl2 = (v4f32) { + pC0[2], pC1[2], pC2[2], pC3[2] + }; + v4f32 _ch2 = (v4f32) { + pC4[2], pC5[2], pC6[2], pC7[2] + }; + v4f32 _cl3 = (v4f32) { + pC0[3], pC1[3], pC2[3], pC3[3] + }; + v4f32 _ch3 = (v4f32) { + pC4[3], pC5[3], pC6[3], pC7[3] + }; if (beta != 1.f) { const v4f32 _beta = __msa_fill_w_f32(beta); @@ -4234,10 +4324,18 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, } if (broadcast_type_C == 3) { - v4f32 _c0 = (v4f32){pC0[0], pC1[0], pC2[0], pC3[0]}; - v4f32 _c4 = (v4f32){pC4[0], pC5[0], pC6[0], pC7[0]}; - v4f32 _c1 = (v4f32){pC0[1], pC1[1], pC2[1], pC3[1]}; - v4f32 _c5 = (v4f32){pC4[1], pC5[1], pC6[1], pC7[1]}; + v4f32 _c0 = (v4f32) { + pC0[0], pC1[0], pC2[0], pC3[0] + }; + v4f32 _c4 = (v4f32) { + pC4[0], pC5[0], pC6[0], pC7[0] + }; + v4f32 _c1 = (v4f32) { + pC0[1], pC1[1], pC2[1], pC3[1] + }; + v4f32 _c5 = (v4f32) { + pC4[1], pC5[1], pC6[1], pC7[1] + }; if (beta != 1.f) { const v4f32 _beta = __msa_fill_w_f32(beta); @@ -4309,8 +4407,12 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, } if (broadcast_type_C == 3) { - v4f32 _c0 = (v4f32){pC0[0], pC1[0], pC2[0], pC3[0]}; - v4f32 _c4 = (v4f32){pC4[0], pC5[0], pC6[0], pC7[0]}; + v4f32 _c0 = (v4f32) { + pC0[0], pC1[0], pC2[0], pC3[0] + }; + v4f32 _c4 = (v4f32) { + pC4[0], pC5[0], pC6[0], pC7[0] + }; if (beta != 1.f) { const v4f32 _beta = __msa_fill_w_f32(beta); @@ -4580,8 +4682,12 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, } if (broadcast_type_C == 3) { - v4f32 _c0 = (v4f32){pC0[0], pC1[0], pC2[0], pC3[0]}; - v4f32 _c1 = (v4f32){pC0[1], pC1[1], pC2[1], pC3[1]}; + v4f32 _c0 = (v4f32) { + pC0[0], pC1[0], pC2[0], pC3[0] + }; + v4f32 _c1 = (v4f32) { + pC0[1], pC1[1], pC2[1], pC3[1] + }; if (beta != 1.f) { v4f32 _beta = __msa_fill_w_f32(beta); @@ -4637,7 +4743,9 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, } if (broadcast_type_C == 3) { - v4f32 _c0 = (v4f32){pC0[0], pC1[0], pC2[0], pC3[0]}; + v4f32 _c0 = (v4f32) { + pC0[0], pC1[0], pC2[0], pC3[0] + }; if (beta != 1.f) _c0 = __msa_fmul_w(_c0, __msa_fill_w_f32(beta)); _f0 = __msa_fadd_w(_f0, _c0); @@ -4710,7 +4818,9 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, { if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) { - v4f32 _c = (v4f32){c0, c1, c0, c1}; + v4f32 _c = (v4f32) { + c0, c1, c0, c1 + }; _f0 = __msa_fadd_w(_f0, _c); _f1 = __msa_fadd_w(_f1, _c); _f2 = __msa_fadd_w(_f2, _c); @@ -4718,10 +4828,18 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, } if (broadcast_type_C == 3) { - v4f32 _c0 = (v4f32){pC0[0], pC1[0], pC0[1], pC1[1]}; - v4f32 _c1 = (v4f32){pC0[2], pC1[2], pC0[3], pC1[3]}; - v4f32 _c2 = (v4f32){pC0[4], pC1[4], pC0[5], pC1[5]}; - v4f32 _c3 = (v4f32){pC0[6], pC1[6], pC0[7], pC1[7]}; + v4f32 _c0 = (v4f32) { + pC0[0], pC1[0], pC0[1], pC1[1] + }; + v4f32 _c1 = (v4f32) { + pC0[2], pC1[2], pC0[3], pC1[3] + }; + v4f32 _c2 = (v4f32) { + pC0[4], pC1[4], pC0[5], pC1[5] + }; + v4f32 _c3 = (v4f32) { + pC0[6], pC1[6], pC0[7], pC1[7] + }; if (beta != 1.f) { v4f32 _beta = __msa_fill_w_f32(beta); @@ -4756,10 +4874,18 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, c06 *= beta; c07 *= beta; } - _f0 = __msa_fadd_w(_f0, (v4f32){c00, c00, c01, c01}); - _f1 = __msa_fadd_w(_f1, (v4f32){c02, c02, c03, c03}); - _f2 = __msa_fadd_w(_f2, (v4f32){c04, c04, c05, c05}); - _f3 = __msa_fadd_w(_f3, (v4f32){c06, c06, c07, c07}); + _f0 = __msa_fadd_w(_f0, (v4f32) { + c00, c00, c01, c01 + }); + _f1 = __msa_fadd_w(_f1, (v4f32) { + c02, c02, c03, c03 + }); + _f2 = __msa_fadd_w(_f2, (v4f32) { + c04, c04, c05, c05 + }); + _f3 = __msa_fadd_w(_f3, (v4f32) { + c06, c06, c07, c07 + }); } } @@ -4799,14 +4925,20 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, { if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) { - v4f32 _c = (v4f32){c0, c1, c0, c1}; + v4f32 _c = (v4f32) { + c0, c1, c0, c1 + }; _f0 = __msa_fadd_w(_f0, _c); _f1 = __msa_fadd_w(_f1, _c); } if (broadcast_type_C == 3) { - v4f32 _c0 = (v4f32){pC0[0], pC1[0], pC0[1], pC1[1]}; - v4f32 _c1 = (v4f32){pC0[2], pC1[2], pC0[3], pC1[3]}; + v4f32 _c0 = (v4f32) { + pC0[0], pC1[0], pC0[1], pC1[1] + }; + v4f32 _c1 = (v4f32) { + pC0[2], pC1[2], pC0[3], pC1[3] + }; if (beta != 1.f) { v4f32 _beta = __msa_fill_w_f32(beta); @@ -4829,8 +4961,12 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, c02 *= beta; c03 *= beta; } - _f0 = __msa_fadd_w(_f0, (v4f32){c00, c00, c01, c01}); - _f1 = __msa_fadd_w(_f1, (v4f32){c02, c02, c03, c03}); + _f0 = __msa_fadd_w(_f0, (v4f32) { + c00, c00, c01, c01 + }); + _f1 = __msa_fadd_w(_f1, (v4f32) { + c02, c02, c03, c03 + }); } } @@ -4862,10 +4998,14 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, if (pC) { if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) - _f = __msa_fadd_w(_f, (v4f32){c0, c1, c0, c1}); + _f = __msa_fadd_w(_f, (v4f32) { + c0, c1, c0, c1 + }); if (broadcast_type_C == 3) { - v4f32 _c = (v4f32){pC0[0], pC1[0], pC0[1], pC1[1]}; + v4f32 _c = (v4f32) { + pC0[0], pC1[0], pC0[1], pC1[1] + }; if (beta != 1.f) _c = __msa_fmul_w(_c, __msa_fill_w_f32(beta)); _f = __msa_fadd_w(_f, _c); @@ -4879,7 +5019,9 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, cc0 *= beta; cc1 *= beta; } - _f = __msa_fadd_w(_f, (v4f32){cc0, cc0, cc1, cc1}); + _f = __msa_fadd_w(_f, (v4f32) { + cc0, cc0, cc1, cc1 + }); } } diff --git a/src/layer/riscv/gemm_wq_int8.h b/src/layer/riscv/gemm_wq_int8.h index 2d585acbb5b8..9edd10147703 100644 --- a/src/layer/riscv/gemm_wq_int8.h +++ b/src/layer/riscv/gemm_wq_int8.h @@ -328,7 +328,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales { v0 *= input_scale_ptr[k]; v1 *= input_scale_ptr[k]; - asm volatile("" : "+f"(v0), "+f"(v1)); + asm volatile("" + : "+f"(v0), "+f"(v1)); } pp[0] = float2int8(v0 * scale0); pp[1] = float2int8(v1 * scale1); @@ -395,14 +396,14 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales if (input_scale_ptr) { v *= input_scale_ptr[k]; - asm volatile("" : "+f"(v)); + asm volatile("" + : "+f"(v)); } *pp++ = float2int8(v * scale); } #endif } } - } // K-major, row-interleaved MR-packn/MR2/MR1 @@ -543,7 +544,8 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A const float s = input_scale_ptr[k]; v0 *= s; v1 *= s; - asm volatile("" : "+f"(v0), "+f"(v1)); + asm volatile("" + : "+f"(v0), "+f"(v1)); } pp[0] = float2int8(v0 * scale0); pp[1] = float2int8(v1 * scale1); @@ -609,7 +611,8 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A if (input_scale_ptr) { v *= input_scale_ptr[k]; - asm volatile("" : "+f"(v)); + asm volatile("" + : "+f"(v)); } *pp++ = float2int8(v * scale); } diff --git a/src/layer/riscv/multiheadattention_riscv.cpp b/src/layer/riscv/multiheadattention_riscv.cpp index af9d641f66fb..b5e3123659e0 100644 --- a/src/layer/riscv/multiheadattention_riscv.cpp +++ b/src/layer/riscv/multiheadattention_riscv.cpp @@ -918,4 +918,3 @@ int MultiHeadAttention_riscv::forward(const std::vector& bottom_blobs, std: } } // namespace ncnn - diff --git a/src/layer/x86/gemm_wq_int8.h b/src/layer/x86/gemm_wq_int8.h index 1a4e78d0e6cc..f098e6dbcc08 100644 --- a/src/layer/x86/gemm_wq_int8.h +++ b/src/layer/x86/gemm_wq_int8.h @@ -170,8 +170,8 @@ static void pack_B_tile_wq_int8(const Mat& B, const Mat& B_scales, unsigned char _mm_storel_epi64((__m128i*)(pp + 8), _p23); pp += 16; } -#endif // __AVX512VNNI__ || __AVXVNNI__ - // K2/K1 are always signed and compact, including classic VNNI. +#endif // __AVX512VNNI__ || __AVXVNNI__ \ +// K2/K1 are always signed and compact, including classic VNNI. for (; kk + 1 < max_kk; kk += 2) { const signed char* p0 = B.row(j + jj) + k0 + kk; @@ -235,8 +235,8 @@ static void pack_B_tile_wq_int8(const Mat& B, const Mat& B_scales, unsigned char pp += 8; } #endif // __SSE2__ -#endif // __AVX512VNNI__ || __AVXVNNI__ - // K2/K1 are always signed and compact, including classic VNNI. +#endif // __AVX512VNNI__ || __AVXVNNI__ \ +// K2/K1 are always signed and compact, including classic VNNI. #if __SSE2__ for (; kk + 1 < max_kk; kk += 2) { @@ -289,8 +289,8 @@ static void pack_B_tile_wq_int8(const Mat& B, const Mat& B_scales, unsigned char *(int*)pp = *(const int*)p0; pp += 4; } -#endif // __AVX512VNNI__ || __AVXVNNI__ - // K2/K1 are always signed and compact, including classic VNNI. +#endif // __AVX512VNNI__ || __AVXVNNI__ \ +// K2/K1 are always signed and compact, including classic VNNI. for (; kk + 1 < max_kk; kk += 2) { const signed char* p0 = B.row(j + jj) + k0 + kk; @@ -552,8 +552,10 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales _pe = _mm512_mul_ps(_pe, _s); _pf = _mm512_mul_ps(_pf, _s); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+x"(_p0), "+x"(_p1), "+x"(_p2), "+x"(_p3), "+x"(_p4), "+x"(_p5), "+x"(_p6), "+x"(_p7)); - asm volatile("" : "+x"(_p8), "+x"(_p9), "+x"(_pa), "+x"(_pb), "+x"(_pc), "+x"(_pd), "+x"(_pe), "+x"(_pf)); + asm volatile("" + : "+x"(_p0), "+x"(_p1), "+x"(_p2), "+x"(_p3), "+x"(_p4), "+x"(_p5), "+x"(_p6), "+x"(_p7)); + asm volatile("" + : "+x"(_p8), "+x"(_p9), "+x"(_pa), "+x"(_pb), "+x"(_pc), "+x"(_pd), "+x"(_pe), "+x"(_pf)); #else volatile __m512 _p0_ordered = _p0; volatile __m512 _p1_ordered = _p1; @@ -696,8 +698,10 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales _pe = _mm_mul_ps(_pe, _s); _pf = _mm_mul_ps(_pf, _s); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+x"(_p0), "+x"(_p1), "+x"(_p2), "+x"(_p3), "+x"(_p4), "+x"(_p5), "+x"(_p6), "+x"(_p7)); - asm volatile("" : "+x"(_p8), "+x"(_p9), "+x"(_pa), "+x"(_pb), "+x"(_pc), "+x"(_pd), "+x"(_pe), "+x"(_pf)); + asm volatile("" + : "+x"(_p0), "+x"(_p1), "+x"(_p2), "+x"(_p3), "+x"(_p4), "+x"(_p5), "+x"(_p6), "+x"(_p7)); + asm volatile("" + : "+x"(_p8), "+x"(_p9), "+x"(_pa), "+x"(_pb), "+x"(_pc), "+x"(_pd), "+x"(_pe), "+x"(_pf)); #else volatile __m128 _p0_ordered = _p0; volatile __m128 _p1_ordered = _p1; @@ -778,7 +782,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales _p0 = _mm512_mul_ps(_p0, _mm512_set1_ps(input_scale_ptr[k0 + kk])); _p1 = _mm512_mul_ps(_p1, _mm512_set1_ps(input_scale_ptr[k0 + kk + 1])); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+x"(_p0), "+x"(_p1)); + asm volatile("" + : "+x"(_p0), "+x"(_p1)); #else volatile __m512 _p0_ordered = _p0; volatile __m512 _p1_ordered = _p1; @@ -799,7 +804,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales { _p = _mm512_mul_ps(_p, _mm512_set1_ps(input_scale_ptr[k0 + kk])); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+x"(_p)); + asm volatile("" + : "+x"(_p)); #else volatile __m512 _p_ordered = _p; _p = _p_ordered; @@ -923,7 +929,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales _p6 = _mm_mul_ps(_p6, _s); _p7 = _mm_mul_ps(_p7, _s); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+x"(_p0), "+x"(_p1), "+x"(_p2), "+x"(_p3), "+x"(_p4), "+x"(_p5), "+x"(_p6), "+x"(_p7)); + asm volatile("" + : "+x"(_p0), "+x"(_p1), "+x"(_p2), "+x"(_p3), "+x"(_p4), "+x"(_p5), "+x"(_p6), "+x"(_p7)); #else volatile __m128 _p0_ordered = _p0; volatile __m128 _p1_ordered = _p1; @@ -997,7 +1004,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales _p0 = _mm256_mul_ps(_p0, _mm256_set1_ps(input_scale_ptr[k0 + kk])); _p1 = _mm256_mul_ps(_p1, _mm256_set1_ps(input_scale_ptr[k0 + kk + 1])); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+x"(_p0), "+x"(_p1)); + asm volatile("" + : "+x"(_p0), "+x"(_p1)); #else volatile __m256 _p0_ordered = _p0; volatile __m256 _p1_ordered = _p1; @@ -1022,7 +1030,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales { _p = _mm256_mul_ps(_p, _mm256_set1_ps(input_scale_ptr[k0 + kk])); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+x"(_p)); + asm volatile("" + : "+x"(_p)); #else volatile __m256 _p_ordered = _p; _p = _p_ordered; @@ -1115,7 +1124,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales _p2 = _mm_mul_ps(_p2, _s); _p3 = _mm_mul_ps(_p3, _s); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+x"(_p0), "+x"(_p1), "+x"(_p2), "+x"(_p3)); + asm volatile("" + : "+x"(_p0), "+x"(_p1), "+x"(_p2), "+x"(_p3)); #else volatile __m128 _p0_ordered = _p0; volatile __m128 _p1_ordered = _p1; @@ -1160,7 +1170,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales _p0 = _mm_mul_ps(_p0, _mm_set1_ps(input_scale_ptr[k0 + kk])); _p1 = _mm_mul_ps(_p1, _mm_set1_ps(input_scale_ptr[k0 + kk + 1])); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+x"(_p0), "+x"(_p1)); + asm volatile("" + : "+x"(_p0), "+x"(_p1)); #else volatile __m128 _p0_ordered = _p0; volatile __m128 _p1_ordered = _p1; @@ -1180,7 +1191,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales { _p = _mm_mul_ps(_p, _mm_set1_ps(input_scale_ptr[k0 + kk])); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+x"(_p)); + asm volatile("" + : "+x"(_p)); #else volatile __m128 _p_ordered = _p; _p = _p_ordered; @@ -1254,7 +1266,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales _p0 = _mm_mul_ps(_p0, _s); _p1 = _mm_mul_ps(_p1, _s); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+x"(_p0), "+x"(_p1)); + asm volatile("" + : "+x"(_p0), "+x"(_p1)); #else volatile __m128 _p0_ordered = _p0; volatile __m128 _p1_ordered = _p1; @@ -1291,7 +1304,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales _p0 = _mm_mul_ps(_p0, _mm_set1_ps(input_scale_ptr[k0 + kk])); _p1 = _mm_mul_ps(_p1, _mm_set1_ps(input_scale_ptr[k0 + kk + 1])); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+x"(_p0), "+x"(_p1)); + asm volatile("" + : "+x"(_p0), "+x"(_p1)); #else volatile __m128 _p0_ordered = _p0; volatile __m128 _p1_ordered = _p1; @@ -1311,7 +1325,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales { _p = _mm_mul_ps(_p, _mm_set1_ps(input_scale_ptr[k0 + kk])); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+x"(_p)); + asm volatile("" + : "+x"(_p)); #else volatile __m128 _p_ordered = _p; _p = _p_ordered; @@ -1474,7 +1489,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales { _p = _mm512_mul_ps(_p, _mm512_loadu_ps(input_scale_ptr + k0 + kk)); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+x"(_p)); + asm volatile("" + : "+x"(_p)); #else volatile __m512 _p_ordered = _p; _p = _p_ordered; @@ -1499,7 +1515,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales { _p = _mm256_mul_ps(_p, _mm256_loadu_ps(input_scale_ptr + k0 + kk)); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+x"(_p)); + asm volatile("" + : "+x"(_p)); #else volatile __m256 _p_ordered = _p; _p = _p_ordered; @@ -1527,7 +1544,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales { _p = _mm_mul_ps(_p, _mm_loadu_ps(input_scale_ptr + k0 + kk)); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+x"(_p)); + asm volatile("" + : "+x"(_p)); #else volatile __m128 _p_ordered = _p; _p = _p_ordered; @@ -1557,7 +1575,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales { v *= input_scale_ptr[k0 + kk]; #if NCNN_GNU_INLINE_ASM && __SSE2__ - asm volatile("" : "+x"(v)); + asm volatile("" + : "+x"(v)); #else volatile float v_ordered = v; v = v_ordered; @@ -1664,7 +1683,8 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A _p2 = _mm512_mul_ps(_p2, _mm512_set1_ps(input_scale_ptr[k0 + kk + 2])); _p3 = _mm512_mul_ps(_p3, _mm512_set1_ps(input_scale_ptr[k0 + kk + 3])); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+x"(_p0), "+x"(_p1), "+x"(_p2), "+x"(_p3)); + asm volatile("" + : "+x"(_p0), "+x"(_p1), "+x"(_p2), "+x"(_p3)); #else volatile __m512 _p0_ordered = _p0; volatile __m512 _p1_ordered = _p1; @@ -1703,7 +1723,8 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A _p0 = _mm512_mul_ps(_p0, _mm512_set1_ps(input_scale_ptr[k0 + kk])); _p1 = _mm512_mul_ps(_p1, _mm512_set1_ps(input_scale_ptr[k0 + kk + 1])); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+x"(_p0), "+x"(_p1)); + asm volatile("" + : "+x"(_p0), "+x"(_p1)); #else volatile __m512 _p0_ordered = _p0; volatile __m512 _p1_ordered = _p1; @@ -1724,7 +1745,8 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A { _p = _mm512_mul_ps(_p, _mm512_set1_ps(input_scale_ptr[k0 + kk])); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+x"(_p)); + asm volatile("" + : "+x"(_p)); #else volatile __m512 _p_ordered = _p; _p = _p_ordered; @@ -1784,7 +1806,8 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A _p2 = _mm256_mul_ps(_p2, _mm256_set1_ps(input_scale_ptr[k0 + kk + 2])); _p3 = _mm256_mul_ps(_p3, _mm256_set1_ps(input_scale_ptr[k0 + kk + 3])); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+x"(_p0), "+x"(_p1), "+x"(_p2), "+x"(_p3)); + asm volatile("" + : "+x"(_p0), "+x"(_p1), "+x"(_p2), "+x"(_p3)); #else volatile __m256 _p0_ordered = _p0; volatile __m256 _p1_ordered = _p1; @@ -1835,7 +1858,8 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A _p0 = _mm256_mul_ps(_p0, _mm256_set1_ps(input_scale_ptr[k0 + kk])); _p1 = _mm256_mul_ps(_p1, _mm256_set1_ps(input_scale_ptr[k0 + kk + 1])); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+x"(_p0), "+x"(_p1)); + asm volatile("" + : "+x"(_p0), "+x"(_p1)); #else volatile __m256 _p0_ordered = _p0; volatile __m256 _p1_ordered = _p1; @@ -1858,7 +1882,8 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A { _p = _mm256_mul_ps(_p, _mm256_set1_ps(input_scale_ptr[k0 + kk])); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+x"(_p)); + asm volatile("" + : "+x"(_p)); #else volatile __m256 _p_ordered = _p; _p = _p_ordered; @@ -1919,7 +1944,8 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A _p2 = _mm_mul_ps(_p2, _mm_set1_ps(input_scale_ptr[k0 + kk + 2])); _p3 = _mm_mul_ps(_p3, _mm_set1_ps(input_scale_ptr[k0 + kk + 3])); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+x"(_p0), "+x"(_p1), "+x"(_p2), "+x"(_p3)); + asm volatile("" + : "+x"(_p0), "+x"(_p1), "+x"(_p2), "+x"(_p3)); #else volatile __m128 _p0_ordered = _p0; volatile __m128 _p1_ordered = _p1; @@ -1964,7 +1990,8 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A _p0 = _mm_mul_ps(_p0, _mm_set1_ps(input_scale_ptr[k0 + kk])); _p1 = _mm_mul_ps(_p1, _mm_set1_ps(input_scale_ptr[k0 + kk + 1])); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+x"(_p0), "+x"(_p1)); + asm volatile("" + : "+x"(_p0), "+x"(_p1)); #else volatile __m128 _p0_ordered = _p0; volatile __m128 _p1_ordered = _p1; @@ -1984,7 +2011,8 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A { _p = _mm_mul_ps(_p, _mm_set1_ps(input_scale_ptr[k0 + kk])); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+x"(_p)); + asm volatile("" + : "+x"(_p)); #else volatile __m128 _p_ordered = _p; _p = _p_ordered; @@ -2042,7 +2070,8 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A _p2 = _mm_mul_ps(_p2, _mm_set1_ps(input_scale_ptr[k0 + kk + 2])); _p3 = _mm_mul_ps(_p3, _mm_set1_ps(input_scale_ptr[k0 + kk + 3])); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+x"(_p0), "+x"(_p1), "+x"(_p2), "+x"(_p3)); + asm volatile("" + : "+x"(_p0), "+x"(_p1), "+x"(_p2), "+x"(_p3)); #else volatile __m128 _p0_ordered = _p0; volatile __m128 _p1_ordered = _p1; @@ -2087,7 +2116,8 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A _p0 = _mm_mul_ps(_p0, _mm_set1_ps(input_scale_ptr[k0 + kk])); _p1 = _mm_mul_ps(_p1, _mm_set1_ps(input_scale_ptr[k0 + kk + 1])); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+x"(_p0), "+x"(_p1)); + asm volatile("" + : "+x"(_p0), "+x"(_p1)); #else volatile __m128 _p0_ordered = _p0; volatile __m128 _p1_ordered = _p1; @@ -2107,7 +2137,8 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A { _p = _mm_mul_ps(_p, _mm_set1_ps(input_scale_ptr[k0 + kk])); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+x"(_p)); + asm volatile("" + : "+x"(_p)); #else volatile __m128 _p_ordered = _p; _p = _p_ordered; @@ -2288,7 +2319,8 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A { _p = _mm512_mul_ps(_p, _mm512_loadu_ps(input_scale_ptr + k0 + kk)); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+x"(_p)); + asm volatile("" + : "+x"(_p)); #else volatile __m512 _p_ordered = _p; _p = _p_ordered; @@ -2314,7 +2346,8 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A { _p = _mm256_mul_ps(_p, _mm256_loadu_ps(input_scale_ptr + k0 + kk)); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+x"(_p)); + asm volatile("" + : "+x"(_p)); #else volatile __m256 _p_ordered = _p; _p = _p_ordered; @@ -2343,7 +2376,8 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A { _p = _mm_mul_ps(_p, _mm_loadu_ps(input_scale_ptr + k0 + kk)); #if NCNN_GNU_INLINE_ASM - asm volatile("" : "+x"(_p)); + asm volatile("" + : "+x"(_p)); #else volatile __m128 _p_ordered = _p; _p = _p_ordered; @@ -2374,7 +2408,8 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A { v *= input_scale_ptr[k]; #if NCNN_GNU_INLINE_ASM && __SSE2__ - asm volatile("" : "+x"(v)); + asm volatile("" + : "+x"(v)); #else volatile float v_ordered = v; v = v_ordered; @@ -4621,7 +4656,6 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _mm256_storeu_ps(p0 + out_hstep * 14, _mm512_castps512_ps256(_f7)); _mm256_storeu_ps(p0 + out_hstep * 15, _mm512_extractf32x8_ps(_f7, 1)); p0 += 8; - } for (; jj + 3 < max_jj; jj += 4) { @@ -4710,7 +4744,6 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _mm_storeu_ps(p0 + out_hstep * 14, _mm512_extractf32x4_ps(_f3, 2)); _mm_storeu_ps(p0 + out_hstep * 15, _mm512_extractf32x4_ps(_f3, 3)); p0 += 4; - } #endif // defined(__x86_64__) || defined(_M_X64) for (; jj + 1 < max_jj; jj += 2) @@ -4799,7 +4832,6 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _mm_storeh_pi((__m64*)(p0 + out_hstep * 15), _r); } p0 += 2; - } for (; jj < max_jj; jj++) { @@ -4835,7 +4867,6 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _vindex = _mm512_mullo_epi32(_vindex, _mm512_set1_epi32((int)out_hstep)); _mm512_i32scatter_ps(p0, _vindex, _f0, sizeof(float)); p0++; - } } #endif // __AVX512F__ @@ -6658,7 +6689,6 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _mm512_storeu_ps(p0 + out_hstep * 6, _f6); _mm512_storeu_ps(p0 + out_hstep * 7, _f7); p0 += out_hstep * 8; - } for (; jj + 3 < max_jj; jj += 4) { @@ -6734,7 +6764,6 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _mm512_storeu_ps(p0 + out_hstep * 2, _f2); _mm512_storeu_ps(p0 + out_hstep * 3, _f3); p0 += out_hstep * 4; - } #endif // defined(__x86_64__) || defined(_M_X64) for (; jj + 1 < max_jj; jj += 2) @@ -6784,7 +6813,6 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _mm512_storeu_ps(p0, _f0); _mm512_storeu_ps(p0 + out_hstep, _f1); p0 += out_hstep * 2; - } for (; jj < max_jj; jj++) { @@ -6818,7 +6846,6 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, } _mm512_storeu_ps(p0, _f0); p0 += out_hstep; - } } #endif // __AVX512F__ From 95ce491f28eca0ffe9ffa074b7804fcbd4c97d34 Mon Sep 17 00:00:00 2001 From: nihui Date: Mon, 20 Jul 2026 14:44:55 +0800 Subject: [PATCH 3/4] w --- src/layer/arm/gemm_arm.cpp | 68 +- src/layer/arm/gemm_arm_asimddp.cpp | 4 +- src/layer/arm/gemm_arm_i8mm.cpp | 4 +- src/layer/arm/gemm_arm_svei8mm.cpp | 4 +- src/layer/arm/gemm_wq_int8.h | 3596 +++++++++++++++--------- src/layer/gemm.cpp | 7 +- src/layer/loongarch/gemm_loongarch.cpp | 63 +- src/layer/loongarch/gemm_wq_int8.h | 3248 +++++++++++++-------- src/layer/mips/gemm_mips.cpp | 62 +- src/layer/mips/gemm_mips_mmi.cpp | 12 +- src/layer/mips/gemm_wq_int8.h | 2559 ++++++++++------- src/layer/multiheadattention.cpp | 7 +- src/layer/riscv/gemm_riscv.cpp | 62 +- src/layer/riscv/gemm_wq_int8.h | 1781 ++++++++---- src/layer/x86/gemm_wq_int8.h | 2992 +++++++++----------- src/layer/x86/gemm_x86.cpp | 125 +- src/layer/x86/gemm_x86_avx2.cpp | 4 +- src/layer/x86/gemm_x86_avx512vnni.cpp | 4 +- src/layer/x86/gemm_x86_avxvnni.cpp | 4 +- src/layer/x86/gemm_x86_avxvnniint8.cpp | 4 +- src/layer/x86/gemm_x86_xop.cpp | 4 +- 21 files changed, 8744 insertions(+), 5870 deletions(-) diff --git a/src/layer/arm/gemm_arm.cpp b/src/layer/arm/gemm_arm.cpp index 1822c15396a0..3f8c791a7cc6 100644 --- a/src/layer/arm/gemm_arm.cpp +++ b/src/layer/arm/gemm_arm.cpp @@ -4594,14 +4594,15 @@ static int gemm_BT_arm_wq_int8(const Mat& A, const Mat& packed_B, const Mat& pac const Mat BT = packed_B.reshape(K, N); const Mat BT_descales = packed_B_descales.reshape(block_count, N); int TILE_M, TILE_N, TILE_K; - get_optimal_tile_mnk_wq_int8(M, N, K, constant_TILE_M, constant_TILE_N, constant_TILE_K, TILE_M, TILE_N, TILE_K, nT); + get_optimal_tile_mnk_wq_int8(M, N, K, block_size, constant_TILE_M, constant_TILE_N, constant_TILE_K, TILE_M, TILE_N, TILE_K, nT); - (void)TILE_K; const int mr = std::min(M, TILE_M); const int nr = std::min(N, TILE_N); const int nn_M = (M + TILE_M - 1) / TILE_M; const int nn_N = (N + TILE_N - 1) / TILE_N; + const int nn_K = (K + TILE_K - 1) / TILE_K; const float* input_scale_ptr = input_scales; + const size_t A_hstep = A.dims == 3 ? A.cstep : (size_t)A.w; Mat topT(nr * mr, 1, nT, (size_t)4u, opt.workspace_allocator); if (topT.empty()) @@ -4614,19 +4615,31 @@ static int gemm_BT_arm_wq_int8(const Mat& A, const Mat& packed_B, const Mat& pac if (AT.empty() || AT_descales.empty()) return -100; + const int nn_MK = nn_M * nn_K; #pragma omp parallel for num_threads(nT) - for (int ppi = 0; ppi < nn_M; ppi++) + for (int ppik = 0; ppik < nn_MK; ppik++) { + const int ppi = ppik / nn_K; + const int ppk = ppik % nn_K; const int i = ppi * TILE_M; + const int k = ppk * TILE_K; const int max_ii = std::min(M - i, TILE_M); + const int max_kk = std::min(K - k, TILE_K); + const int max_block_count = (max_kk + block_size - 1) / block_size; - Mat AT_tile = AT.channel(i / TILE_M); - Mat AT_descales_tile = AT_descales.channel(i / TILE_M); + Mat AT_channel = AT.channel(i / TILE_M); + Mat AT_descales_channel = AT_descales.channel(i / TILE_M); + Mat AT_tile(max_kk, max_ii, (signed char*)AT_channel + (size_t)k * mr, (size_t)1u); + Mat AT_descales_tile(max_block_count, max_ii, (float*)AT_descales_channel + (size_t)(k / block_size) * mr, (size_t)4u); + + Mat A_tile = A; + A_tile.data = (unsigned char*)A_tile.data + (transA ? (size_t)k * A_hstep : (size_t)k) * sizeof(float); + const float* input_scale_tile = input_scale_ptr ? input_scale_ptr + k : 0; if (transA) - transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); + transpose_quantize_A_tile_wq_int8(A_tile, AT_tile, AT_descales_tile, i, max_ii, max_kk, block_size, input_scale_tile); else - quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); + quantize_A_tile_wq_int8(A_tile, AT_tile, AT_descales_tile, i, max_ii, max_kk, block_size, input_scale_tile); } const int nn_MN = nn_M * nn_N; @@ -4642,13 +4655,21 @@ static int gemm_BT_arm_wq_int8(const Mat& A, const Mat& packed_B, const Mat& pac const int max_ii = std::min(M - i, TILE_M); const int max_jj = std::min(N - j, TILE_N); - Mat AT_tile = AT.channel(i / TILE_M); - Mat AT_descales_tile = AT_descales.channel(i / TILE_M); Mat BT_tile = BT.row_range(j, max_jj); Mat BT_descales_tile = BT_descales.row_range(j, max_jj); Mat topT_tile = topT.channel(get_omp_thread_num()); - gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); + Mat AT_channel = AT.channel(i / TILE_M); + Mat AT_descales_channel = AT_descales.channel(i / TILE_M); + for (int k = 0; k < K; k += TILE_K) + { + const int max_kk = std::min(K - k, TILE_K); + const int max_block_count = (max_kk + block_size - 1) / block_size; + Mat AT_tile(max_kk, max_ii, (signed char*)AT_channel + (size_t)k * mr, (size_t)1u); + Mat AT_descales_tile(max_block_count, max_ii, (float*)AT_descales_channel + (size_t)(k / block_size) * mr, (size_t)4u); + + gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, k, max_kk, block_size); + } if (output_transpose) transpose_unpack_output_tile_wq_int8(topT_tile, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, N, alpha, beta); else @@ -4672,18 +4693,33 @@ static int gemm_BT_arm_wq_int8(const Mat& A, const Mat& packed_B, const Mat& pac Mat AT_descales_tile = ATX_descales.channel(get_omp_thread_num()); Mat topT_tile = topT.channel(get_omp_thread_num()); - if (transA) - transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); - else - quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); - for (int j = 0; j < N; j += TILE_N) { const int max_jj = std::min(N - j, TILE_N); Mat BT_tile = BT.row_range(j, max_jj); Mat BT_descales_tile = BT_descales.row_range(j, max_jj); - gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); + for (int k = 0; k < K; k += TILE_K) + { + const int max_kk = std::min(K - k, TILE_K); + const int max_block_count = (max_kk + block_size - 1) / block_size; + Mat AT_tile_k(max_kk, max_ii, (signed char*)AT_tile + (size_t)k * mr, (size_t)1u); + Mat AT_descales_tile_k(max_block_count, max_ii, (float*)AT_descales_tile + (size_t)(k / block_size) * mr, (size_t)4u); + + if (j == 0) + { + Mat A_tile = A; + A_tile.data = (unsigned char*)A_tile.data + (transA ? (size_t)k * A_hstep : (size_t)k) * sizeof(float); + const float* input_scale_tile = input_scale_ptr ? input_scale_ptr + k : 0; + + if (transA) + transpose_quantize_A_tile_wq_int8(A_tile, AT_tile_k, AT_descales_tile_k, i, max_ii, max_kk, block_size, input_scale_tile); + else + quantize_A_tile_wq_int8(A_tile, AT_tile_k, AT_descales_tile_k, i, max_ii, max_kk, block_size, input_scale_tile); + } + + gemm_transB_packed_tile_wq_int8(AT_tile_k, AT_descales_tile_k, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, k, max_kk, block_size); + } if (output_transpose) transpose_unpack_output_tile_wq_int8(topT_tile, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, N, alpha, beta); else diff --git a/src/layer/arm/gemm_arm_asimddp.cpp b/src/layer/arm/gemm_arm_asimddp.cpp index 240ac68d0a17..12f3c852867b 100644 --- a/src/layer/arm/gemm_arm_asimddp.cpp +++ b/src/layer/arm/gemm_arm_asimddp.cpp @@ -34,9 +34,9 @@ void transpose_quantize_A_tile_wq_int8_asimddp(const Mat& A, Mat& AT_tile, Mat& transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); } -void gemm_transB_packed_tile_wq_int8_asimddp(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int block_size) +void gemm_transB_packed_tile_wq_int8_asimddp(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int k, int max_kk, int block_size) { - gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); + gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, k, max_kk, block_size); } #endif // NCNN_WEIGHT_QUANT diff --git a/src/layer/arm/gemm_arm_i8mm.cpp b/src/layer/arm/gemm_arm_i8mm.cpp index 882f6d57f5f1..902c8f56b6e9 100644 --- a/src/layer/arm/gemm_arm_i8mm.cpp +++ b/src/layer/arm/gemm_arm_i8mm.cpp @@ -34,9 +34,9 @@ void transpose_quantize_A_tile_wq_int8_i8mm(const Mat& A, Mat& AT_tile, Mat& AT_ transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); } -void gemm_transB_packed_tile_wq_int8_i8mm(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int block_size) +void gemm_transB_packed_tile_wq_int8_i8mm(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int k, int max_kk, int block_size) { - gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); + gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, k, max_kk, block_size); } #endif // NCNN_WEIGHT_QUANT diff --git a/src/layer/arm/gemm_arm_svei8mm.cpp b/src/layer/arm/gemm_arm_svei8mm.cpp index 73512ebca56e..ef28b536e050 100644 --- a/src/layer/arm/gemm_arm_svei8mm.cpp +++ b/src/layer/arm/gemm_arm_svei8mm.cpp @@ -17,9 +17,9 @@ int pack_B_wq_int8_svei8mm(const Mat& B, const Mat& B_scales, Mat& BT, Mat& BT_d return pack_B_wq_int8(B, B_scales, BT, BT_descales, N, K, block_size, num_threads); } -void gemm_transB_packed_tile_wq_int8_svei8mm(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int block_size) +void gemm_transB_packed_tile_wq_int8_svei8mm(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int k, int max_kk, int block_size) { - gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); + gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, k, max_kk, block_size); } #endif // NCNN_WEIGHT_QUANT diff --git a/src/layer/arm/gemm_wq_int8.h b/src/layer/arm/gemm_wq_int8.h index 09fb487e41fd..d6210dc149a2 100644 --- a/src/layer/arm/gemm_wq_int8.h +++ b/src/layer/arm/gemm_wq_int8.h @@ -4,22 +4,24 @@ #include #if NCNN_RUNTIME_CPU && NCNN_ARM86SVEI8MM && __aarch64__ && !__ARM_FEATURE_SVE_MATMUL_INT8 +// The svei8mm translation unit reuses the aarch64 i8mm v-register kernel and +// its persistent B layout. There is no SVE z-register WQ kernel layout yet. int pack_B_wq_int8_svei8mm(const Mat& B, const Mat& B_scales, Mat& BT, Mat& BT_descales, int N, int K, int block_size, int num_threads); -void gemm_transB_packed_tile_wq_int8_svei8mm(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int block_size); +void gemm_transB_packed_tile_wq_int8_svei8mm(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int k, int max_kk, int block_size); #endif #if NCNN_RUNTIME_CPU && NCNN_ARM84I8MM && __aarch64__ && !__ARM_FEATURE_MATMUL_INT8 int pack_B_wq_int8_i8mm(const Mat& B, const Mat& B_scales, Mat& BT, Mat& BT_descales, int N, int K, int block_size, int num_threads); void quantize_A_tile_wq_int8_i8mm(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr); void transpose_quantize_A_tile_wq_int8_i8mm(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr); -void gemm_transB_packed_tile_wq_int8_i8mm(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int block_size); +void gemm_transB_packed_tile_wq_int8_i8mm(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int k, int max_kk, int block_size); #endif #if NCNN_RUNTIME_CPU && NCNN_ARM82DOT && __aarch64__ && !__ARM_FEATURE_DOTPROD && !__ARM_FEATURE_MATMUL_INT8 int pack_B_wq_int8_asimddp(const Mat& B, const Mat& B_scales, Mat& BT, Mat& BT_descales, int N, int K, int block_size, int num_threads); void quantize_A_tile_wq_int8_asimddp(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr); void transpose_quantize_A_tile_wq_int8_asimddp(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr); -void gemm_transB_packed_tile_wq_int8_asimddp(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int block_size); +void gemm_transB_packed_tile_wq_int8_asimddp(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int k, int max_kk, int block_size); #endif static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr) @@ -46,101 +48,98 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales const int block_count = (K + block_size - 1) / block_size; const size_t A_hstep = A.dims == 3 ? A.cstep : (size_t)A.w; -#if __ARM_NEON && __aarch64__ + int ii = 0; +#if __ARM_NEON +#if __aarch64__ if (max_ii >= 8) { signed char* pp = AT_tile; + const float* pA0 = (const float*)A + i * A_hstep; + const float* ps = input_scale_ptr; + float* pd = descales; for (int g = 0; g < block_count; g++) { - const int k0 = g * block_size; - const int max_kk = std::min(K - k0, block_size); + const int max_kk = std::min(K - g * block_size, block_size); float absmax[8]; float scales[8]; for (int r = 0; r < 8; r++) { - const float* ptrA = (const float*)A + (i + r) * A_hstep + k0; + const float* ptrA = pA0 + r * A_hstep; float32x4_t _absmax = vdupq_n_f32(0.f); int kk = 0; for (; kk + 3 < max_kk; kk += 4) { float32x4_t _v = vld1q_f32(ptrA + kk); - if (input_scale_ptr) - _v = vmulq_f32(_v, vld1q_f32(input_scale_ptr + k0 + kk)); + if (ps) + _v = vmulq_f32(_v, vld1q_f32(ps + kk)); _absmax = vmaxq_f32(_absmax, vabsq_f32(_v)); } -#if __aarch64__ absmax[r] = vmaxvq_f32(_absmax); -#else - float32x2_t _max2 = vmax_f32(vget_low_f32(_absmax), vget_high_f32(_absmax)); - _max2 = vpmax_f32(_max2, _max2); - absmax[r] = vget_lane_f32(_max2, 0); -#endif for (; kk < max_kk; kk++) { float v = ptrA[kk]; - if (input_scale_ptr) - v *= input_scale_ptr[k0 + kk]; + if (ps) + v *= ps[kk]; absmax[r] = std::max(absmax[r], fabsf(v)); } - descales[g * 8 + r] = absmax[r] / 127.f; - volatile double scale_fp64 = absmax[r] == 0.f ? 0.0 : 127.0 / (double)absmax[r]; - scales[r] = (float)scale_fp64; + pd[r] = absmax[r] / 127.f; + scales[r] = absmax[r] == 0.f ? 0.f : 127.f / absmax[r]; } int kk = 0; +#if __ARM_FEATURE_DOTPROD #if __ARM_FEATURE_MATMUL_INT8 for (; kk + 7 < max_kk; kk += 8) { + float32x4_t _s0; + float32x4_t _s1; + if (ps) + { + _s0 = vld1q_f32(ps + kk); + _s1 = vld1q_f32(ps + kk + 4); + } for (int r = 0; r < 8; r += 2) { - const float* ptrA0 = (const float*)A + (i + r) * A_hstep + k0 + kk; + const float* ptrA0 = pA0 + r * A_hstep + kk; const float* ptrA1 = ptrA0 + A_hstep; float32x4_t _v00 = vld1q_f32(ptrA0); float32x4_t _v01 = vld1q_f32(ptrA0 + 4); float32x4_t _v10 = vld1q_f32(ptrA1); float32x4_t _v11 = vld1q_f32(ptrA1 + 4); - if (input_scale_ptr) + if (ps) { - const float32x4_t _s0 = vld1q_f32(input_scale_ptr + k0 + kk); - const float32x4_t _s1 = vld1q_f32(input_scale_ptr + k0 + kk + 4); _v00 = vmulq_f32(_v00, _s0); _v01 = vmulq_f32(_v01, _s1); _v10 = vmulq_f32(_v10, _s0); _v11 = vmulq_f32(_v11, _s1); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+w"(_v00), "+w"(_v01), "+w"(_v10), "+w"(_v11)); -#endif } - const float32x4_t _scale0 = vdupq_n_f32(scales[r]); - const float32x4_t _scale1 = vdupq_n_f32(scales[r + 1]); - const int8x8_t _q0 = float2int8(vmulq_f32(_v00, _scale0), vmulq_f32(_v01, _scale0)); - const int8x8_t _q1 = float2int8(vmulq_f32(_v10, _scale1), vmulq_f32(_v11, _scale1)); + float32x4_t _scale0 = vdupq_n_f32(scales[r]); + float32x4_t _scale1 = vdupq_n_f32(scales[r + 1]); + int8x8_t _q0 = float2int8(vmulq_f32(_v00, _scale0), vmulq_f32(_v01, _scale0)); + int8x8_t _q1 = float2int8(vmulq_f32(_v10, _scale1), vmulq_f32(_v11, _scale1)); vst1q_s8(pp, vcombine_s8(_q0, _q1)); pp += 16; } } #endif // __ARM_FEATURE_MATMUL_INT8 -#if __ARM_FEATURE_DOTPROD for (; kk + 3 < max_kk; kk += 4) { + float32x4_t _s; + if (ps) + _s = vld1q_f32(ps + kk); for (int r = 0; r < 8; r++) { - const float* ptrA = (const float*)A + (i + r) * A_hstep + k0 + kk; + const float* ptrA = pA0 + r * A_hstep + kk; float32x4_t _v = vld1q_f32(ptrA); - if (input_scale_ptr) + if (ps) { - _v = vmulq_f32(_v, vld1q_f32(input_scale_ptr + k0 + kk)); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+w"(_v)); -#endif + _v = vmulq_f32(_v, _s); } - const float32x4_t _scale = vdupq_n_f32(scales[r]); - const int8x8_t _q = float2int8(vmulq_f32(_v, _scale), vmulq_f32(_v, _scale)); + float32x4_t _scale = vdupq_n_f32(scales[r]); + int8x8_t _q = float2int8(vmulq_f32(_v, _scale), vmulq_f32(_v, _scale)); vst1_lane_s32((int*)pp, vreinterpret_s32_s8(_q), 0); pp += 4; } @@ -148,19 +147,22 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales #endif // __ARM_FEATURE_DOTPROD for (; kk + 1 < max_kk; kk += 2) { + float s0 = 0.f; + float s1 = 0.f; + if (ps) + { + s0 = ps[kk]; + s1 = ps[kk + 1]; + } for (int r = 0; r < 8; r++) { - const float* ptrA = (const float*)A + (i + r) * A_hstep + k0 + kk; + const float* ptrA = pA0 + r * A_hstep + kk; float v0 = ptrA[0]; float v1 = ptrA[1]; - if (input_scale_ptr) + if (ps) { - v0 *= input_scale_ptr[k0 + kk]; - v1 *= input_scale_ptr[k0 + kk + 1]; -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+w"(v0), "+w"(v1)); -#endif + v0 *= s0; + v1 *= s1; } *pp++ = float2int8(v0 * scales[r]); *pp++ = float2int8(v1 * scales[r]); @@ -168,51 +170,52 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales } if (kk < max_kk) { + float s = 0.f; + if (ps) + s = ps[kk]; for (int r = 0; r < 8; r++) { - const float* ptrA = (const float*)A + (i + r) * A_hstep + k0 + kk; + const float* ptrA = pA0 + r * A_hstep + kk; float v = ptrA[0]; - if (input_scale_ptr) + if (ps) { - v *= input_scale_ptr[k0 + kk]; -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+w"(v)); -#endif + v *= s; } *pp++ = float2int8(v * scales[r]); } } + pA0 += max_kk; + if (ps) + ps += max_kk; + pd += 8; } return; } -#endif // __ARM_NEON && __aarch64__ - - int ii = 0; -#if __ARM_NEON +#endif // __aarch64__ for (; ii + 3 < max_ii; ii += 4) { signed char* pp = outptr + ii * out_hstep; float* descale_ptr = descales + ii * descales_hstep; + const float* pA0 = (const float*)A + (i + ii) * A_hstep; + const float* ps = input_scale_ptr; for (int g = 0; g < block_count; g++) { - const int k0 = g * block_size; - const int max_kk = std::min(K - k0, block_size); + const int max_kk = std::min(K - g * block_size, block_size); float absmax[4]; float scales[4]; for (int r = 0; r < 4; r++) { - const float* ptrA = (const float*)A + (i + ii + r) * A_hstep + k0; + const float* ptrA = pA0 + r * A_hstep; float32x4_t _absmax = vdupq_n_f32(0.f); int kk = 0; for (; kk + 3 < max_kk; kk += 4) { float32x4_t _v = vld1q_f32(ptrA + kk); - if (input_scale_ptr) - _v = vmulq_f32(_v, vld1q_f32(input_scale_ptr + k0 + kk)); + if (ps) + _v = vmulq_f32(_v, vld1q_f32(ps + kk)); _absmax = vmaxq_f32(_absmax, vabsq_f32(_v)); } #if __aarch64__ @@ -225,65 +228,64 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales for (; kk < max_kk; kk++) { float v = ptrA[kk]; - if (input_scale_ptr) - v *= input_scale_ptr[k0 + kk]; + if (ps) + v *= ps[kk]; absmax[r] = std::max(absmax[r], fabsf(v)); } - descale_ptr[g * 4 + r] = absmax[r] / 127.f; - volatile double scale_fp64 = absmax[r] == 0.f ? 0.0 : 127.0 / (double)absmax[r]; - scales[r] = (float)scale_fp64; + descale_ptr[r] = absmax[r] / 127.f; + scales[r] = absmax[r] == 0.f ? 0.f : 127.f / absmax[r]; } int kk = 0; +#if __ARM_FEATURE_DOTPROD #if __ARM_FEATURE_MATMUL_INT8 for (; kk + 7 < max_kk; kk += 8) { + float32x4_t _s0; + float32x4_t _s1; + if (ps) + { + _s0 = vld1q_f32(ps + kk); + _s1 = vld1q_f32(ps + kk + 4); + } for (int r = 0; r < 4; r += 2) { - const float* ptrA0 = (const float*)A + (i + ii + r) * A_hstep + k0 + kk; + const float* ptrA0 = pA0 + r * A_hstep + kk; const float* ptrA1 = ptrA0 + A_hstep; float32x4_t _v00 = vld1q_f32(ptrA0); float32x4_t _v01 = vld1q_f32(ptrA0 + 4); float32x4_t _v10 = vld1q_f32(ptrA1); float32x4_t _v11 = vld1q_f32(ptrA1 + 4); - if (input_scale_ptr) + if (ps) { - const float32x4_t _s0 = vld1q_f32(input_scale_ptr + k0 + kk); - const float32x4_t _s1 = vld1q_f32(input_scale_ptr + k0 + kk + 4); _v00 = vmulq_f32(_v00, _s0); _v01 = vmulq_f32(_v01, _s1); _v10 = vmulq_f32(_v10, _s0); _v11 = vmulq_f32(_v11, _s1); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+w"(_v00), "+w"(_v01), "+w"(_v10), "+w"(_v11)); -#endif } - const float32x4_t _scale0 = vdupq_n_f32(scales[r]); - const float32x4_t _scale1 = vdupq_n_f32(scales[r + 1]); - const int8x8_t _q0 = float2int8(vmulq_f32(_v00, _scale0), vmulq_f32(_v01, _scale0)); - const int8x8_t _q1 = float2int8(vmulq_f32(_v10, _scale1), vmulq_f32(_v11, _scale1)); + float32x4_t _scale0 = vdupq_n_f32(scales[r]); + float32x4_t _scale1 = vdupq_n_f32(scales[r + 1]); + int8x8_t _q0 = float2int8(vmulq_f32(_v00, _scale0), vmulq_f32(_v01, _scale0)); + int8x8_t _q1 = float2int8(vmulq_f32(_v10, _scale1), vmulq_f32(_v11, _scale1)); vst1q_s8(pp, vcombine_s8(_q0, _q1)); pp += 16; } } #endif // __ARM_FEATURE_MATMUL_INT8 -#if __ARM_FEATURE_DOTPROD for (; kk + 3 < max_kk; kk += 4) { + float32x4_t _s; + if (ps) + _s = vld1q_f32(ps + kk); for (int r = 0; r < 4; r++) { - const float* ptrA = (const float*)A + (i + ii + r) * A_hstep + k0 + kk; + const float* ptrA = pA0 + r * A_hstep + kk; float32x4_t _v = vld1q_f32(ptrA); - if (input_scale_ptr) + if (ps) { - _v = vmulq_f32(_v, vld1q_f32(input_scale_ptr + k0 + kk)); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+w"(_v)); -#endif + _v = vmulq_f32(_v, _s); } - const int8x8_t _q = float2int8(vmulq_n_f32(_v, scales[r]), vmulq_n_f32(_v, scales[r])); + int8x8_t _q = float2int8(vmulq_n_f32(_v, scales[r]), vmulq_n_f32(_v, scales[r])); vst1_lane_s32((int*)pp, vreinterpret_s32_s8(_q), 0); pp += 4; } @@ -291,19 +293,22 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales #endif // __ARM_FEATURE_DOTPROD for (; kk + 1 < max_kk; kk += 2) { + float s0 = 0.f; + float s1 = 0.f; + if (ps) + { + s0 = ps[kk]; + s1 = ps[kk + 1]; + } for (int r = 0; r < 4; r++) { - const float* ptrA = (const float*)A + (i + ii + r) * A_hstep + k0 + kk; + const float* ptrA = pA0 + r * A_hstep + kk; float v0 = ptrA[0]; float v1 = ptrA[1]; - if (input_scale_ptr) + if (ps) { - v0 *= input_scale_ptr[k0 + kk]; - v1 *= input_scale_ptr[k0 + kk + 1]; -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+w"(v0), "+w"(v1)); -#endif + v0 *= s0; + v1 *= s1; } *pp++ = float2int8(v0 * scales[r]); *pp++ = float2int8(v1 * scales[r]); @@ -311,45 +316,49 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales } if (kk < max_kk) { + float s = 0.f; + if (ps) + s = ps[kk]; for (int r = 0; r < 4; r++) { - const float* ptrA = (const float*)A + (i + ii + r) * A_hstep + k0 + kk; + const float* ptrA = pA0 + r * A_hstep + kk; float v = ptrA[0]; - if (input_scale_ptr) + if (ps) { - v *= input_scale_ptr[k0 + kk]; -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+w"(v)); -#endif + v *= s; } *pp++ = float2int8(v * scales[r]); } } + pA0 += max_kk; + if (ps) + ps += max_kk; + descale_ptr += 4; } } for (; ii + 1 < max_ii; ii += 2) { signed char* pp = outptr + ii * out_hstep; float* descale_ptr = descales + ii * descales_hstep; + const float* pA0 = (const float*)A + (i + ii) * A_hstep; + const float* ps = input_scale_ptr; for (int g = 0; g < block_count; g++) { - const int k0 = g * block_size; - const int max_kk = std::min(K - k0, block_size); + const int max_kk = std::min(K - g * block_size, block_size); float absmax[2]; float scales[2]; for (int r = 0; r < 2; r++) { - const float* ptrA = (const float*)A + (i + ii + r) * A_hstep + k0; + const float* ptrA = pA0 + r * A_hstep; float32x4_t _absmax = vdupq_n_f32(0.f); int kk = 0; for (; kk + 3 < max_kk; kk += 4) { float32x4_t _v = vld1q_f32(ptrA + kk); - if (input_scale_ptr) - _v = vmulq_f32(_v, vld1q_f32(input_scale_ptr + k0 + kk)); + if (ps) + _v = vmulq_f32(_v, vld1q_f32(ps + kk)); _absmax = vmaxq_f32(_absmax, vabsq_f32(_v)); } #if __aarch64__ @@ -362,61 +371,55 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales for (; kk < max_kk; kk++) { float v = ptrA[kk]; - if (input_scale_ptr) - v *= input_scale_ptr[k0 + kk]; + if (ps) + v *= ps[kk]; absmax[r] = std::max(absmax[r], fabsf(v)); } - descale_ptr[g * 2 + r] = absmax[r] / 127.f; - volatile double scale_fp64 = absmax[r] == 0.f ? 0.0 : 127.0 / (double)absmax[r]; - scales[r] = (float)scale_fp64; + descale_ptr[r] = absmax[r] / 127.f; + scales[r] = absmax[r] == 0.f ? 0.f : 127.f / absmax[r]; } int kk = 0; +#if __ARM_FEATURE_DOTPROD #if __ARM_FEATURE_MATMUL_INT8 for (; kk + 7 < max_kk; kk += 8) { - const float* ptrA0 = (const float*)A + (i + ii) * A_hstep + k0 + kk; + const float* ptrA0 = pA0 + kk; const float* ptrA1 = ptrA0 + A_hstep; float32x4_t _v00 = vld1q_f32(ptrA0); float32x4_t _v01 = vld1q_f32(ptrA0 + 4); float32x4_t _v10 = vld1q_f32(ptrA1); float32x4_t _v11 = vld1q_f32(ptrA1 + 4); - if (input_scale_ptr) + if (ps) { - const float32x4_t _s0 = vld1q_f32(input_scale_ptr + k0 + kk); - const float32x4_t _s1 = vld1q_f32(input_scale_ptr + k0 + kk + 4); + float32x4_t _s0 = vld1q_f32(ps + kk); + float32x4_t _s1 = vld1q_f32(ps + kk + 4); _v00 = vmulq_f32(_v00, _s0); _v01 = vmulq_f32(_v01, _s1); _v10 = vmulq_f32(_v10, _s0); _v11 = vmulq_f32(_v11, _s1); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+w"(_v00), "+w"(_v01), "+w"(_v10), "+w"(_v11)); -#endif } - const int8x8_t _q0 = float2int8(vmulq_n_f32(_v00, scales[0]), vmulq_n_f32(_v01, scales[0])); - const int8x8_t _q1 = float2int8(vmulq_n_f32(_v10, scales[1]), vmulq_n_f32(_v11, scales[1])); + int8x8_t _q0 = float2int8(vmulq_n_f32(_v00, scales[0]), vmulq_n_f32(_v01, scales[0])); + int8x8_t _q1 = float2int8(vmulq_n_f32(_v10, scales[1]), vmulq_n_f32(_v11, scales[1])); vst1q_s8(pp, vcombine_s8(_q0, _q1)); pp += 16; } #endif // __ARM_FEATURE_MATMUL_INT8 -#if __ARM_FEATURE_DOTPROD for (; kk + 3 < max_kk; kk += 4) { int32x2_t _q01 = vdup_n_s32(0); + float32x4_t _s; + if (ps) + _s = vld1q_f32(ps + kk); for (int r = 0; r < 2; r++) { - const float* ptrA = (const float*)A + (i + ii + r) * A_hstep + k0 + kk; + const float* ptrA = pA0 + r * A_hstep + kk; float32x4_t _v = vld1q_f32(ptrA); - if (input_scale_ptr) + if (ps) { - _v = vmulq_f32(_v, vld1q_f32(input_scale_ptr + k0 + kk)); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+w"(_v)); -#endif + _v = vmulq_f32(_v, _s); } - const int8x8_t _q = float2int8(vmulq_n_f32(_v, scales[r]), vmulq_n_f32(_v, scales[r])); + int8x8_t _q = float2int8(vmulq_n_f32(_v, scales[r]), vmulq_n_f32(_v, scales[r])); _q01 = vset_lane_s32(vget_lane_s32(vreinterpret_s32_s8(_q), 0), _q01, r); } vst1_s8(pp, vreinterpret_s8_s32(_q01)); @@ -425,15 +428,22 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales #endif // __ARM_FEATURE_DOTPROD for (; kk + 1 < max_kk; kk += 2) { + float s0 = 0.f; + float s1 = 0.f; + if (ps) + { + s0 = ps[kk]; + s1 = ps[kk + 1]; + } for (int r = 0; r < 2; r++) { - const float* ptrA = (const float*)A + (i + ii + r) * A_hstep + k0 + kk; + const float* ptrA = pA0 + r * A_hstep + kk; float v0 = ptrA[0]; float v1 = ptrA[1]; - if (input_scale_ptr) + if (ps) { - v0 *= input_scale_ptr[k0 + kk]; - v1 *= input_scale_ptr[k0 + kk + 1]; + v0 *= s0; + v1 *= s1; } *pp++ = float2int8(v0 * scales[r]); *pp++ = float2int8(v1 * scales[r]); @@ -441,15 +451,22 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales } if (kk < max_kk) { + float s = 0.f; + if (ps) + s = ps[kk]; for (int r = 0; r < 2; r++) { - const float* ptrA = (const float*)A + (i + ii + r) * A_hstep + k0 + kk; + const float* ptrA = pA0 + r * A_hstep + kk; float v = ptrA[0]; - if (input_scale_ptr) - v *= input_scale_ptr[k0 + kk]; + if (ps) + v *= s; *pp++ = float2int8(v * scales[r]); } } + pA0 += max_kk; + if (ps) + ps += max_kk; + descale_ptr += 2; } } #elif __ARM_FEATURE_SIMD32 && NCNN_GNU_INLINE_ASM @@ -457,63 +474,75 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales { signed char* pp = outptr + ii * out_hstep; float* descale_ptr = descales + ii * descales_hstep; + const float* pA0g = (const float*)A + (i + ii) * A_hstep; + const float* pA1g = pA0g + A_hstep; + const float* ps = input_scale_ptr; for (int g = 0; g < block_count; g++) { - const int k0 = g * block_size; - const int max_kk = std::min(K - k0, block_size); - const float* ptrA0 = (const float*)A + (i + ii) * A_hstep + k0; - const float* ptrA1 = ptrA0 + A_hstep; + const int max_kk = std::min(K - g * block_size, block_size); + const float* pA0 = pA0g; + const float* pA1 = pA1g; + const float* pscale = ps; float absmax0 = 0.f; float absmax1 = 0.f; for (int kk = 0; kk < max_kk; kk++) { - float v0 = ptrA0[kk]; - float v1 = ptrA1[kk]; - if (input_scale_ptr) + float v0 = *pA0++; + float v1 = *pA1++; + if (pscale) { - v0 *= input_scale_ptr[k0 + kk]; - v1 *= input_scale_ptr[k0 + kk]; + const float s = *pscale++; + v0 *= s; + v1 *= s; } absmax0 = std::max(absmax0, fabsf(v0)); absmax1 = std::max(absmax1, fabsf(v1)); } - descale_ptr[g * 2] = absmax0 / 127.f; - descale_ptr[g * 2 + 1] = absmax1 / 127.f; - volatile double scale0_fp64 = absmax0 == 0.f ? 0.0 : 127.0 / (double)absmax0; - volatile double scale1_fp64 = absmax1 == 0.f ? 0.0 : 127.0 / (double)absmax1; - const float scale0 = (float)scale0_fp64; - const float scale1 = (float)scale1_fp64; + descale_ptr[0] = absmax0 / 127.f; + descale_ptr[1] = absmax1 / 127.f; + const float scale0 = absmax0 == 0.f ? 0.f : 127.f / absmax0; + const float scale1 = absmax1 == 0.f ? 0.f : 127.f / absmax1; + pA0 = pA0g; + pA1 = pA1g; + pscale = ps; for (int kk = 0; kk < max_kk; kk++) { - float v0 = ptrA0[kk]; - float v1 = ptrA1[kk]; - if (input_scale_ptr) + float v0 = *pA0++; + float v1 = *pA1++; + if (pscale) { - v0 *= input_scale_ptr[k0 + kk]; - v1 *= input_scale_ptr[k0 + kk]; - asm volatile("" - : "+w"(v0), "+w"(v1)); + const float s = *pscale++; + v0 *= s; + v1 *= s; } *pp++ = float2int8(v0 * scale0); *pp++ = float2int8(v1 * scale1); } + + pA0g += max_kk; + pA1g += max_kk; + if (ps) + ps += max_kk; + descale_ptr += 2; } } #endif // __ARM_NEON for (; ii < max_ii; ii++) { - const float* ptrA = (const float*)A + (i + ii) * A_hstep; - signed char* outptr0 = outptr + ii * out_hstep; - float* descale_ptr = descales + ii * descales_hstep; + const float* pAg = (const float*)A + (i + ii) * A_hstep; + const float* ps = input_scale_ptr; + signed char* pp = outptr + ii * out_hstep; + float* pd = descales + ii * descales_hstep; for (int g = 0; g < block_count; g++) { - const int k0 = g * block_size; - const int max_kk = std::min(K - k0, block_size); + const int max_kk = std::min(K - g * block_size, block_size); + const float* pA = pAg; + const float* pscale = ps; float absmax = 0.f; int kk = 0; @@ -521,9 +550,13 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales float32x4_t _absmax = vdupq_n_f32(0.f); for (; kk + 3 < max_kk; kk += 4) { - float32x4_t _v = vld1q_f32(ptrA + k0 + kk); - if (input_scale_ptr) - _v = vmulq_f32(_v, vld1q_f32(input_scale_ptr + k0 + kk)); + float32x4_t _v = vld1q_f32(pA); + pA += 4; + if (pscale) + { + _v = vmulq_f32(_v, vld1q_f32(pscale)); + pscale += 4; + } _absmax = vmaxq_f32(_absmax, vabsq_f32(_v)); } #if __aarch64__ @@ -536,83 +569,70 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales #endif // __ARM_NEON for (; kk < max_kk; kk++) { - const int k = k0 + kk; - float v = ptrA[k]; - if (input_scale_ptr) - v *= input_scale_ptr[k]; + float v = *pA++; + if (pscale) + v *= *pscale++; absmax = std::max(absmax, fabsf(v)); } if (absmax == 0.f) { - descale_ptr[g] = 0.f; + *pd++ = 0.f; for (int k = 0; k < max_kk; k++) - outptr0[k0 + k] = 0; + *pp++ = 0; + pAg += max_kk; + if (ps) + ps += max_kk; continue; } - volatile double scale_fp64 = 127.0 / (double)absmax; - const float scale = (float)scale_fp64; - descale_ptr[g] = absmax / 127.f; + const float scale = 127.f / absmax; + *pd++ = absmax / 127.f; kk = 0; + pA = pAg; + pscale = ps; #if __ARM_NEON - const float32x4_t _scale = vdupq_n_f32(scale); + float32x4_t _scale = vdupq_n_f32(scale); for (; kk + 7 < max_kk; kk += 8) { - float32x4_t _v0 = vld1q_f32(ptrA + k0 + kk); - float32x4_t _v1 = vld1q_f32(ptrA + k0 + kk + 4); - if (input_scale_ptr) + float32x4_t _v0 = vld1q_f32(pA); + float32x4_t _v1 = vld1q_f32(pA + 4); + pA += 8; + if (pscale) { - _v0 = vmulq_f32(_v0, vld1q_f32(input_scale_ptr + k0 + kk)); - _v1 = vmulq_f32(_v1, vld1q_f32(input_scale_ptr + k0 + kk + 4)); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+w"(_v0), "+w"(_v1)); -#else - volatile float32x4_t _v0_ordered = _v0; - volatile float32x4_t _v1_ordered = _v1; - _v0 = _v0_ordered; - _v1 = _v1_ordered; -#endif + _v0 = vmulq_f32(_v0, vld1q_f32(pscale)); + _v1 = vmulq_f32(_v1, vld1q_f32(pscale + 4)); + pscale += 8; } - vst1_s8(outptr0 + k0 + kk, float2int8(vmulq_f32(_v0, _scale), vmulq_f32(_v1, _scale))); + vst1_s8(pp, float2int8(vmulq_f32(_v0, _scale), vmulq_f32(_v1, _scale))); + pp += 8; } for (; kk + 3 < max_kk; kk += 4) { - float32x4_t _v = vld1q_f32(ptrA + k0 + kk); - if (input_scale_ptr) + float32x4_t _v = vld1q_f32(pA); + pA += 4; + if (pscale) { - _v = vmulq_f32(_v, vld1q_f32(input_scale_ptr + k0 + kk)); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+w"(_v)); -#else - volatile float32x4_t _v_ordered = _v; - _v = _v_ordered; -#endif + _v = vmulq_f32(_v, vld1q_f32(pscale)); + pscale += 4; } - const int8x8_t _q = float2int8(vmulq_f32(_v, _scale), vmulq_f32(_v, _scale)); - vst1_lane_s32((int*)(outptr0 + k0 + kk), vreinterpret_s32_s8(_q), 0); + int8x8_t _q = float2int8(vmulq_f32(_v, _scale), vmulq_f32(_v, _scale)); + vst1_lane_s32((int*)pp, vreinterpret_s32_s8(_q), 0); + pp += 4; } #endif // __ARM_NEON for (; kk < max_kk; kk++) { - const int k = k0 + kk; - float v = ptrA[k]; - if (input_scale_ptr) - { - v *= input_scale_ptr[k]; -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+w"(v)); -#else - volatile float v_ordered = v; - v = v_ordered; -#endif - } - outptr0[k] = float2int8(v * scale); + float v = *pA++; + if (pscale) + v *= *pscale++; + *pp++ = float2int8(v * scale); } + + pAg += max_kk; + if (ps) + ps += max_kk; } } } @@ -648,67 +668,67 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A { const float* ptrA = (const float*)A + i + ii; signed char* pp = outptr + ii * out_hstep; + const float* ptrAg = ptrA; + const float* ps = input_scale_ptr; + float* pd = descales; for (int g = 0; g < block_count; g++) { - const int k0 = g * block_size; - const int max_kk = std::min(K - k0, block_size); + const int max_kk = std::min(K - g * block_size, block_size); float32x4_t _absmax0 = vdupq_n_f32(0.f); float32x4_t _absmax1 = vdupq_n_f32(0.f); + const float* ptrAk = ptrAg; for (int kk = 0; kk < max_kk; kk++) { - const int k = k0 + kk; - float32x4_t _v0 = vld1q_f32(ptrA + (size_t)k * A_hstep); - float32x4_t _v1 = vld1q_f32(ptrA + (size_t)k * A_hstep + 4); - if (input_scale_ptr) + float32x4_t _v0 = vld1q_f32(ptrAk); + float32x4_t _v1 = vld1q_f32(ptrAk + 4); + if (ps) { - _v0 = vmulq_n_f32(_v0, input_scale_ptr[k]); - _v1 = vmulq_n_f32(_v1, input_scale_ptr[k]); + _v0 = vmulq_n_f32(_v0, ps[kk]); + _v1 = vmulq_n_f32(_v1, ps[kk]); } _absmax0 = vmaxq_f32(_absmax0, vabsq_f32(_v0)); _absmax1 = vmaxq_f32(_absmax1, vabsq_f32(_v1)); + ptrAk += A_hstep; } float absmax[8]; vst1q_f32(absmax, _absmax0); vst1q_f32(absmax + 4, _absmax1); - descales[g * 8] = absmax[0] / 127.f; - descales[g * 8 + 1] = absmax[1] / 127.f; - descales[g * 8 + 2] = absmax[2] / 127.f; - descales[g * 8 + 3] = absmax[3] / 127.f; - descales[g * 8 + 4] = absmax[4] / 127.f; - descales[g * 8 + 5] = absmax[5] / 127.f; - descales[g * 8 + 6] = absmax[6] / 127.f; - descales[g * 8 + 7] = absmax[7] / 127.f; + pd[0] = absmax[0] / 127.f; + pd[1] = absmax[1] / 127.f; + pd[2] = absmax[2] / 127.f; + pd[3] = absmax[3] / 127.f; + pd[4] = absmax[4] / 127.f; + pd[5] = absmax[5] / 127.f; + pd[6] = absmax[6] / 127.f; + pd[7] = absmax[7] / 127.f; float scales[8]; for (int r = 0; r < 8; r++) { - volatile double scale_fp64 = absmax[r] == 0.f ? 0.0 : 127.0 / (double)absmax[r]; - scales[r] = (float)scale_fp64; + scales[r] = absmax[r] == 0.f ? 0.f : 127.f / absmax[r]; } - const float32x4_t _scale0 = vld1q_f32(scales); - const float32x4_t _scale1 = vld1q_f32(scales + 4); + float32x4_t _scale0 = vld1q_f32(scales); + float32x4_t _scale1 = vld1q_f32(scales + 4); int kk = 0; +#if __ARM_FEATURE_DOTPROD #if __ARM_FEATURE_MATMUL_INT8 for (; kk + 7 < max_kk; kk += 8) { int8x8_t _q[8]; + ptrAk = ptrAg + (size_t)kk * A_hstep; for (int t = 0; t < 8; t++) { - const int k = k0 + kk + t; - float32x4_t _v0 = vld1q_f32(ptrA + (size_t)k * A_hstep); - float32x4_t _v1 = vld1q_f32(ptrA + (size_t)k * A_hstep + 4); - if (input_scale_ptr) + float32x4_t _v0 = vld1q_f32(ptrAk); + float32x4_t _v1 = vld1q_f32(ptrAk + 4); + if (ps) { - _v0 = vmulq_n_f32(_v0, input_scale_ptr[k]); - _v1 = vmulq_n_f32(_v1, input_scale_ptr[k]); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+w"(_v0), "+w"(_v1)); -#endif + _v0 = vmulq_n_f32(_v0, ps[kk + t]); + _v1 = vmulq_n_f32(_v1, ps[kk + t]); } _q[t] = float2int8(vmulq_f32(_v0, _scale0), vmulq_f32(_v1, _scale1)); + ptrAk += A_hstep; } int8x8x2_t _r04 = vzip_s8(_q[0], _q[4]); int8x8x2_t _r15 = vzip_s8(_q[1], _q[5]); @@ -729,25 +749,21 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A pp += 64; } #endif // __ARM_FEATURE_MATMUL_INT8 -#if __ARM_FEATURE_DOTPROD for (; kk + 3 < max_kk; kk += 4) { int8x8x4_t _q; + ptrAk = ptrAg + (size_t)kk * A_hstep; for (int t = 0; t < 4; t++) { - const int k = k0 + kk + t; - float32x4_t _v0 = vld1q_f32(ptrA + (size_t)k * A_hstep); - float32x4_t _v1 = vld1q_f32(ptrA + (size_t)k * A_hstep + 4); - if (input_scale_ptr) + float32x4_t _v0 = vld1q_f32(ptrAk); + float32x4_t _v1 = vld1q_f32(ptrAk + 4); + if (ps) { - _v0 = vmulq_n_f32(_v0, input_scale_ptr[k]); - _v1 = vmulq_n_f32(_v1, input_scale_ptr[k]); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+w"(_v0), "+w"(_v1)); -#endif + _v0 = vmulq_n_f32(_v0, ps[kk + t]); + _v1 = vmulq_n_f32(_v1, ps[kk + t]); } _q.val[t] = float2int8(vmulq_f32(_v0, _scale0), vmulq_f32(_v1, _scale1)); + ptrAk += A_hstep; } vst4_s8(pp, _q); pp += 32; @@ -756,42 +772,39 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A for (; kk + 1 < max_kk; kk += 2) { int8x8x2_t _q; + ptrAk = ptrAg + (size_t)kk * A_hstep; for (int t = 0; t < 2; t++) { - const int k = k0 + kk + t; - float32x4_t _v0 = vld1q_f32(ptrA + (size_t)k * A_hstep); - float32x4_t _v1 = vld1q_f32(ptrA + (size_t)k * A_hstep + 4); - if (input_scale_ptr) + float32x4_t _v0 = vld1q_f32(ptrAk); + float32x4_t _v1 = vld1q_f32(ptrAk + 4); + if (ps) { - _v0 = vmulq_n_f32(_v0, input_scale_ptr[k]); - _v1 = vmulq_n_f32(_v1, input_scale_ptr[k]); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+w"(_v0), "+w"(_v1)); -#endif + _v0 = vmulq_n_f32(_v0, ps[kk + t]); + _v1 = vmulq_n_f32(_v1, ps[kk + t]); } _q.val[t] = float2int8(vmulq_f32(_v0, _scale0), vmulq_f32(_v1, _scale1)); + ptrAk += A_hstep; } vst2_s8(pp, _q); pp += 16; } if (kk < max_kk) { - const int k = k0 + kk; - float32x4_t _v0 = vld1q_f32(ptrA + (size_t)k * A_hstep); - float32x4_t _v1 = vld1q_f32(ptrA + (size_t)k * A_hstep + 4); - if (input_scale_ptr) + ptrAk = ptrAg + (size_t)kk * A_hstep; + float32x4_t _v0 = vld1q_f32(ptrAk); + float32x4_t _v1 = vld1q_f32(ptrAk + 4); + if (ps) { - _v0 = vmulq_n_f32(_v0, input_scale_ptr[k]); - _v1 = vmulq_n_f32(_v1, input_scale_ptr[k]); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+w"(_v0), "+w"(_v1)); -#endif + _v0 = vmulq_n_f32(_v0, ps[kk]); + _v1 = vmulq_n_f32(_v1, ps[kk]); } vst1_s8(pp, float2int8(vmulq_f32(_v0, _scale0), vmulq_f32(_v1, _scale1))); pp += 8; } + ptrAg += (size_t)max_kk * A_hstep; + if (ps) + ps += max_kk; + pd += 8; } } #endif // __aarch64__ @@ -800,59 +813,52 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A const float* ptrA = (const float*)A + i + ii; signed char* pp = outptr + ii * out_hstep; float* descale_ptr = descales + ii * descales_hstep; + const float* ptrAg = ptrA; + const float* ps = input_scale_ptr; for (int g = 0; g < block_count; g++) { - const int k0 = g * block_size; - const int max_kk = std::min(K - k0, block_size); + const int max_kk = std::min(K - g * block_size, block_size); float32x4_t _absmax = vdupq_n_f32(0.f); + const float* ptrAk = ptrAg; for (int kk = 0; kk < max_kk; kk++) { - const int k = k0 + kk; - float32x4_t _v = vld1q_f32(ptrA + (size_t)k * A_hstep); - if (input_scale_ptr) - _v = vmulq_f32(_v, vdupq_n_f32(input_scale_ptr[k])); + float32x4_t _v = vld1q_f32(ptrAk); + if (ps) + _v = vmulq_f32(_v, vdupq_n_f32(ps[kk])); _absmax = vmaxq_f32(_absmax, vabsq_f32(_v)); + ptrAk += A_hstep; } float absmax[4]; vst1q_f32(absmax, _absmax); - vst1q_f32(descale_ptr + g * 4, vmulq_n_f32(_absmax, 1.f / 127.f)); + vst1q_f32(descale_ptr, vmulq_n_f32(_absmax, 1.f / 127.f)); - volatile double scale0_fp64 = absmax[0] == 0.f ? 0.0 : 127.0 / (double)absmax[0]; - volatile double scale1_fp64 = absmax[1] == 0.f ? 0.0 : 127.0 / (double)absmax[1]; - volatile double scale2_fp64 = absmax[2] == 0.f ? 0.0 : 127.0 / (double)absmax[2]; - volatile double scale3_fp64 = absmax[3] == 0.f ? 0.0 : 127.0 / (double)absmax[3]; const float scales[4] = { - (float)scale0_fp64, - (float)scale1_fp64, - (float)scale2_fp64, - (float)scale3_fp64 + absmax[0] == 0.f ? 0.f : 127.f / absmax[0], + absmax[1] == 0.f ? 0.f : 127.f / absmax[1], + absmax[2] == 0.f ? 0.f : 127.f / absmax[2], + absmax[3] == 0.f ? 0.f : 127.f / absmax[3] }; - const float32x4_t _scale = vld1q_f32(scales); + float32x4_t _scale = vld1q_f32(scales); int kk = 0; +#if __ARM_FEATURE_DOTPROD #if __ARM_FEATURE_MATMUL_INT8 for (; kk + 7 < max_kk; kk += 8) { int8x8_t _q[8]; + ptrAk = ptrAg + (size_t)kk * A_hstep; for (int t = 0; t < 8; t++) { - const int k = k0 + kk + t; - float32x4_t _v = vld1q_f32(ptrA + (size_t)k * A_hstep); - if (input_scale_ptr) + float32x4_t _v = vld1q_f32(ptrAk); + if (ps) { - _v = vmulq_n_f32(_v, input_scale_ptr[k]); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+w"(_v)); -#else - volatile float32x4_t _v_ordered = _v; - _v = _v_ordered; -#endif + _v = vmulq_n_f32(_v, ps[kk + t]); } _q[t] = float2int8(vmulq_f32(_v, _scale), vmulq_f32(_v, _scale)); + ptrAk += A_hstep; } int8x8x2_t _r04 = vzip_s8(_q[0], _q[4]); int8x8x2_t _r15 = vzip_s8(_q[1], _q[5]); @@ -867,26 +873,19 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A pp += 32; } #endif // __ARM_FEATURE_MATMUL_INT8 -#if __ARM_FEATURE_DOTPROD for (; kk + 3 < max_kk; kk += 4) { int8x8x4_t _q; + ptrAk = ptrAg + (size_t)kk * A_hstep; for (int t = 0; t < 4; t++) { - const int k = k0 + kk + t; - float32x4_t _v = vld1q_f32(ptrA + (size_t)k * A_hstep); - if (input_scale_ptr) + float32x4_t _v = vld1q_f32(ptrAk); + if (ps) { - _v = vmulq_n_f32(_v, input_scale_ptr[k]); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+w"(_v)); -#else - volatile float32x4_t _v_ordered = _v; - _v = _v_ordered; -#endif + _v = vmulq_n_f32(_v, ps[kk + t]); } _q.val[t] = float2int8(vmulq_f32(_v, _scale), vmulq_f32(_v, _scale)); + ptrAk += A_hstep; } vst4_lane_s8(pp, _q, 0); vst4_lane_s8(pp + 4, _q, 1); @@ -898,22 +897,16 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A for (; kk + 1 < max_kk; kk += 2) { int8x8x2_t _q; + ptrAk = ptrAg + (size_t)kk * A_hstep; for (int t = 0; t < 2; t++) { - const int k = k0 + kk + t; - float32x4_t _v = vld1q_f32(ptrA + (size_t)k * A_hstep); - if (input_scale_ptr) + float32x4_t _v = vld1q_f32(ptrAk); + if (ps) { - _v = vmulq_n_f32(_v, input_scale_ptr[k]); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+w"(_v)); -#else - volatile float32x4_t _v_ordered = _v; - _v = _v_ordered; -#endif + _v = vmulq_n_f32(_v, ps[kk + t]); } _q.val[t] = float2int8(vmulq_f32(_v, _scale), vmulq_f32(_v, _scale)); + ptrAk += A_hstep; } vst2_lane_s8(pp, _q, 0); vst2_lane_s8(pp + 2, _q, 1); @@ -923,22 +916,19 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A } if (kk < max_kk) { - const int k = k0 + kk; - float32x4_t _v = vld1q_f32(ptrA + (size_t)k * A_hstep); - if (input_scale_ptr) + ptrAk = ptrAg + (size_t)kk * A_hstep; + float32x4_t _v = vld1q_f32(ptrAk); + if (ps) { - _v = vmulq_n_f32(_v, input_scale_ptr[k]); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+w"(_v)); -#else - volatile float32x4_t _v_ordered = _v; - _v = _v_ordered; -#endif + _v = vmulq_n_f32(_v, ps[kk]); } vst1_lane_s32((int*)pp, vreinterpret_s32_s8(float2int8(vmulq_f32(_v, _scale), vmulq_f32(_v, _scale))), 0); pp += 4; } + ptrAg += (size_t)max_kk * A_hstep; + if (ps) + ps += max_kk; + descale_ptr += 4; } } for (; ii + 1 < max_ii; ii += 2) @@ -946,54 +936,58 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A const float* ptrA = (const float*)A + i + ii; signed char* pp = outptr + ii * out_hstep; float* descale_ptr = descales + ii * descales_hstep; + const float* ptrAg = ptrA; + const float* ps = input_scale_ptr; for (int g = 0; g < block_count; g++) { - const int k0 = g * block_size; - const int max_kk = std::min(K - k0, block_size); + const int max_kk = std::min(K - g * block_size, block_size); float32x2_t _absmax = vdup_n_f32(0.f); + const float* ptrAk = ptrAg; for (int kk = 0; kk < max_kk; kk++) { - const int k = k0 + kk; - float32x2_t _v = vld1_f32(ptrA + (size_t)k * A_hstep); - if (input_scale_ptr) - _v = vmul_n_f32(_v, input_scale_ptr[k]); + float32x2_t _v = vld1_f32(ptrAk); + if (ps) + _v = vmul_n_f32(_v, ps[kk]); _absmax = vmax_f32(_absmax, vabs_f32(_v)); + ptrAk += A_hstep; } - vst1_f32(descale_ptr + g * 2, vmul_n_f32(_absmax, 1.f / 127.f)); + vst1_f32(descale_ptr, vmul_n_f32(_absmax, 1.f / 127.f)); float absmax[2]; vst1_f32(absmax, _absmax); float scales[2]; for (int r = 0; r < 2; r++) { - volatile double scale_fp64 = absmax[r] == 0.f ? 0.0 : 127.0 / (double)absmax[r]; - scales[r] = (float)scale_fp64; + scales[r] = absmax[r] == 0.f ? 0.f : 127.f / absmax[r]; } - const float32x2_t _scale = vld1_f32(scales); + float32x2_t _scale = vld1_f32(scales); int kk = 0; +#if __ARM_FEATURE_DOTPROD #if __ARM_FEATURE_MATMUL_INT8 for (; kk + 7 < max_kk; kk += 8) { int8x8x4_t _q0; int8x8x4_t _q1; + const float* ptrAk0 = ptrAg + (size_t)kk * A_hstep; + const float* ptrAk1 = ptrAk0 + (size_t)4 * A_hstep; for (int t = 0; t < 4; t++) { - const int k0t = k0 + kk + t; - const int k1t = k0 + kk + 4 + t; - float32x2_t _v0 = vld1_f32(ptrA + (size_t)k0t * A_hstep); - float32x2_t _v1 = vld1_f32(ptrA + (size_t)k1t * A_hstep); - if (input_scale_ptr) + float32x2_t _v0 = vld1_f32(ptrAk0); + float32x2_t _v1 = vld1_f32(ptrAk1); + if (ps) { - _v0 = vmul_n_f32(_v0, input_scale_ptr[k0t]); - _v1 = vmul_n_f32(_v1, input_scale_ptr[k1t]); + _v0 = vmul_n_f32(_v0, ps[kk + t]); + _v1 = vmul_n_f32(_v1, ps[kk + 4 + t]); } - const float32x4_t _s = vcombine_f32(_scale, _scale); - const float32x4_t _v0q = vmulq_f32(vcombine_f32(_v0, _v0), _s); - const float32x4_t _v1q = vmulq_f32(vcombine_f32(_v1, _v1), _s); + float32x4_t _s = vcombine_f32(_scale, _scale); + float32x4_t _v0q = vmulq_f32(vcombine_f32(_v0, _v0), _s); + float32x4_t _v1q = vmulq_f32(vcombine_f32(_v1, _v1), _s); _q0.val[t] = float2int8(_v0q, _v0q); _q1.val[t] = float2int8(_v1q, _v1q); + ptrAk0 += A_hstep; + ptrAk1 += A_hstep; } vst4_lane_s8(pp, _q0, 0); vst4_lane_s8(pp + 4, _q1, 0); @@ -1002,18 +996,18 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A pp += 16; } #endif // __ARM_FEATURE_MATMUL_INT8 -#if __ARM_FEATURE_DOTPROD for (; kk + 3 < max_kk; kk += 4) { int8x8x4_t _q; + ptrAk = ptrAg + (size_t)kk * A_hstep; for (int t = 0; t < 4; t++) { - const int k = k0 + kk + t; - float32x2_t _v = vld1_f32(ptrA + (size_t)k * A_hstep); - if (input_scale_ptr) - _v = vmul_n_f32(_v, input_scale_ptr[k]); - const float32x4_t _vq = vmulq_f32(vcombine_f32(_v, _v), vcombine_f32(_scale, _scale)); + float32x2_t _v = vld1_f32(ptrAk); + if (ps) + _v = vmul_n_f32(_v, ps[kk + t]); + float32x4_t _vq = vmulq_f32(vcombine_f32(_v, _v), vcombine_f32(_scale, _scale)); _q.val[t] = float2int8(_vq, _vq); + ptrAk += A_hstep; } vst4_lane_s8(pp, _q, 0); vst4_lane_s8(pp + 4, _q, 1); @@ -1023,14 +1017,15 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A for (; kk + 1 < max_kk; kk += 2) { int8x8x2_t _q; + ptrAk = ptrAg + (size_t)kk * A_hstep; for (int t = 0; t < 2; t++) { - const int k = k0 + kk + t; - float32x2_t _v = vld1_f32(ptrA + (size_t)k * A_hstep); - if (input_scale_ptr) - _v = vmul_n_f32(_v, input_scale_ptr[k]); - const float32x4_t _vq = vmulq_f32(vcombine_f32(_v, _v), vcombine_f32(_scale, _scale)); + float32x2_t _v = vld1_f32(ptrAk); + if (ps) + _v = vmul_n_f32(_v, ps[kk + t]); + float32x4_t _vq = vmulq_f32(vcombine_f32(_v, _v), vcombine_f32(_scale, _scale)); _q.val[t] = float2int8(_vq, _vq); + ptrAk += A_hstep; } vst2_lane_s8(pp, _q, 0); vst2_lane_s8(pp + 2, _q, 1); @@ -1038,14 +1033,18 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A } if (kk < max_kk) { - const int k = k0 + kk; - float32x2_t _v = vld1_f32(ptrA + (size_t)k * A_hstep); - if (input_scale_ptr) - _v = vmul_n_f32(_v, input_scale_ptr[k]); - const float32x4_t _vq = vmulq_f32(vcombine_f32(_v, _v), vcombine_f32(_scale, _scale)); + ptrAk = ptrAg + (size_t)kk * A_hstep; + float32x2_t _v = vld1_f32(ptrAk); + if (ps) + _v = vmul_n_f32(_v, ps[kk]); + float32x4_t _vq = vmulq_f32(vcombine_f32(_v, _v), vcombine_f32(_scale, _scale)); vst1_lane_s16((short*)pp, vreinterpret_s16_s8(float2int8(_vq, _vq)), 0); pp += 2; } + ptrAg += (size_t)max_kk * A_hstep; + if (ps) + ps += max_kk; + descale_ptr += 2; } } #elif __ARM_FEATURE_SIMD32 && NCNN_GNU_INLINE_ASM @@ -1054,12 +1053,13 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A const float* ptrA = (const float*)A + i + ii; signed char* pp = outptr + ii * out_hstep; float* descale_ptr = descales + ii * descales_hstep; + const float* ptrAg = ptrA; + const float* ps = input_scale_ptr; for (int g = 0; g < block_count; g++) { - const int k0 = g * block_size; - const int max_kk = std::min(K - k0, block_size); - const float* ptrAk = ptrA + (size_t)k0 * A_hstep; + const int max_kk = std::min(K - g * block_size, block_size); + const float* ptrAk = ptrAg; float absmax0 = 0.f; float absmax1 = 0.f; @@ -1067,92 +1067,93 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A { float v0 = ptrAk[0]; float v1 = ptrAk[1]; - if (input_scale_ptr) + if (ps) { - v0 *= input_scale_ptr[k0 + kk]; - v1 *= input_scale_ptr[k0 + kk]; + v0 *= ps[kk]; + v1 *= ps[kk]; } absmax0 = std::max(absmax0, fabsf(v0)); absmax1 = std::max(absmax1, fabsf(v1)); ptrAk += A_hstep; } - descale_ptr[g * 2] = absmax0 / 127.f; - descale_ptr[g * 2 + 1] = absmax1 / 127.f; - volatile double scale0_fp64 = absmax0 == 0.f ? 0.0 : 127.0 / (double)absmax0; - volatile double scale1_fp64 = absmax1 == 0.f ? 0.0 : 127.0 / (double)absmax1; - const float scale0 = (float)scale0_fp64; - const float scale1 = (float)scale1_fp64; + descale_ptr[0] = absmax0 / 127.f; + descale_ptr[1] = absmax1 / 127.f; + const float scale0 = absmax0 == 0.f ? 0.f : 127.f / absmax0; + const float scale1 = absmax1 == 0.f ? 0.f : 127.f / absmax1; - ptrAk = ptrA + (size_t)k0 * A_hstep; + ptrAk = ptrAg; for (int kk = 0; kk < max_kk; kk++) { float v0 = ptrAk[0]; float v1 = ptrAk[1]; - if (input_scale_ptr) + if (ps) { - v0 *= input_scale_ptr[k0 + kk]; - v1 *= input_scale_ptr[k0 + kk]; - asm volatile("" - : "+w"(v0), "+w"(v1)); + v0 *= ps[kk]; + v1 *= ps[kk]; } *pp++ = float2int8(v0 * scale0); *pp++ = float2int8(v1 * scale1); ptrAk += A_hstep; } + ptrAg += (size_t)max_kk * A_hstep; + if (ps) + ps += max_kk; + descale_ptr += 2; } } #endif // __ARM_NEON for (; ii < max_ii; ii++) { const float* ptrA = (const float*)A + i + ii; - signed char* outptr0 = outptr + ii * out_hstep; + signed char* pp = outptr + ii * out_hstep; float* descale_ptr = descales + ii * descales_hstep; + const float* ptrAg = ptrA; + const float* ps = input_scale_ptr; for (int g = 0; g < block_count; g++) { - const int k0 = g * block_size; - const int max_kk = std::min(K - k0, block_size); + const int max_kk = std::min(K - g * block_size, block_size); float absmax = 0.f; + const float* ptrAk = ptrAg; for (int kk = 0; kk < max_kk; kk++) { - const int k = k0 + kk; - float v = ptrA[(size_t)k * A_hstep]; - if (input_scale_ptr) - v *= input_scale_ptr[k]; + float v = *ptrAk; + if (ps) + v *= ps[kk]; absmax = std::max(absmax, fabsf(v)); + ptrAk += A_hstep; } if (absmax == 0.f) { - descale_ptr[g] = 0.f; + *descale_ptr++ = 0.f; for (int k = 0; k < max_kk; k++) - outptr0[k0 + k] = 0; + *pp++ = 0; + ptrAg += (size_t)max_kk * A_hstep; + if (ps) + ps += max_kk; continue; } - volatile double scale_fp64 = 127.0 / (double)absmax; - const float scale = (float)scale_fp64; - descale_ptr[g] = absmax / 127.f; + const float scale = 127.f / absmax; + *descale_ptr++ = absmax / 127.f; + ptrAk = ptrAg; for (int kk = 0; kk < max_kk; kk++) { - const int k = k0 + kk; - float v = ptrA[(size_t)k * A_hstep]; - if (input_scale_ptr) + float v = *ptrAk; + if (ps) { - v *= input_scale_ptr[k]; -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+w"(v)); -#else - volatile float v_ordered = v; - v = v_ordered; -#endif + v *= ps[kk]; } - outptr0[k] = float2int8(v * scale); + *pp++ = float2int8(v * scale); + ptrAk += A_hstep; } + ptrAg += (size_t)max_kk * A_hstep; + if (ps) + ps += max_kk; } } } @@ -1205,22 +1206,26 @@ static int pack_B_wq_int8(const Mat& B, const Mat& B_scales, Mat& BT, Mat& BT_de const int j = panel_start4 + p * 4; signed char* pp = (signed char*)BT_packed + j * K; float* pd = (float*)BT_packed_descales + j * block_count; + const signed char* p0 = B.row(j); + const signed char* p1 = B.row(j + 1); + const signed char* p2 = B.row(j + 2); + const signed char* p3 = B.row(j + 3); + const float* ps0 = B_scales.row(j); + const float* ps1 = B_scales.row(j + 1); + const float* ps2 = B_scales.row(j + 2); + const float* ps3 = B_scales.row(j + 3); for (int g = 0; g < block_count; g++) { const int k0 = g * block_size; const int max_kk = std::min(K - k0, block_size); - const signed char* p0 = B.row(j) + k0; - const signed char* p1 = B.row(j + 1) + k0; - const signed char* p2 = B.row(j + 2) + k0; - const signed char* p3 = B.row(j + 3) + k0; int kk = 0; for (; kk + 15 < max_kk; kk += 16) { - const int8x16_t _p0 = vld1q_s8(p0); - const int8x16_t _p1 = vld1q_s8(p1); - const int8x16_t _p2 = vld1q_s8(p2); - const int8x16_t _p3 = vld1q_s8(p3); + int8x16_t _p0 = vld1q_s8(p0); + int8x16_t _p1 = vld1q_s8(p1); + int8x16_t _p2 = vld1q_s8(p2); + int8x16_t _p3 = vld1q_s8(p3); #if __ARM_FEATURE_DOTPROD #if __ARM_FEATURE_MATMUL_INT8 int64x2x4_t _r0123; @@ -1253,10 +1258,10 @@ static int pack_B_wq_int8(const Mat& B, const Mat& B_scales, Mat& BT, Mat& BT_de } for (; kk + 7 < max_kk; kk += 8) { - const int8x8_t _p0 = vld1_s8(p0); - const int8x8_t _p1 = vld1_s8(p1); - const int8x8_t _p2 = vld1_s8(p2); - const int8x8_t _p3 = vld1_s8(p3); + int8x8_t _p0 = vld1_s8(p0); + int8x8_t _p1 = vld1_s8(p1); + int8x8_t _p2 = vld1_s8(p2); + int8x8_t _p3 = vld1_s8(p3); #if __ARM_FEATURE_DOTPROD #if __ARM_FEATURE_MATMUL_INT8 vst1q_s8(pp, vcombine_s8(_p0, _p1)); @@ -1332,10 +1337,16 @@ static int pack_B_wq_int8(const Mat& B, const Mat& B_scales, Mat& BT, Mat& BT_de pp[2] = p2[0]; pp[3] = p3[0]; pp += 4; + p0++; + p1++; + p2++; + p3++; } - for (int jj = 0; jj < 4; jj++) - pd[g * 4 + jj] = 1.f / B_scales.row(j + jj)[g]; + *pd++ = 1.f / *ps0++; + *pd++ = 1.f / *ps1++; + *pd++ = 1.f / *ps2++; + *pd++ = 1.f / *ps3++; } } #endif @@ -1345,19 +1356,21 @@ static int pack_B_wq_int8(const Mat& B, const Mat& B_scales, Mat& BT, Mat& BT_de const int j = panel_start2 + p * 2; signed char* pp = (signed char*)BT_packed + j * K; float* pd = (float*)BT_packed_descales + j * block_count; + const signed char* p0 = B.row(j); + const signed char* p1 = B.row(j + 1); + const float* ps0 = B_scales.row(j); + const float* ps1 = B_scales.row(j + 1); for (int g = 0; g < block_count; g++) { const int k0 = g * block_size; const int max_kk = std::min(K - k0, block_size); - const signed char* p0 = B.row(j) + k0; - const signed char* p1 = B.row(j + 1) + k0; int kk = 0; #if __ARM_NEON for (; kk + 15 < max_kk; kk += 16) { - const int8x16_t _p0 = vld1q_s8(p0); - const int8x16_t _p1 = vld1q_s8(p1); + int8x16_t _p0 = vld1q_s8(p0); + int8x16_t _p1 = vld1q_s8(p1); #if __ARM_FEATURE_DOTPROD #if __ARM_FEATURE_MATMUL_INT8 int64x2x2_t _r01; @@ -1382,8 +1395,8 @@ static int pack_B_wq_int8(const Mat& B, const Mat& B_scales, Mat& BT, Mat& BT_de } for (; kk + 7 < max_kk; kk += 8) { - const int8x8_t _p0 = vld1_s8(p0); - const int8x8_t _p1 = vld1_s8(p1); + int8x8_t _p0 = vld1_s8(p0); + int8x8_t _p1 = vld1_s8(p1); #if __ARM_FEATURE_DOTPROD #if __ARM_FEATURE_MATMUL_INT8 vst1q_s8(pp, vcombine_s8(_p0, _p1)); @@ -1403,7 +1416,6 @@ static int pack_B_wq_int8(const Mat& B, const Mat& B_scales, Mat& BT, Mat& BT_de p0 += 8; p1 += 8; } -#endif // __ARM_NEON #if __ARM_FEATURE_DOTPROD for (; kk + 3 < max_kk; kk += 4) { @@ -1420,6 +1432,7 @@ static int pack_B_wq_int8(const Mat& B, const Mat& B_scales, Mat& BT, Mat& BT_de p1 += 4; } #endif // __ARM_FEATURE_DOTPROD +#endif // __ARM_NEON for (; kk + 1 < max_kk; kk += 2) { #if !__ARM_NEON && __ARM_FEATURE_SIMD32 && NCNN_GNU_INLINE_ASM @@ -1443,10 +1456,12 @@ static int pack_B_wq_int8(const Mat& B, const Mat& B_scales, Mat& BT, Mat& BT_de pp[0] = p0[0]; pp[1] = p1[0]; pp += 2; + p0++; + p1++; } - for (int jj = 0; jj < 2; jj++) - pd[g * 2 + jj] = 1.f / B_scales.row(j + jj)[g]; + *pd++ = 1.f / *ps0++; + *pd++ = 1.f / *ps1++; } } #pragma omp for @@ -1455,12 +1470,13 @@ static int pack_B_wq_int8(const Mat& B, const Mat& B_scales, Mat& BT, Mat& BT_de const int j = panel_start1 + p * 1; signed char* pp = (signed char*)BT_packed + j * K; float* pd = (float*)BT_packed_descales + j * block_count; + const signed char* p0 = B.row(j); + const float* ps0 = B_scales.row(j); for (int g = 0; g < block_count; g++) { const int k0 = g * block_size; const int max_kk = std::min(K - k0, block_size); - const signed char* p0 = B.row(j) + k0; int kk = 0; #if __ARM_NEON for (; kk + 15 < max_kk; kk += 16) @@ -1475,7 +1491,6 @@ static int pack_B_wq_int8(const Mat& B, const Mat& B_scales, Mat& BT, Mat& BT_de pp += 8; p0 += 8; } -#endif // __ARM_NEON #if __ARM_FEATURE_DOTPROD for (; kk + 3 < max_kk; kk += 4) { @@ -1487,6 +1502,7 @@ static int pack_B_wq_int8(const Mat& B, const Mat& B_scales, Mat& BT, Mat& BT_de p0 += 4; } #endif // __ARM_FEATURE_DOTPROD +#endif // __ARM_NEON for (; kk + 1 < max_kk; kk += 2) { pp[0] = p0[0]; @@ -1495,10 +1511,9 @@ static int pack_B_wq_int8(const Mat& B, const Mat& B_scales, Mat& BT, Mat& BT_de p0 += 2; } if (kk < max_kk) - *pp++ = p0[0]; + *pp++ = *p0++; - for (int jj = 0; jj < 1; jj++) - pd[g * 1 + jj] = 1.f / B_scales.row(j + jj)[g]; + *pd++ = 1.f / *ps0++; } } } @@ -1508,26 +1523,26 @@ static int pack_B_wq_int8(const Mat& B, const Mat& B_scales, Mat& BT, Mat& BT_de return 0; } -static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int block_size) +static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int full_K, int k, int max_kk, int block_size) { #if NCNN_RUNTIME_CPU && NCNN_ARM86SVEI8MM && __aarch64__ && !__ARM_FEATURE_SVE_MATMUL_INT8 if (ncnn::cpu_support_arm_svei8mm()) { - gemm_transB_packed_tile_wq_int8_svei8mm(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); + gemm_transB_packed_tile_wq_int8_svei8mm(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, full_K, k, max_kk, block_size); return; } #endif #if NCNN_RUNTIME_CPU && NCNN_ARM84I8MM && __aarch64__ && !__ARM_FEATURE_MATMUL_INT8 if (ncnn::cpu_support_arm_i8mm()) { - gemm_transB_packed_tile_wq_int8_i8mm(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); + gemm_transB_packed_tile_wq_int8_i8mm(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, full_K, k, max_kk, block_size); return; } #endif #if NCNN_RUNTIME_CPU && NCNN_ARM82DOT && __aarch64__ && !__ARM_FEATURE_DOTPROD && !__ARM_FEATURE_MATMUL_INT8 if (ncnn::cpu_support_arm_asimddp()) { - gemm_transB_packed_tile_wq_int8_asimddp(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); + gemm_transB_packed_tile_wq_int8_asimddp(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, full_K, k, max_kk, block_size); return; } #endif @@ -1538,6 +1553,10 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de const int A_descales_hstep = AT_descales_tile.w; const signed char* pBT = BT_tile; const float* pBT_descales = BT_descales_tile; + const int block_count = (full_K + block_size - 1) / block_size; + const int block_start = k / block_size; + const int k0 = k; + const int K = max_kk; float* outptr = topT_tile; @@ -1548,22 +1567,23 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de { const signed char* pA_block = pAT; const float* pA_descales_block = pAT_descales; - const signed char* pB = pBT; - const float* pB_descales = pBT_descales; - int jj = 0; + const signed char* pB4 = pBT + 4 * k0; + const float* pB_descales4 = pBT_descales + 4 * block_start; for (; jj + 3 < max_jj; jj += 4) { + const signed char* pB = pB4; + const float* pB_descales = pB_descales4; const signed char* pA = pA_block; const float* pA_descales = pA_descales_block; - float32x4_t _fsum0 = vdupq_n_f32(0.f); - float32x4_t _fsum1 = vdupq_n_f32(0.f); - float32x4_t _fsum2 = vdupq_n_f32(0.f); - float32x4_t _fsum3 = vdupq_n_f32(0.f); - float32x4_t _fsum4 = vdupq_n_f32(0.f); - float32x4_t _fsum5 = vdupq_n_f32(0.f); - float32x4_t _fsum6 = vdupq_n_f32(0.f); - float32x4_t _fsum7 = vdupq_n_f32(0.f); + float32x4_t _fsum0 = k0 == 0 ? vdupq_n_f32(0.f) : vld1q_f32(outptr); + float32x4_t _fsum1 = k0 == 0 ? vdupq_n_f32(0.f) : vld1q_f32(outptr + 4); + float32x4_t _fsum2 = k0 == 0 ? vdupq_n_f32(0.f) : vld1q_f32(outptr + 8); + float32x4_t _fsum3 = k0 == 0 ? vdupq_n_f32(0.f) : vld1q_f32(outptr + 12); + float32x4_t _fsum4 = k0 == 0 ? vdupq_n_f32(0.f) : vld1q_f32(outptr + 16); + float32x4_t _fsum5 = k0 == 0 ? vdupq_n_f32(0.f) : vld1q_f32(outptr + 20); + float32x4_t _fsum6 = k0 == 0 ? vdupq_n_f32(0.f) : vld1q_f32(outptr + 24); + float32x4_t _fsum7 = k0 == 0 ? vdupq_n_f32(0.f) : vld1q_f32(outptr + 28); for (int k = 0; k < K; k += block_size) { @@ -1577,6 +1597,7 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de int32x4_t _sum7 = vdupq_n_s32(0); const int max_kk = std::min(K - k, block_size); int kk = 0; +#if __ARM_FEATURE_DOTPROD #if __ARM_FEATURE_MATMUL_INT8 int32x4_t _msum0 = vdupq_n_s32(0); int32x4_t _msum1 = vdupq_n_s32(0); @@ -1588,12 +1609,12 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de int32x4_t _msum7 = vdupq_n_s32(0); for (; kk + 7 < max_kk; kk += 8) { - const int8x16_t _b0 = vld1q_s8(pB); - const int8x16_t _b1 = vld1q_s8(pB + 16); - const int8x16_t _a01 = vld1q_s8(pA); - const int8x16_t _a23 = vld1q_s8(pA + 16); - const int8x16_t _a45 = vld1q_s8(pA + 32); - const int8x16_t _a67 = vld1q_s8(pA + 48); + int8x16_t _b0 = vld1q_s8(pB); + int8x16_t _b1 = vld1q_s8(pB + 16); + int8x16_t _a01 = vld1q_s8(pA); + int8x16_t _a23 = vld1q_s8(pA + 16); + int8x16_t _a45 = vld1q_s8(pA + 32); + int8x16_t _a67 = vld1q_s8(pA + 48); _msum0 = vmmlaq_s32(_msum0, _a01, _b0); _msum1 = vmmlaq_s32(_msum1, _a23, _b0); _msum2 = vmmlaq_s32(_msum2, _a01, _b1); @@ -1614,12 +1635,11 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de _sum6 = vcombine_s32(vget_low_s32(_msum5), vget_low_s32(_msum7)); _sum7 = vcombine_s32(vget_high_s32(_msum5), vget_high_s32(_msum7)); #endif // __ARM_FEATURE_MATMUL_INT8 -#if __ARM_FEATURE_DOTPROD for (; kk + 3 < max_kk; kk += 4) { - const int8x16_t _b0 = vld1q_s8(pB); - const int8x16_t _a0 = vld1q_s8(pA); - const int8x16_t _a1 = vld1q_s8(pA + 16); + int8x16_t _b0 = vld1q_s8(pB); + int8x16_t _a0 = vld1q_s8(pA); + int8x16_t _a1 = vld1q_s8(pA + 16); _sum0 = vdotq_laneq_s32(_sum0, _b0, _a0, 0); _sum1 = vdotq_laneq_s32(_sum1, _b0, _a0, 1); _sum2 = vdotq_laneq_s32(_sum2, _b0, _a0, 2); @@ -1631,13 +1651,112 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pA += 32; pB += 16; } +#else // __ARM_FEATURE_DOTPROD +#if NCNN_GNU_INLINE_ASM + { + int nn = (max_kk - kk) >> 2; + const int remain = nn; + asm volatile( + "cmp %w2, #0 \n" + "beq 1f \n" + "0: \n" + "ld1 {v0.16b, v1.16b}, [%0], #32 \n" + "ld1 {v2.16b}, [%1], #16 \n" + "dup v3.8h, v0.h[0] \n" + "smull v4.8h, v2.8b, v3.8b \n" + "dup v3.8h, v1.h[0] \n" + "smlal2 v4.8h, v2.16b, v3.16b \n" + "sadalp %3.4s, v4.8h \n" + "dup v3.8h, v0.h[1] \n" + "smull v4.8h, v2.8b, v3.8b \n" + "dup v3.8h, v1.h[1] \n" + "smlal2 v4.8h, v2.16b, v3.16b \n" + "sadalp %4.4s, v4.8h \n" + "dup v3.8h, v0.h[2] \n" + "smull v4.8h, v2.8b, v3.8b \n" + "dup v3.8h, v1.h[2] \n" + "smlal2 v4.8h, v2.16b, v3.16b \n" + "sadalp %5.4s, v4.8h \n" + "dup v3.8h, v0.h[3] \n" + "smull v4.8h, v2.8b, v3.8b \n" + "dup v3.8h, v1.h[3] \n" + "smlal2 v4.8h, v2.16b, v3.16b \n" + "sadalp %6.4s, v4.8h \n" + "dup v3.8h, v0.h[4] \n" + "smull v4.8h, v2.8b, v3.8b \n" + "dup v3.8h, v1.h[4] \n" + "smlal2 v4.8h, v2.16b, v3.16b \n" + "sadalp %7.4s, v4.8h \n" + "dup v3.8h, v0.h[5] \n" + "smull v4.8h, v2.8b, v3.8b \n" + "dup v3.8h, v1.h[5] \n" + "smlal2 v4.8h, v2.16b, v3.16b \n" + "sadalp %8.4s, v4.8h \n" + "dup v3.8h, v0.h[6] \n" + "smull v4.8h, v2.8b, v3.8b \n" + "dup v3.8h, v1.h[6] \n" + "smlal2 v4.8h, v2.16b, v3.16b \n" + "sadalp %9.4s, v4.8h \n" + "dup v3.8h, v0.h[7] \n" + "smull v4.8h, v2.8b, v3.8b \n" + "dup v3.8h, v1.h[7] \n" + "smlal2 v4.8h, v2.16b, v3.16b \n" + "subs %w2, %w2, #1 \n" + "sadalp %10.4s, v4.8h \n" + "bne 0b \n" + "1: \n" + : "+r"(pA), "+r"(pB), "+r"(nn), "+w"(_sum0), "+w"(_sum1), "+w"(_sum2), "+w"(_sum3), "+w"(_sum4), "+w"(_sum5), "+w"(_sum6), "+w"(_sum7) + : + : "cc", "memory", "v0", "v1", "v2", "v3", "v4"); + kk += remain * 4; + } +#else // NCNN_GNU_INLINE_ASM + for (; kk + 3 < max_kk; kk += 4) + { + int16x8_t _a01 = vreinterpretq_s16_s8(vld1q_s8(pA)); + int16x8_t _a23 = vreinterpretq_s16_s8(vld1q_s8(pA + 16)); + int8x16_t _b = vld1q_s8(pB); + int8x8_t _b01 = vget_low_s8(_b); + int8x8_t _b23 = vget_high_s8(_b); + int16x4_t _a010 = vget_low_s16(_a01); + int16x4_t _a011 = vget_high_s16(_a01); + int16x4_t _a230 = vget_low_s16(_a23); + int16x4_t _a231 = vget_high_s16(_a23); + int16x8_t _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a010, 0))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a230, 0))); + _sum0 = vpadalq_s16(_sum0, _s); + _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a010, 1))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a230, 1))); + _sum1 = vpadalq_s16(_sum1, _s); + _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a010, 2))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a230, 2))); + _sum2 = vpadalq_s16(_sum2, _s); + _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a010, 3))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a230, 3))); + _sum3 = vpadalq_s16(_sum3, _s); + _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a011, 0))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a231, 0))); + _sum4 = vpadalq_s16(_sum4, _s); + _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a011, 1))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a231, 1))); + _sum5 = vpadalq_s16(_sum5, _s); + _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a011, 2))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a231, 2))); + _sum6 = vpadalq_s16(_sum6, _s); + _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a011, 3))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a231, 3))); + _sum7 = vpadalq_s16(_sum7, _s); + pA += 32; + pB += 16; + } +#endif // NCNN_GNU_INLINE_ASM #endif // __ARM_FEATURE_DOTPROD for (; kk + 1 < max_kk; kk += 2) { - const int8x8_t _b0 = vld1_s8(pB); - const int16x8_t _a = vreinterpretq_s16_s8(vld1q_s8(pA)); - const int16x4_t _a0 = vget_low_s16(_a); - const int16x4_t _a1 = vget_high_s16(_a); + int8x8_t _b0 = vld1_s8(pB); + int16x8_t _a = vreinterpretq_s16_s8(vld1q_s8(pA)); + int16x4_t _a0 = vget_low_s16(_a); + int16x4_t _a1 = vget_high_s16(_a); _sum0 = vaddq_s32(_sum0, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a0, 0))))); _sum1 = vaddq_s32(_sum1, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a0, 1))))); _sum2 = vaddq_s32(_sum2, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a0, 2))))); @@ -1651,31 +1770,31 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } if (kk < max_kk) { - const int8x8_t _b0 = vreinterpret_s8_s32(vld1_dup_s32((const int*)pB)); - const int8x8_t _a = vld1_s8(pA); - const int16x8_t _p0 = vmull_s8(_b0, vdup_lane_s8(_a, 0)); - const int16x8_t _p1 = vmull_s8(_b0, vdup_lane_s8(_a, 1)); + int8x8_t _b0 = vreinterpret_s8_s32(vld1_dup_s32((const int*)pB)); + int8x8_t _a = vld1_s8(pA); + int16x8_t _p0 = vmull_s8(_b0, vdup_lane_s8(_a, 0)); + int16x8_t _p1 = vmull_s8(_b0, vdup_lane_s8(_a, 1)); _sum0 = vaddq_s32(_sum0, vmovl_s16(vget_low_s16(_p0))); _sum1 = vaddq_s32(_sum1, vmovl_s16(vget_low_s16(_p1))); - const int16x8_t _p2 = vmull_s8(_b0, vdup_lane_s8(_a, 2)); - const int16x8_t _p3 = vmull_s8(_b0, vdup_lane_s8(_a, 3)); + int16x8_t _p2 = vmull_s8(_b0, vdup_lane_s8(_a, 2)); + int16x8_t _p3 = vmull_s8(_b0, vdup_lane_s8(_a, 3)); _sum2 = vaddq_s32(_sum2, vmovl_s16(vget_low_s16(_p2))); _sum3 = vaddq_s32(_sum3, vmovl_s16(vget_low_s16(_p3))); - const int16x8_t _p4 = vmull_s8(_b0, vdup_lane_s8(_a, 4)); - const int16x8_t _p5 = vmull_s8(_b0, vdup_lane_s8(_a, 5)); + int16x8_t _p4 = vmull_s8(_b0, vdup_lane_s8(_a, 4)); + int16x8_t _p5 = vmull_s8(_b0, vdup_lane_s8(_a, 5)); _sum4 = vaddq_s32(_sum4, vmovl_s16(vget_low_s16(_p4))); _sum5 = vaddq_s32(_sum5, vmovl_s16(vget_low_s16(_p5))); - const int16x8_t _p6 = vmull_s8(_b0, vdup_lane_s8(_a, 6)); - const int16x8_t _p7 = vmull_s8(_b0, vdup_lane_s8(_a, 7)); + int16x8_t _p6 = vmull_s8(_b0, vdup_lane_s8(_a, 6)); + int16x8_t _p7 = vmull_s8(_b0, vdup_lane_s8(_a, 7)); _sum6 = vaddq_s32(_sum6, vmovl_s16(vget_low_s16(_p6))); _sum7 = vaddq_s32(_sum7, vmovl_s16(vget_low_s16(_p7))); pA += 8; pB += 4; } - const float32x4_t _bd0 = vld1q_f32(pB_descales); - const float32x4_t _ad0 = vld1q_f32(pA_descales); - const float32x4_t _ad1 = vld1q_f32(pA_descales + 4); + float32x4_t _bd0 = vld1q_f32(pB_descales); + float32x4_t _ad0 = vld1q_f32(pA_descales); + float32x4_t _ad1 = vld1q_f32(pA_descales + 4); _fsum0 = vmlaq_f32(_fsum0, vcvtq_f32_s32(_sum0), vmulq_laneq_f32(_bd0, _ad0, 0)); _fsum1 = vmlaq_f32(_fsum1, vcvtq_f32_s32(_sum1), vmulq_laneq_f32(_bd0, _ad0, 1)); _fsum2 = vmlaq_f32(_fsum2, vcvtq_f32_s32(_sum2), vmulq_laneq_f32(_bd0, _ad0, 2)); @@ -1697,19 +1816,25 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de vst1q_f32(outptr + 24, _fsum6); vst1q_f32(outptr + 28, _fsum7); outptr += 32; + pB4 += (size_t)4 * full_K; + pB_descales4 += (size_t)4 * block_count; } + const signed char* pB2 = pB4 - 2 * k0; + const float* pB_descales2 = pB_descales4 - 2 * block_start; for (; jj + 1 < max_jj; jj += 2) { + const signed char* pB = pB2; + const float* pB_descales = pB_descales2; const signed char* pA = pA_block; const float* pA_descales = pA_descales_block; - float32x4_t _fsum0 = vdupq_n_f32(0.f); - float32x4_t _fsum1 = vdupq_n_f32(0.f); - float32x4_t _fsum2 = vdupq_n_f32(0.f); - float32x4_t _fsum3 = vdupq_n_f32(0.f); - float32x4_t _fsum4 = vdupq_n_f32(0.f); - float32x4_t _fsum5 = vdupq_n_f32(0.f); - float32x4_t _fsum6 = vdupq_n_f32(0.f); - float32x4_t _fsum7 = vdupq_n_f32(0.f); + float32x4_t _fsum0 = k0 == 0 ? vdupq_n_f32(0.f) : vcombine_f32(vld1_f32(outptr), vdup_n_f32(0.f)); + float32x4_t _fsum1 = k0 == 0 ? vdupq_n_f32(0.f) : vcombine_f32(vld1_f32(outptr + 2), vdup_n_f32(0.f)); + float32x4_t _fsum2 = k0 == 0 ? vdupq_n_f32(0.f) : vcombine_f32(vld1_f32(outptr + 4), vdup_n_f32(0.f)); + float32x4_t _fsum3 = k0 == 0 ? vdupq_n_f32(0.f) : vcombine_f32(vld1_f32(outptr + 6), vdup_n_f32(0.f)); + float32x4_t _fsum4 = k0 == 0 ? vdupq_n_f32(0.f) : vcombine_f32(vld1_f32(outptr + 8), vdup_n_f32(0.f)); + float32x4_t _fsum5 = k0 == 0 ? vdupq_n_f32(0.f) : vcombine_f32(vld1_f32(outptr + 10), vdup_n_f32(0.f)); + float32x4_t _fsum6 = k0 == 0 ? vdupq_n_f32(0.f) : vcombine_f32(vld1_f32(outptr + 12), vdup_n_f32(0.f)); + float32x4_t _fsum7 = k0 == 0 ? vdupq_n_f32(0.f) : vcombine_f32(vld1_f32(outptr + 14), vdup_n_f32(0.f)); for (int k = 0; k < K; k += block_size) { @@ -1723,6 +1848,7 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de int32x4_t _sum7 = vdupq_n_s32(0); const int max_kk = std::min(K - k, block_size); int kk = 0; +#if __ARM_FEATURE_DOTPROD #if __ARM_FEATURE_MATMUL_INT8 int32x4_t _msum0 = vdupq_n_s32(0); int32x4_t _msum1 = vdupq_n_s32(0); @@ -1730,11 +1856,11 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de int32x4_t _msum3 = vdupq_n_s32(0); for (; kk + 7 < max_kk; kk += 8) { - const int8x16_t _b0 = vld1q_s8(pB); - const int8x16_t _a01 = vld1q_s8(pA); - const int8x16_t _a23 = vld1q_s8(pA + 16); - const int8x16_t _a45 = vld1q_s8(pA + 32); - const int8x16_t _a67 = vld1q_s8(pA + 48); + int8x16_t _b0 = vld1q_s8(pB); + int8x16_t _a01 = vld1q_s8(pA); + int8x16_t _a23 = vld1q_s8(pA + 16); + int8x16_t _a45 = vld1q_s8(pA + 32); + int8x16_t _a67 = vld1q_s8(pA + 48); _msum0 = vmmlaq_s32(_msum0, _a01, _b0); _msum1 = vmmlaq_s32(_msum1, _a23, _b0); _msum2 = vmmlaq_s32(_msum2, _a45, _b0); @@ -1751,12 +1877,11 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de _sum6 = vcombine_s32(vget_low_s32(_msum3), vdup_n_s32(0)); _sum7 = vcombine_s32(vget_high_s32(_msum3), vdup_n_s32(0)); #endif // __ARM_FEATURE_MATMUL_INT8 -#if __ARM_FEATURE_DOTPROD for (; kk + 3 < max_kk; kk += 4) { - const int8x16_t _b0 = vcombine_s8(vld1_s8(pB), vdup_n_s8(0)); - const int8x16_t _a0 = vld1q_s8(pA); - const int8x16_t _a1 = vld1q_s8(pA + 16); + int8x16_t _b0 = vcombine_s8(vld1_s8(pB), vdup_n_s8(0)); + int8x16_t _a0 = vld1q_s8(pA); + int8x16_t _a1 = vld1q_s8(pA + 16); _sum0 = vdotq_laneq_s32(_sum0, _b0, _a0, 0); _sum1 = vdotq_laneq_s32(_sum1, _b0, _a0, 1); _sum2 = vdotq_laneq_s32(_sum2, _b0, _a0, 2); @@ -1768,13 +1893,115 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pA += 32; pB += 8; } +#else // __ARM_FEATURE_DOTPROD +#if NCNN_GNU_INLINE_ASM + { + int nn = (max_kk - kk) >> 2; + const int remain = nn; + asm volatile( + "cmp %w2, #0 \n" + "beq 1f \n" + "0: \n" + "ld1 {v0.16b, v1.16b}, [%0], #32 \n" + "ld1 {v2.8b}, [%1], #8 \n" + "dup v3.2s, v2.s[0] \n" + "dup v4.2s, v2.s[1] \n" + "dup v5.8h, v0.h[0] \n" + "smull v6.8h, v3.8b, v5.8b \n" + "dup v5.8h, v1.h[0] \n" + "smlal v6.8h, v4.8b, v5.8b \n" + "sadalp %3.4s, v6.8h \n" + "dup v5.8h, v0.h[1] \n" + "smull v6.8h, v3.8b, v5.8b \n" + "dup v5.8h, v1.h[1] \n" + "smlal v6.8h, v4.8b, v5.8b \n" + "sadalp %4.4s, v6.8h \n" + "dup v5.8h, v0.h[2] \n" + "smull v6.8h, v3.8b, v5.8b \n" + "dup v5.8h, v1.h[2] \n" + "smlal v6.8h, v4.8b, v5.8b \n" + "sadalp %5.4s, v6.8h \n" + "dup v5.8h, v0.h[3] \n" + "smull v6.8h, v3.8b, v5.8b \n" + "dup v5.8h, v1.h[3] \n" + "smlal v6.8h, v4.8b, v5.8b \n" + "sadalp %6.4s, v6.8h \n" + "dup v5.8h, v0.h[4] \n" + "smull v6.8h, v3.8b, v5.8b \n" + "dup v5.8h, v1.h[4] \n" + "smlal v6.8h, v4.8b, v5.8b \n" + "sadalp %7.4s, v6.8h \n" + "dup v5.8h, v0.h[5] \n" + "smull v6.8h, v3.8b, v5.8b \n" + "dup v5.8h, v1.h[5] \n" + "smlal v6.8h, v4.8b, v5.8b \n" + "sadalp %8.4s, v6.8h \n" + "dup v5.8h, v0.h[6] \n" + "smull v6.8h, v3.8b, v5.8b \n" + "dup v5.8h, v1.h[6] \n" + "smlal v6.8h, v4.8b, v5.8b \n" + "sadalp %9.4s, v6.8h \n" + "dup v5.8h, v0.h[7] \n" + "smull v6.8h, v3.8b, v5.8b \n" + "dup v5.8h, v1.h[7] \n" + "smlal v6.8h, v4.8b, v5.8b \n" + "subs %w2, %w2, #1 \n" + "sadalp %10.4s, v6.8h \n" + "bne 0b \n" + "1: \n" + : "+r"(pA), "+r"(pB), "+r"(nn), "+w"(_sum0), "+w"(_sum1), "+w"(_sum2), "+w"(_sum3), "+w"(_sum4), "+w"(_sum5), "+w"(_sum6), "+w"(_sum7) + : + : "cc", "memory", "v0", "v1", "v2", "v3", "v4", "v5", "v6"); + kk += remain * 4; + } +#else // NCNN_GNU_INLINE_ASM + for (; kk + 3 < max_kk; kk += 4) + { + int16x8_t _a01 = vreinterpretq_s16_s8(vld1q_s8(pA)); + int16x8_t _a23 = vreinterpretq_s16_s8(vld1q_s8(pA + 16)); + int8x8_t _b = vld1_s8(pB); + int32x2_t _b02 = vreinterpret_s32_s8(_b); + int8x8_t _b01 = vreinterpret_s8_s32(vdup_lane_s32(_b02, 0)); + int8x8_t _b23 = vreinterpret_s8_s32(vdup_lane_s32(_b02, 1)); + int16x4_t _a010 = vget_low_s16(_a01); + int16x4_t _a011 = vget_high_s16(_a01); + int16x4_t _a230 = vget_low_s16(_a23); + int16x4_t _a231 = vget_high_s16(_a23); + int16x8_t _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a010, 0))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a230, 0))); + _sum0 = vpadalq_s16(_sum0, _s); + _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a010, 1))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a230, 1))); + _sum1 = vpadalq_s16(_sum1, _s); + _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a010, 2))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a230, 2))); + _sum2 = vpadalq_s16(_sum2, _s); + _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a010, 3))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a230, 3))); + _sum3 = vpadalq_s16(_sum3, _s); + _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a011, 0))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a231, 0))); + _sum4 = vpadalq_s16(_sum4, _s); + _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a011, 1))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a231, 1))); + _sum5 = vpadalq_s16(_sum5, _s); + _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a011, 2))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a231, 2))); + _sum6 = vpadalq_s16(_sum6, _s); + _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a011, 3))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a231, 3))); + _sum7 = vpadalq_s16(_sum7, _s); + pA += 32; + pB += 8; + } +#endif // NCNN_GNU_INLINE_ASM #endif // __ARM_FEATURE_DOTPROD for (; kk + 1 < max_kk; kk += 2) { - const int8x8_t _b0 = vreinterpret_s8_s32(vld1_dup_s32((const int*)pB)); - const int16x8_t _a = vreinterpretq_s16_s8(vld1q_s8(pA)); - const int16x4_t _a0 = vget_low_s16(_a); - const int16x4_t _a1 = vget_high_s16(_a); + int8x8_t _b0 = vreinterpret_s8_s32(vld1_dup_s32((const int*)pB)); + int16x8_t _a = vreinterpretq_s16_s8(vld1q_s8(pA)); + int16x4_t _a0 = vget_low_s16(_a); + int16x4_t _a1 = vget_high_s16(_a); _sum0 = vaddq_s32(_sum0, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a0, 0))))); _sum1 = vaddq_s32(_sum1, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a0, 1))))); _sum2 = vaddq_s32(_sum2, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a0, 2))))); @@ -1788,16 +2015,16 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } if (kk < max_kk) { - const int8x8_t _b0 = vreinterpret_s8_s16(vld1_dup_s16((const short*)pB)); - const int8x8_t _a = vld1_s8(pA); - const int16x8_t _p0 = vmull_s8(_b0, vdup_lane_s8(_a, 0)); - const int16x8_t _p1 = vmull_s8(_b0, vdup_lane_s8(_a, 1)); - const int16x8_t _p2 = vmull_s8(_b0, vdup_lane_s8(_a, 2)); - const int16x8_t _p3 = vmull_s8(_b0, vdup_lane_s8(_a, 3)); - const int16x8_t _p4 = vmull_s8(_b0, vdup_lane_s8(_a, 4)); - const int16x8_t _p5 = vmull_s8(_b0, vdup_lane_s8(_a, 5)); - const int16x8_t _p6 = vmull_s8(_b0, vdup_lane_s8(_a, 6)); - const int16x8_t _p7 = vmull_s8(_b0, vdup_lane_s8(_a, 7)); + int8x8_t _b0 = vreinterpret_s8_s16(vld1_dup_s16((const short*)pB)); + int8x8_t _a = vld1_s8(pA); + int16x8_t _p0 = vmull_s8(_b0, vdup_lane_s8(_a, 0)); + int16x8_t _p1 = vmull_s8(_b0, vdup_lane_s8(_a, 1)); + int16x8_t _p2 = vmull_s8(_b0, vdup_lane_s8(_a, 2)); + int16x8_t _p3 = vmull_s8(_b0, vdup_lane_s8(_a, 3)); + int16x8_t _p4 = vmull_s8(_b0, vdup_lane_s8(_a, 4)); + int16x8_t _p5 = vmull_s8(_b0, vdup_lane_s8(_a, 5)); + int16x8_t _p6 = vmull_s8(_b0, vdup_lane_s8(_a, 6)); + int16x8_t _p7 = vmull_s8(_b0, vdup_lane_s8(_a, 7)); _sum0 = vaddq_s32(_sum0, vmovl_s16(vget_low_s16(_p0))); _sum1 = vaddq_s32(_sum1, vmovl_s16(vget_low_s16(_p1))); _sum2 = vaddq_s32(_sum2, vmovl_s16(vget_low_s16(_p2))); @@ -1810,9 +2037,9 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB += 2; } - const float32x4_t _bd = vcombine_f32(vld1_f32(pB_descales), vdup_n_f32(0.f)); - const float32x4_t _ad0 = vld1q_f32(pA_descales); - const float32x4_t _ad1 = vld1q_f32(pA_descales + 4); + float32x4_t _bd = vcombine_f32(vld1_f32(pB_descales), vdup_n_f32(0.f)); + float32x4_t _ad0 = vld1q_f32(pA_descales); + float32x4_t _ad1 = vld1q_f32(pA_descales + 4); _fsum0 = vmlaq_f32(_fsum0, vcvtq_f32_s32(_sum0), vmulq_laneq_f32(_bd, _ad0, 0)); _fsum1 = vmlaq_f32(_fsum1, vcvtq_f32_s32(_sum1), vmulq_laneq_f32(_bd, _ad0, 1)); _fsum2 = vmlaq_f32(_fsum2, vcvtq_f32_s32(_sum2), vmulq_laneq_f32(_bd, _ad0, 2)); @@ -1834,19 +2061,25 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de vst1_f32(outptr + 12, vget_low_f32(_fsum6)); vst1_f32(outptr + 14, vget_low_f32(_fsum7)); outptr += 16; + pB2 += (size_t)2 * full_K; + pB_descales2 += (size_t)2 * block_count; } + const signed char* pB1 = pB2 - k0; + const float* pB_descales1 = pB_descales2 - block_start; for (; jj < max_jj; jj++) { + const signed char* pB = pB1; + const float* pB_descales = pB_descales1; const signed char* pA = pA_block; const float* pA_descales = pA_descales_block; - float32x4_t _fsum0 = vdupq_n_f32(0.f); - float32x4_t _fsum1 = vdupq_n_f32(0.f); - float32x4_t _fsum2 = vdupq_n_f32(0.f); - float32x4_t _fsum3 = vdupq_n_f32(0.f); - float32x4_t _fsum4 = vdupq_n_f32(0.f); - float32x4_t _fsum5 = vdupq_n_f32(0.f); - float32x4_t _fsum6 = vdupq_n_f32(0.f); - float32x4_t _fsum7 = vdupq_n_f32(0.f); + float32x4_t _fsum0 = k0 == 0 ? vdupq_n_f32(0.f) : vsetq_lane_f32(outptr[0], vdupq_n_f32(0.f), 0); + float32x4_t _fsum1 = k0 == 0 ? vdupq_n_f32(0.f) : vsetq_lane_f32(outptr[1], vdupq_n_f32(0.f), 0); + float32x4_t _fsum2 = k0 == 0 ? vdupq_n_f32(0.f) : vsetq_lane_f32(outptr[2], vdupq_n_f32(0.f), 0); + float32x4_t _fsum3 = k0 == 0 ? vdupq_n_f32(0.f) : vsetq_lane_f32(outptr[3], vdupq_n_f32(0.f), 0); + float32x4_t _fsum4 = k0 == 0 ? vdupq_n_f32(0.f) : vsetq_lane_f32(outptr[4], vdupq_n_f32(0.f), 0); + float32x4_t _fsum5 = k0 == 0 ? vdupq_n_f32(0.f) : vsetq_lane_f32(outptr[5], vdupq_n_f32(0.f), 0); + float32x4_t _fsum6 = k0 == 0 ? vdupq_n_f32(0.f) : vsetq_lane_f32(outptr[6], vdupq_n_f32(0.f), 0); + float32x4_t _fsum7 = k0 == 0 ? vdupq_n_f32(0.f) : vsetq_lane_f32(outptr[7], vdupq_n_f32(0.f), 0); for (int k = 0; k < K; k += block_size) { @@ -1860,15 +2093,16 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de int32x4_t _sum7 = vdupq_n_s32(0); const int max_kk = std::min(K - k, block_size); int kk = 0; +#if __ARM_FEATURE_DOTPROD #if __ARM_FEATURE_MATMUL_INT8 for (; kk + 7 < max_kk; kk += 8) { - const int8x16_t _a01 = vld1q_s8(pA); - const int8x16_t _a23 = vld1q_s8(pA + 16); - const int8x16_t _a45 = vld1q_s8(pA + 32); - const int8x16_t _a67 = vld1q_s8(pA + 48); - const int8x16_t _b0 = vcombine_s8(vreinterpret_s8_s32(vld1_dup_s32((const int*)pB)), vdup_n_s8(0)); - const int8x16_t _b1 = vcombine_s8(vreinterpret_s8_s32(vld1_dup_s32((const int*)(pB + 4))), vdup_n_s8(0)); + int8x16_t _a01 = vld1q_s8(pA); + int8x16_t _a23 = vld1q_s8(pA + 16); + int8x16_t _a45 = vld1q_s8(pA + 32); + int8x16_t _a67 = vld1q_s8(pA + 48); + int8x16_t _b0 = vcombine_s8(vreinterpret_s8_s32(vld1_dup_s32((const int*)pB)), vdup_n_s8(0)); + int8x16_t _b1 = vcombine_s8(vreinterpret_s8_s32(vld1_dup_s32((const int*)(pB + 4))), vdup_n_s8(0)); _sum0 = vdotq_laneq_s32(_sum0, _b0, _a01, 0); _sum0 = vdotq_laneq_s32(_sum0, _b1, _a01, 1); _sum1 = vdotq_laneq_s32(_sum1, _b0, _a01, 2); @@ -1889,12 +2123,11 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB += 8; } #endif // __ARM_FEATURE_MATMUL_INT8 -#if __ARM_FEATURE_DOTPROD for (; kk + 3 < max_kk; kk += 4) { - const int8x16_t _b = vcombine_s8(vreinterpret_s8_s32(vld1_dup_s32((const int*)pB)), vdup_n_s8(0)); - const int8x16_t _a0 = vld1q_s8(pA); - const int8x16_t _a1 = vld1q_s8(pA + 16); + int8x16_t _b = vcombine_s8(vreinterpret_s8_s32(vld1_dup_s32((const int*)pB)), vdup_n_s8(0)); + int8x16_t _a0 = vld1q_s8(pA); + int8x16_t _a1 = vld1q_s8(pA + 16); _sum0 = vdotq_laneq_s32(_sum0, _b, _a0, 0); _sum1 = vdotq_laneq_s32(_sum1, _b, _a0, 1); _sum2 = vdotq_laneq_s32(_sum2, _b, _a0, 2); @@ -1906,13 +2139,112 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pA += 32; pB += 4; } +#else // __ARM_FEATURE_DOTPROD +#if NCNN_GNU_INLINE_ASM + { + int nn = (max_kk - kk) >> 2; + const int remain = nn; + asm volatile( + "cmp %w2, #0 \n" + "beq 1f \n" + "0: \n" + "ld1 {v0.16b, v1.16b}, [%0], #32 \n" + "ld1r {v3.2s}, [%1], #4 \n" + "ext v4.8b, v3.8b, v3.8b, #2 \n" + "dup v5.8h, v0.h[0] \n" + "smull v6.8h, v3.8b, v5.8b \n" + "dup v5.8h, v1.h[0] \n" + "smlal v6.8h, v4.8b, v5.8b \n" + "sadalp %3.4s, v6.8h \n" + "dup v5.8h, v0.h[1] \n" + "smull v6.8h, v3.8b, v5.8b \n" + "dup v5.8h, v1.h[1] \n" + "smlal v6.8h, v4.8b, v5.8b \n" + "sadalp %4.4s, v6.8h \n" + "dup v5.8h, v0.h[2] \n" + "smull v6.8h, v3.8b, v5.8b \n" + "dup v5.8h, v1.h[2] \n" + "smlal v6.8h, v4.8b, v5.8b \n" + "sadalp %5.4s, v6.8h \n" + "dup v5.8h, v0.h[3] \n" + "smull v6.8h, v3.8b, v5.8b \n" + "dup v5.8h, v1.h[3] \n" + "smlal v6.8h, v4.8b, v5.8b \n" + "sadalp %6.4s, v6.8h \n" + "dup v5.8h, v0.h[4] \n" + "smull v6.8h, v3.8b, v5.8b \n" + "dup v5.8h, v1.h[4] \n" + "smlal v6.8h, v4.8b, v5.8b \n" + "sadalp %7.4s, v6.8h \n" + "dup v5.8h, v0.h[5] \n" + "smull v6.8h, v3.8b, v5.8b \n" + "dup v5.8h, v1.h[5] \n" + "smlal v6.8h, v4.8b, v5.8b \n" + "sadalp %8.4s, v6.8h \n" + "dup v5.8h, v0.h[6] \n" + "smull v6.8h, v3.8b, v5.8b \n" + "dup v5.8h, v1.h[6] \n" + "smlal v6.8h, v4.8b, v5.8b \n" + "sadalp %9.4s, v6.8h \n" + "dup v5.8h, v0.h[7] \n" + "smull v6.8h, v3.8b, v5.8b \n" + "dup v5.8h, v1.h[7] \n" + "smlal v6.8h, v4.8b, v5.8b \n" + "subs %w2, %w2, #1 \n" + "sadalp %10.4s, v6.8h \n" + "bne 0b \n" + "1: \n" + : "+r"(pA), "+r"(pB), "+r"(nn), "+w"(_sum0), "+w"(_sum1), "+w"(_sum2), "+w"(_sum3), "+w"(_sum4), "+w"(_sum5), "+w"(_sum6), "+w"(_sum7) + : + : "cc", "memory", "v0", "v1", "v3", "v4", "v5", "v6"); + kk += remain * 4; + } +#else // NCNN_GNU_INLINE_ASM + for (; kk + 3 < max_kk; kk += 4) + { + int16x8_t _a01 = vreinterpretq_s16_s8(vld1q_s8(pA)); + int16x8_t _a23 = vreinterpretq_s16_s8(vld1q_s8(pA + 16)); + int8x8_t _b01 = vreinterpret_s8_s32(vld1_dup_s32((const int*)pB)); + int8x8_t _b23 = vext_s8(_b01, _b01, 2); + int16x4_t _a010 = vget_low_s16(_a01); + int16x4_t _a011 = vget_high_s16(_a01); + int16x4_t _a230 = vget_low_s16(_a23); + int16x4_t _a231 = vget_high_s16(_a23); + int16x8_t _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a010, 0))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a230, 0))); + _sum0 = vpadalq_s16(_sum0, _s); + _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a010, 1))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a230, 1))); + _sum1 = vpadalq_s16(_sum1, _s); + _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a010, 2))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a230, 2))); + _sum2 = vpadalq_s16(_sum2, _s); + _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a010, 3))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a230, 3))); + _sum3 = vpadalq_s16(_sum3, _s); + _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a011, 0))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a231, 0))); + _sum4 = vpadalq_s16(_sum4, _s); + _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a011, 1))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a231, 1))); + _sum5 = vpadalq_s16(_sum5, _s); + _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a011, 2))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a231, 2))); + _sum6 = vpadalq_s16(_sum6, _s); + _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a011, 3))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a231, 3))); + _sum7 = vpadalq_s16(_sum7, _s); + pA += 32; + pB += 4; + } +#endif // NCNN_GNU_INLINE_ASM #endif // __ARM_FEATURE_DOTPROD for (; kk + 1 < max_kk; kk += 2) { - const int8x8_t _b = vreinterpret_s8_s16(vld1_dup_s16((const short*)pB)); - const int16x8_t _a = vreinterpretq_s16_s8(vld1q_s8(pA)); - const int16x4_t _a0 = vget_low_s16(_a); - const int16x4_t _a1 = vget_high_s16(_a); + int8x8_t _b = vreinterpret_s8_s16(vld1_dup_s16((const short*)pB)); + int16x8_t _a = vreinterpretq_s16_s8(vld1q_s8(pA)); + int16x4_t _a0 = vget_low_s16(_a); + int16x4_t _a1 = vget_high_s16(_a); _sum0 = vaddq_s32(_sum0, vpaddlq_s16(vmull_s8(_b, vreinterpret_s8_s16(vdup_lane_s16(_a0, 0))))); _sum1 = vaddq_s32(_sum1, vpaddlq_s16(vmull_s8(_b, vreinterpret_s8_s16(vdup_lane_s16(_a0, 1))))); _sum2 = vaddq_s32(_sum2, vpaddlq_s16(vmull_s8(_b, vreinterpret_s8_s16(vdup_lane_s16(_a0, 2))))); @@ -1926,16 +2258,16 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } if (kk < max_kk) { - const int8x8_t _b = vset_lane_s8(pB[0], vdup_n_s8(0), 0); - const int8x8_t _a = vld1_s8(pA); - const int16x8_t _p0 = vmull_s8(_b, vdup_lane_s8(_a, 0)); - const int16x8_t _p1 = vmull_s8(_b, vdup_lane_s8(_a, 1)); - const int16x8_t _p2 = vmull_s8(_b, vdup_lane_s8(_a, 2)); - const int16x8_t _p3 = vmull_s8(_b, vdup_lane_s8(_a, 3)); - const int16x8_t _p4 = vmull_s8(_b, vdup_lane_s8(_a, 4)); - const int16x8_t _p5 = vmull_s8(_b, vdup_lane_s8(_a, 5)); - const int16x8_t _p6 = vmull_s8(_b, vdup_lane_s8(_a, 6)); - const int16x8_t _p7 = vmull_s8(_b, vdup_lane_s8(_a, 7)); + int8x8_t _b = vset_lane_s8(pB[0], vdup_n_s8(0), 0); + int8x8_t _a = vld1_s8(pA); + int16x8_t _p0 = vmull_s8(_b, vdup_lane_s8(_a, 0)); + int16x8_t _p1 = vmull_s8(_b, vdup_lane_s8(_a, 1)); + int16x8_t _p2 = vmull_s8(_b, vdup_lane_s8(_a, 2)); + int16x8_t _p3 = vmull_s8(_b, vdup_lane_s8(_a, 3)); + int16x8_t _p4 = vmull_s8(_b, vdup_lane_s8(_a, 4)); + int16x8_t _p5 = vmull_s8(_b, vdup_lane_s8(_a, 5)); + int16x8_t _p6 = vmull_s8(_b, vdup_lane_s8(_a, 6)); + int16x8_t _p7 = vmull_s8(_b, vdup_lane_s8(_a, 7)); _sum0 = vaddq_s32(_sum0, vmovl_s16(vget_low_s16(_p0))); _sum1 = vaddq_s32(_sum1, vmovl_s16(vget_low_s16(_p1))); _sum2 = vaddq_s32(_sum2, vmovl_s16(vget_low_s16(_p2))); @@ -1948,9 +2280,9 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB++; } - const float32x4_t _bd = vdupq_n_f32(pB_descales[0]); - const float32x4_t _ad0 = vld1q_f32(pA_descales); - const float32x4_t _ad1 = vld1q_f32(pA_descales + 4); + float32x4_t _bd = vdupq_n_f32(pB_descales[0]); + float32x4_t _ad0 = vld1q_f32(pA_descales); + float32x4_t _ad1 = vld1q_f32(pA_descales + 4); _fsum0 = vmlaq_f32(_fsum0, vcvtq_f32_s32(_sum0), vmulq_laneq_f32(_bd, _ad0, 0)); _fsum1 = vmlaq_f32(_fsum1, vcvtq_f32_s32(_sum1), vmulq_laneq_f32(_bd, _ad0, 1)); _fsum2 = vmlaq_f32(_fsum2, vcvtq_f32_s32(_sum2), vmulq_laneq_f32(_bd, _ad0, 2)); @@ -1979,6 +2311,8 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de outptr++; vst1q_lane_f32(outptr, _fsum7, 0); outptr++; + pB1 += full_K; + pB_descales1 += block_count; } } #endif // __aarch64__ @@ -1986,25 +2320,24 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de { const signed char* pA0_block = pAT + ii * A_hstep; const float* pA_descales0_block = pAT_descales + ii * A_descales_hstep; - const signed char* pB = pBT; - const float* pB_descales = pBT_descales; - int jj = 0; + const signed char* pB8 = pBT + 4 * k0; + const float* pB_descales8 = pBT_descales + 4 * block_start; #if __aarch64__ for (; jj + 7 < max_jj; jj += 8) { - const signed char* pB0 = pB; - const signed char* pB1 = pB + 4 * K; - const float* pB_descales0 = pB_descales; - const float* pB_descales1 = pB_descales + 4 * ((K + block_size - 1) / block_size); - float32x4_t _fsum0 = vdupq_n_f32(0.f); - float32x4_t _fsum1 = vdupq_n_f32(0.f); - float32x4_t _fsum2 = vdupq_n_f32(0.f); - float32x4_t _fsum3 = vdupq_n_f32(0.f); - float32x4_t _fsum4 = vdupq_n_f32(0.f); - float32x4_t _fsum5 = vdupq_n_f32(0.f); - float32x4_t _fsum6 = vdupq_n_f32(0.f); - float32x4_t _fsum7 = vdupq_n_f32(0.f); + const signed char* pB0 = pB8; + const signed char* pB1 = pB8 + (size_t)4 * full_K; + const float* pB_descales0 = pB_descales8; + const float* pB_descales1 = pB_descales8 + (size_t)4 * block_count; + float32x4_t _fsum0 = k0 == 0 ? vdupq_n_f32(0.f) : vld1q_f32(outptr); + float32x4_t _fsum1 = k0 == 0 ? vdupq_n_f32(0.f) : vld1q_f32(outptr + 4); + float32x4_t _fsum2 = k0 == 0 ? vdupq_n_f32(0.f) : vld1q_f32(outptr + 8); + float32x4_t _fsum3 = k0 == 0 ? vdupq_n_f32(0.f) : vld1q_f32(outptr + 12); + float32x4_t _fsum4 = k0 == 0 ? vdupq_n_f32(0.f) : vld1q_f32(outptr + 16); + float32x4_t _fsum5 = k0 == 0 ? vdupq_n_f32(0.f) : vld1q_f32(outptr + 20); + float32x4_t _fsum6 = k0 == 0 ? vdupq_n_f32(0.f) : vld1q_f32(outptr + 24); + float32x4_t _fsum7 = k0 == 0 ? vdupq_n_f32(0.f) : vld1q_f32(outptr + 28); const signed char* pA = pA0_block; const float* pA_descales = pA_descales0_block; @@ -2020,6 +2353,7 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de int32x4_t _sum7 = vdupq_n_s32(0); const int max_kk = std::min(K - k, block_size); int kk = 0; +#if __ARM_FEATURE_DOTPROD #if __ARM_FEATURE_MATMUL_INT8 int32x4_t _msum0 = vdupq_n_s32(0); int32x4_t _msum1 = vdupq_n_s32(0); @@ -2031,12 +2365,12 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de int32x4_t _msum7 = vdupq_n_s32(0); for (; kk + 7 < max_kk; kk += 8) { - const int8x16_t _a01 = vld1q_s8(pA); - const int8x16_t _a23 = vld1q_s8(pA + 16); - const int8x16_t _b00 = vld1q_s8(pB0); - const int8x16_t _b01 = vld1q_s8(pB0 + 16); - const int8x16_t _b10 = vld1q_s8(pB1); - const int8x16_t _b11 = vld1q_s8(pB1 + 16); + int8x16_t _a01 = vld1q_s8(pA); + int8x16_t _a23 = vld1q_s8(pA + 16); + int8x16_t _b00 = vld1q_s8(pB0); + int8x16_t _b01 = vld1q_s8(pB0 + 16); + int8x16_t _b10 = vld1q_s8(pB1); + int8x16_t _b11 = vld1q_s8(pB1 + 16); _msum0 = vmmlaq_s32(_msum0, _a01, _b00); _msum1 = vmmlaq_s32(_msum1, _a01, _b01); _msum2 = vmmlaq_s32(_msum2, _a23, _b00); @@ -2058,12 +2392,11 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de _sum6 = vcombine_s32(vget_high_s32(_msum2), vget_high_s32(_msum3)); _sum7 = vcombine_s32(vget_high_s32(_msum6), vget_high_s32(_msum7)); #endif // __ARM_FEATURE_MATMUL_INT8 -#if __ARM_FEATURE_DOTPROD for (; kk + 3 < max_kk; kk += 4) { - const int8x16_t _a = vld1q_s8(pA); - const int8x16_t _b0 = vld1q_s8(pB0); - const int8x16_t _b1 = vld1q_s8(pB1); + int8x16_t _a = vld1q_s8(pA); + int8x16_t _b0 = vld1q_s8(pB0); + int8x16_t _b1 = vld1q_s8(pB1); _sum0 = vdotq_laneq_s32(_sum0, _b0, _a, 0); _sum1 = vdotq_laneq_s32(_sum1, _b1, _a, 0); _sum2 = vdotq_laneq_s32(_sum2, _b0, _a, 1); @@ -2076,12 +2409,57 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB0 += 16; pB1 += 16; } +#else // __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk; kk += 4) + { + int16x8_t _a = vreinterpretq_s16_s8(vld1q_s8(pA)); + int16x4_t _a01 = vget_low_s16(_a); + int16x4_t _a23 = vget_high_s16(_a); +#if NCNN_GNU_INLINE_ASM && __aarch64__ + asm volatile("" + : + : "w"(_a)); +#endif + int8x16_t _b0 = vld1q_s8(pB0); + int8x16_t _b1 = vld1q_s8(pB1); + int8x8_t _b001 = vget_low_s8(_b0); + int8x8_t _b023 = vget_high_s8(_b0); + int8x8_t _b101 = vget_low_s8(_b1); + int8x8_t _b123 = vget_high_s8(_b1); + int16x8_t _s = vmull_s8(_b001, vreinterpret_s8_s16(vdup_lane_s16(_a01, 0))); + _s = vmlal_s8(_s, _b023, vreinterpret_s8_s16(vdup_lane_s16(_a23, 0))); + _sum0 = vpadalq_s16(_sum0, _s); + _s = vmull_s8(_b101, vreinterpret_s8_s16(vdup_lane_s16(_a01, 0))); + _s = vmlal_s8(_s, _b123, vreinterpret_s8_s16(vdup_lane_s16(_a23, 0))); + _sum1 = vpadalq_s16(_sum1, _s); + _s = vmull_s8(_b001, vreinterpret_s8_s16(vdup_lane_s16(_a01, 1))); + _s = vmlal_s8(_s, _b023, vreinterpret_s8_s16(vdup_lane_s16(_a23, 1))); + _sum2 = vpadalq_s16(_sum2, _s); + _s = vmull_s8(_b101, vreinterpret_s8_s16(vdup_lane_s16(_a01, 1))); + _s = vmlal_s8(_s, _b123, vreinterpret_s8_s16(vdup_lane_s16(_a23, 1))); + _sum3 = vpadalq_s16(_sum3, _s); + _s = vmull_s8(_b001, vreinterpret_s8_s16(vdup_lane_s16(_a01, 2))); + _s = vmlal_s8(_s, _b023, vreinterpret_s8_s16(vdup_lane_s16(_a23, 2))); + _sum4 = vpadalq_s16(_sum4, _s); + _s = vmull_s8(_b101, vreinterpret_s8_s16(vdup_lane_s16(_a01, 2))); + _s = vmlal_s8(_s, _b123, vreinterpret_s8_s16(vdup_lane_s16(_a23, 2))); + _sum5 = vpadalq_s16(_sum5, _s); + _s = vmull_s8(_b001, vreinterpret_s8_s16(vdup_lane_s16(_a01, 3))); + _s = vmlal_s8(_s, _b023, vreinterpret_s8_s16(vdup_lane_s16(_a23, 3))); + _sum6 = vpadalq_s16(_sum6, _s); + _s = vmull_s8(_b101, vreinterpret_s8_s16(vdup_lane_s16(_a01, 3))); + _s = vmlal_s8(_s, _b123, vreinterpret_s8_s16(vdup_lane_s16(_a23, 3))); + _sum7 = vpadalq_s16(_sum7, _s); + pA += 16; + pB0 += 16; + pB1 += 16; + } #endif // __ARM_FEATURE_DOTPROD for (; kk + 1 < max_kk; kk += 2) { - const int16x4_t _a = vreinterpret_s16_s8(vld1_s8(pA)); - const int8x8_t _b0 = vld1_s8(pB0); - const int8x8_t _b1 = vld1_s8(pB1); + int16x4_t _a = vreinterpret_s16_s8(vld1_s8(pA)); + int8x8_t _b0 = vld1_s8(pB0); + int8x8_t _b1 = vld1_s8(pB1); _sum0 = vaddq_s32(_sum0, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 0))))); _sum1 = vaddq_s32(_sum1, vpaddlq_s16(vmull_s8(_b1, vreinterpret_s8_s16(vdup_lane_s16(_a, 0))))); _sum2 = vaddq_s32(_sum2, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 1))))); @@ -2096,17 +2474,17 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } if (kk < max_kk) { - const int8x8_t _a = vreinterpret_s8_s32(vld1_lane_s32((const int*)pA, vdup_n_s32(0), 0)); - const int8x8_t _b0 = vreinterpret_s8_s32(vld1_dup_s32((const int*)pB0)); - const int8x8_t _b1 = vreinterpret_s8_s32(vld1_dup_s32((const int*)pB1)); - const int16x8_t _p00 = vmull_s8(_b0, vdup_lane_s8(_a, 0)); - const int16x8_t _p01 = vmull_s8(_b1, vdup_lane_s8(_a, 0)); - const int16x8_t _p10 = vmull_s8(_b0, vdup_lane_s8(_a, 1)); - const int16x8_t _p11 = vmull_s8(_b1, vdup_lane_s8(_a, 1)); - const int16x8_t _p20 = vmull_s8(_b0, vdup_lane_s8(_a, 2)); - const int16x8_t _p21 = vmull_s8(_b1, vdup_lane_s8(_a, 2)); - const int16x8_t _p30 = vmull_s8(_b0, vdup_lane_s8(_a, 3)); - const int16x8_t _p31 = vmull_s8(_b1, vdup_lane_s8(_a, 3)); + int8x8_t _a = vreinterpret_s8_s32(vld1_lane_s32((const int*)pA, vdup_n_s32(0), 0)); + int8x8_t _b0 = vreinterpret_s8_s32(vld1_dup_s32((const int*)pB0)); + int8x8_t _b1 = vreinterpret_s8_s32(vld1_dup_s32((const int*)pB1)); + int16x8_t _p00 = vmull_s8(_b0, vdup_lane_s8(_a, 0)); + int16x8_t _p01 = vmull_s8(_b1, vdup_lane_s8(_a, 0)); + int16x8_t _p10 = vmull_s8(_b0, vdup_lane_s8(_a, 1)); + int16x8_t _p11 = vmull_s8(_b1, vdup_lane_s8(_a, 1)); + int16x8_t _p20 = vmull_s8(_b0, vdup_lane_s8(_a, 2)); + int16x8_t _p21 = vmull_s8(_b1, vdup_lane_s8(_a, 2)); + int16x8_t _p30 = vmull_s8(_b0, vdup_lane_s8(_a, 3)); + int16x8_t _p31 = vmull_s8(_b1, vdup_lane_s8(_a, 3)); _sum0 = vaddq_s32(_sum0, vmovl_s16(vget_low_s16(_p00))); _sum1 = vaddq_s32(_sum1, vmovl_s16(vget_low_s16(_p01))); _sum2 = vaddq_s32(_sum2, vmovl_s16(vget_low_s16(_p10))); @@ -2120,9 +2498,9 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB1 += 4; } - const float32x4_t _bd0 = vld1q_f32(pB_descales0); - const float32x4_t _bd1 = vld1q_f32(pB_descales1); - const float32x4_t _ad = vld1q_f32(pA_descales); + float32x4_t _bd0 = vld1q_f32(pB_descales0); + float32x4_t _bd1 = vld1q_f32(pB_descales1); + float32x4_t _ad = vld1q_f32(pA_descales); _fsum0 = vmlaq_f32(_fsum0, vcvtq_f32_s32(_sum0), vmulq_laneq_f32(_bd0, _ad, 0)); _fsum1 = vmlaq_f32(_fsum1, vcvtq_f32_s32(_sum1), vmulq_laneq_f32(_bd1, _ad, 0)); _fsum2 = vmlaq_f32(_fsum2, vcvtq_f32_s32(_sum2), vmulq_laneq_f32(_bd0, _ad, 1)); @@ -2137,9 +2515,6 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB_descales1 += 4; } - pB = pB1; - pB_descales = pB_descales1; - vst1q_f32(outptr, _fsum0); outptr += 4; vst1q_f32(outptr, _fsum1); @@ -2156,14 +2531,20 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de outptr += 4; vst1q_f32(outptr, _fsum7); outptr += 4; + pB8 += (size_t)8 * full_K; + pB_descales8 += (size_t)8 * block_count; } #endif // __aarch64__ + const signed char* pB4 = pB8; + const float* pB_descales4 = pB_descales8; for (; jj + 3 < max_jj; jj += 4) { - float32x4_t _fsum0 = vdupq_n_f32(0.f); - float32x4_t _fsum1 = vdupq_n_f32(0.f); - float32x4_t _fsum2 = vdupq_n_f32(0.f); - float32x4_t _fsum3 = vdupq_n_f32(0.f); + const signed char* pB = pB4; + const float* pB_descales = pB_descales4; + float32x4_t _fsum0 = k0 == 0 ? vdupq_n_f32(0.f) : vld1q_f32(outptr); + float32x4_t _fsum1 = k0 == 0 ? vdupq_n_f32(0.f) : vld1q_f32(outptr + 4); + float32x4_t _fsum2 = k0 == 0 ? vdupq_n_f32(0.f) : vld1q_f32(outptr + 8); + float32x4_t _fsum3 = k0 == 0 ? vdupq_n_f32(0.f) : vld1q_f32(outptr + 12); const signed char* pA = pA0_block; const float* pA_descales = pA_descales0_block; @@ -2175,6 +2556,7 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de int32x4_t _sum3 = vdupq_n_s32(0); const int max_kk = std::min(K - k, block_size); int kk = 0; +#if __ARM_FEATURE_DOTPROD #if __ARM_FEATURE_MATMUL_INT8 int32x4_t _msum0 = vdupq_n_s32(0); int32x4_t _msum1 = vdupq_n_s32(0); @@ -2182,10 +2564,10 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de int32x4_t _msum3 = vdupq_n_s32(0); for (; kk + 7 < max_kk; kk += 8) { - const int8x16_t _b0 = vld1q_s8(pB); - const int8x16_t _b1 = vld1q_s8(pB + 16); - const int8x16_t _a01 = vld1q_s8(pA); - const int8x16_t _a23 = vld1q_s8(pA + 16); + int8x16_t _b0 = vld1q_s8(pB); + int8x16_t _b1 = vld1q_s8(pB + 16); + int8x16_t _a01 = vld1q_s8(pA); + int8x16_t _a23 = vld1q_s8(pA + 16); _msum0 = vmmlaq_s32(_msum0, _a01, _b0); _msum1 = vmmlaq_s32(_msum1, _a23, _b0); _msum2 = vmmlaq_s32(_msum2, _a01, _b1); @@ -2198,11 +2580,10 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de _sum2 = vcombine_s32(vget_low_s32(_msum1), vget_low_s32(_msum3)); _sum3 = vcombine_s32(vget_high_s32(_msum1), vget_high_s32(_msum3)); #endif // __ARM_FEATURE_MATMUL_INT8 -#if __ARM_FEATURE_DOTPROD for (; kk + 3 < max_kk; kk += 4) { - const int8x16_t _b0 = vld1q_s8(pB); - const int8x16_t _a = vld1q_s8(pA); + int8x16_t _b0 = vld1q_s8(pB); + int8x16_t _a = vld1q_s8(pA); _sum0 = vdotq_laneq_s32(_sum0, _b0, _a, 0); _sum1 = vdotq_laneq_s32(_sum1, _b0, _a, 1); _sum2 = vdotq_laneq_s32(_sum2, _b0, _a, 2); @@ -2210,11 +2591,80 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pA += 16; pB += 16; } +#else // __ARM_FEATURE_DOTPROD +#if NCNN_GNU_INLINE_ASM && !__aarch64__ + { + int nn = (max_kk - kk) >> 2; + const int remain = nn; + asm volatile( + "cmp %2, #0 \n" + "beq 1f \n" + "0: \n" + "vld1.s8 {d0-d1}, [%0]! \n" + "vld1.s8 {d2-d3}, [%1]! \n" + "vdup.16 q2, d0[0] \n" + "vmull.s8 q3, d2, d4 \n" + "vdup.16 q2, d1[0] \n" + "vmlal.s8 q3, d3, d4 \n" + "vpadal.s16 %q3, q3 \n" + "vdup.16 q2, d0[1] \n" + "vmull.s8 q3, d2, d4 \n" + "vdup.16 q2, d1[1] \n" + "vmlal.s8 q3, d3, d4 \n" + "vpadal.s16 %q4, q3 \n" + "vdup.16 q2, d0[2] \n" + "vmull.s8 q3, d2, d4 \n" + "vdup.16 q2, d1[2] \n" + "vmlal.s8 q3, d3, d4 \n" + "vpadal.s16 %q5, q3 \n" + "vdup.16 q2, d0[3] \n" + "vmull.s8 q3, d2, d4 \n" + "vdup.16 q2, d1[3] \n" + "vmlal.s8 q3, d3, d4 \n" + "subs %2, %2, #1 \n" + "vpadal.s16 %q6, q3 \n" + "bne 0b \n" + "1: \n" + : "+r"(pA), "+r"(pB), "+r"(nn), "+w"(_sum0), "+w"(_sum1), "+w"(_sum2), "+w"(_sum3) + : + : "cc", "memory", "q0", "q1", "q2", "q3"); + kk += remain * 4; + } +#else // NCNN_GNU_INLINE_ASM && !__aarch64__ + for (; kk + 3 < max_kk; kk += 4) + { + int16x8_t _a = vreinterpretq_s16_s8(vld1q_s8(pA)); + int16x4_t _a01 = vget_low_s16(_a); + int16x4_t _a23 = vget_high_s16(_a); +#if NCNN_GNU_INLINE_ASM && __aarch64__ + asm volatile("" + : + : "w"(_a)); +#endif + int8x16_t _b = vld1q_s8(pB); + int8x8_t _b01 = vget_low_s8(_b); + int8x8_t _b23 = vget_high_s8(_b); + int16x8_t _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a01, 0))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a23, 0))); + _sum0 = vpadalq_s16(_sum0, _s); + _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a01, 1))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a23, 1))); + _sum1 = vpadalq_s16(_sum1, _s); + _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a01, 2))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a23, 2))); + _sum2 = vpadalq_s16(_sum2, _s); + _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a01, 3))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a23, 3))); + _sum3 = vpadalq_s16(_sum3, _s); + pA += 16; + pB += 16; + } +#endif // NCNN_GNU_INLINE_ASM && !__aarch64__ #endif // __ARM_FEATURE_DOTPROD for (; kk + 1 < max_kk; kk += 2) { - const int8x8_t _b0 = vld1_s8(pB); - const int16x4_t _a = vreinterpret_s16_s8(vld1_s8(pA)); + int8x8_t _b0 = vld1_s8(pB); + int16x4_t _a = vreinterpret_s16_s8(vld1_s8(pA)); _sum0 = vaddq_s32(_sum0, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 0))))); _sum1 = vaddq_s32(_sum1, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 1))))); _sum2 = vaddq_s32(_sum2, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 2))))); @@ -2224,12 +2674,12 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } if (kk < max_kk) { - const int8x8_t _b = vreinterpret_s8_s32(vld1_dup_s32((const int*)(pB))); - const int8x8_t _a = vreinterpret_s8_s32(vld1_lane_s32((const int*)pA, vdup_n_s32(0), 0)); - const int16x8_t _p0 = vmull_s8(_b, vdup_lane_s8(_a, 0)); - const int16x8_t _p1 = vmull_s8(_b, vdup_lane_s8(_a, 1)); - const int16x8_t _p2 = vmull_s8(_b, vdup_lane_s8(_a, 2)); - const int16x8_t _p3 = vmull_s8(_b, vdup_lane_s8(_a, 3)); + int8x8_t _b = vreinterpret_s8_s32(vld1_dup_s32((const int*)(pB))); + int8x8_t _a = vreinterpret_s8_s32(vld1_lane_s32((const int*)pA, vdup_n_s32(0), 0)); + int16x8_t _p0 = vmull_s8(_b, vdup_lane_s8(_a, 0)); + int16x8_t _p1 = vmull_s8(_b, vdup_lane_s8(_a, 1)); + int16x8_t _p2 = vmull_s8(_b, vdup_lane_s8(_a, 2)); + int16x8_t _p3 = vmull_s8(_b, vdup_lane_s8(_a, 3)); _sum0 = vaddq_s32(_sum0, vmovl_s16(vget_low_s16(_p0))); _sum1 = vaddq_s32(_sum1, vmovl_s16(vget_low_s16(_p1))); _sum2 = vaddq_s32(_sum2, vmovl_s16(vget_low_s16(_p2))); @@ -2238,8 +2688,8 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB += 4; } - const float32x4_t _bd0 = vld1q_f32(pB_descales); - const float32x4_t _ad = vld1q_f32(pA_descales); + float32x4_t _bd0 = vld1q_f32(pB_descales); + float32x4_t _ad = vld1q_f32(pA_descales); _fsum0 = vmlaq_f32(_fsum0, vcvtq_f32_s32(_sum0), vmulq_n_f32(_bd0, vgetq_lane_f32(_ad, 0))); _fsum1 = vmlaq_f32(_fsum1, vcvtq_f32_s32(_sum1), vmulq_n_f32(_bd0, vgetq_lane_f32(_ad, 1))); _fsum2 = vmlaq_f32(_fsum2, vcvtq_f32_s32(_sum2), vmulq_n_f32(_bd0, vgetq_lane_f32(_ad, 2))); @@ -2257,13 +2707,19 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de outptr += 4; vst1q_f32(outptr, _fsum3); outptr += 4; + pB4 += (size_t)4 * full_K; + pB_descales4 += (size_t)4 * block_count; } + const signed char* pB2 = pB4 - 2 * k0; + const float* pB_descales2 = pB_descales4 - 2 * block_start; for (; jj + 1 < max_jj; jj += 2) { - float32x4_t _fsum0 = vdupq_n_f32(0.f); - float32x4_t _fsum1 = vdupq_n_f32(0.f); - float32x4_t _fsum2 = vdupq_n_f32(0.f); - float32x4_t _fsum3 = vdupq_n_f32(0.f); + const signed char* pB = pB2; + const float* pB_descales = pB_descales2; + float32x4_t _fsum0 = k0 == 0 ? vdupq_n_f32(0.f) : vcombine_f32(vld1_f32(outptr), vdup_n_f32(0.f)); + float32x4_t _fsum1 = k0 == 0 ? vdupq_n_f32(0.f) : vcombine_f32(vld1_f32(outptr + 2), vdup_n_f32(0.f)); + float32x4_t _fsum2 = k0 == 0 ? vdupq_n_f32(0.f) : vcombine_f32(vld1_f32(outptr + 4), vdup_n_f32(0.f)); + float32x4_t _fsum3 = k0 == 0 ? vdupq_n_f32(0.f) : vcombine_f32(vld1_f32(outptr + 6), vdup_n_f32(0.f)); const signed char* pA = pA0_block; const float* pA_descales = pA_descales0_block; @@ -2275,14 +2731,15 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de int32x4_t _sum3 = vdupq_n_s32(0); const int max_kk = std::min(K - k, block_size); int kk = 0; +#if __ARM_FEATURE_DOTPROD #if __ARM_FEATURE_MATMUL_INT8 int32x4_t _msum0 = vdupq_n_s32(0); int32x4_t _msum1 = vdupq_n_s32(0); for (; kk + 7 < max_kk; kk += 8) { - const int8x16_t _b0 = vld1q_s8(pB); - const int8x16_t _a01 = vld1q_s8(pA); - const int8x16_t _a23 = vld1q_s8(pA + 16); + int8x16_t _b0 = vld1q_s8(pB); + int8x16_t _a01 = vld1q_s8(pA); + int8x16_t _a23 = vld1q_s8(pA + 16); _msum0 = vmmlaq_s32(_msum0, _a01, _b0); _msum1 = vmmlaq_s32(_msum1, _a23, _b0); pA += 32; @@ -2293,11 +2750,10 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de _sum2 = vcombine_s32(vget_low_s32(_msum1), vdup_n_s32(0)); _sum3 = vcombine_s32(vget_high_s32(_msum1), vdup_n_s32(0)); #endif // __ARM_FEATURE_MATMUL_INT8 -#if __ARM_FEATURE_DOTPROD for (; kk + 3 < max_kk; kk += 4) { - const int8x16_t _b0 = vcombine_s8(vld1_s8(pB), vdup_n_s8(0)); - const int8x16_t _a = vld1q_s8(pA); + int8x16_t _b0 = vcombine_s8(vld1_s8(pB), vdup_n_s8(0)); + int8x16_t _a = vld1q_s8(pA); _sum0 = vdotq_laneq_s32(_sum0, _b0, _a, 0); _sum1 = vdotq_laneq_s32(_sum1, _b0, _a, 1); _sum2 = vdotq_laneq_s32(_sum2, _b0, _a, 2); @@ -2305,11 +2761,83 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pA += 16; pB += 8; } +#else // __ARM_FEATURE_DOTPROD +#if NCNN_GNU_INLINE_ASM && !__aarch64__ + { + int nn = (max_kk - kk) >> 2; + const int remain = nn; + asm volatile( + "cmp %2, #0 \n" + "beq 1f \n" + "0: \n" + "vld1.s8 {d0-d1}, [%0]! \n" + "vld1.s8 {d2}, [%1]! \n" + "vdup.32 d3, d2[0] \n" + "vdup.32 d4, d2[1] \n" + "vdup.16 q3, d0[0] \n" + "vmull.s8 q4, d3, d6 \n" + "vdup.16 q3, d1[0] \n" + "vmlal.s8 q4, d4, d6 \n" + "vpadal.s16 %q3, q4 \n" + "vdup.16 q3, d0[1] \n" + "vmull.s8 q4, d3, d6 \n" + "vdup.16 q3, d1[1] \n" + "vmlal.s8 q4, d4, d6 \n" + "vpadal.s16 %q4, q4 \n" + "vdup.16 q3, d0[2] \n" + "vmull.s8 q4, d3, d6 \n" + "vdup.16 q3, d1[2] \n" + "vmlal.s8 q4, d4, d6 \n" + "vpadal.s16 %q5, q4 \n" + "vdup.16 q3, d0[3] \n" + "vmull.s8 q4, d3, d6 \n" + "vdup.16 q3, d1[3] \n" + "vmlal.s8 q4, d4, d6 \n" + "subs %2, %2, #1 \n" + "vpadal.s16 %q6, q4 \n" + "bne 0b \n" + "1: \n" + : "+r"(pA), "+r"(pB), "+r"(nn), "+w"(_sum0), "+w"(_sum1), "+w"(_sum2), "+w"(_sum3) + : + : "cc", "memory", "q0", "q1", "q2", "q3", "q4"); + kk += remain * 4; + } +#else // NCNN_GNU_INLINE_ASM && !__aarch64__ + for (; kk + 3 < max_kk; kk += 4) + { + int16x8_t _a = vreinterpretq_s16_s8(vld1q_s8(pA)); + int16x4_t _a01 = vget_low_s16(_a); + int16x4_t _a23 = vget_high_s16(_a); +#if NCNN_GNU_INLINE_ASM && __aarch64__ + asm volatile("" + : + : "w"(_a)); +#endif + int8x8_t _b = vld1_s8(pB); + int32x2_t _b02 = vreinterpret_s32_s8(_b); + int8x8_t _b01 = vreinterpret_s8_s32(vdup_lane_s32(_b02, 0)); + int8x8_t _b23 = vreinterpret_s8_s32(vdup_lane_s32(_b02, 1)); + int16x8_t _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a01, 0))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a23, 0))); + _sum0 = vpadalq_s16(_sum0, _s); + _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a01, 1))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a23, 1))); + _sum1 = vpadalq_s16(_sum1, _s); + _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a01, 2))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a23, 2))); + _sum2 = vpadalq_s16(_sum2, _s); + _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a01, 3))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a23, 3))); + _sum3 = vpadalq_s16(_sum3, _s); + pA += 16; + pB += 8; + } +#endif // NCNN_GNU_INLINE_ASM && !__aarch64__ #endif // __ARM_FEATURE_DOTPROD for (; kk + 1 < max_kk; kk += 2) { - const int8x8_t _b0 = vreinterpret_s8_s32(vld1_dup_s32((const int*)(pB))); - const int16x4_t _a = vreinterpret_s16_s8(vld1_s8(pA)); + int8x8_t _b0 = vreinterpret_s8_s32(vld1_dup_s32((const int*)(pB))); + int16x4_t _a = vreinterpret_s16_s8(vld1_s8(pA)); _sum0 = vaddq_s32(_sum0, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 0))))); _sum1 = vaddq_s32(_sum1, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 1))))); _sum2 = vaddq_s32(_sum2, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 2))))); @@ -2319,12 +2847,12 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } if (kk < max_kk) { - const int8x8_t _b = vreinterpret_s8_s16(vld1_dup_s16((const short*)(pB))); - const int8x8_t _a = vreinterpret_s8_s32(vld1_lane_s32((const int*)pA, vdup_n_s32(0), 0)); - const int16x8_t _p0 = vmull_s8(_b, vdup_lane_s8(_a, 0)); - const int16x8_t _p1 = vmull_s8(_b, vdup_lane_s8(_a, 1)); - const int16x8_t _p2 = vmull_s8(_b, vdup_lane_s8(_a, 2)); - const int16x8_t _p3 = vmull_s8(_b, vdup_lane_s8(_a, 3)); + int8x8_t _b = vreinterpret_s8_s16(vld1_dup_s16((const short*)(pB))); + int8x8_t _a = vreinterpret_s8_s32(vld1_lane_s32((const int*)pA, vdup_n_s32(0), 0)); + int16x8_t _p0 = vmull_s8(_b, vdup_lane_s8(_a, 0)); + int16x8_t _p1 = vmull_s8(_b, vdup_lane_s8(_a, 1)); + int16x8_t _p2 = vmull_s8(_b, vdup_lane_s8(_a, 2)); + int16x8_t _p3 = vmull_s8(_b, vdup_lane_s8(_a, 3)); _sum0 = vaddq_s32(_sum0, vmovl_s16(vget_low_s16(_p0))); _sum1 = vaddq_s32(_sum1, vmovl_s16(vget_low_s16(_p1))); _sum2 = vaddq_s32(_sum2, vmovl_s16(vget_low_s16(_p2))); @@ -2333,8 +2861,8 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB += 2; } - const float32x4_t _bd0 = vcombine_f32(vld1_f32(pB_descales), vdup_n_f32(0.f)); - const float32x4_t _ad = vld1q_f32(pA_descales); + float32x4_t _bd0 = vcombine_f32(vld1_f32(pB_descales), vdup_n_f32(0.f)); + float32x4_t _ad = vld1q_f32(pA_descales); _fsum0 = vmlaq_f32(_fsum0, vcvtq_f32_s32(_sum0), vmulq_n_f32(_bd0, vgetq_lane_f32(_ad, 0))); _fsum1 = vmlaq_f32(_fsum1, vcvtq_f32_s32(_sum1), vmulq_n_f32(_bd0, vgetq_lane_f32(_ad, 1))); _fsum2 = vmlaq_f32(_fsum2, vcvtq_f32_s32(_sum2), vmulq_n_f32(_bd0, vgetq_lane_f32(_ad, 2))); @@ -2352,13 +2880,19 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de outptr += 2; vst1_f32(outptr, vget_low_f32(_fsum3)); outptr += 2; + pB2 += (size_t)2 * full_K; + pB_descales2 += (size_t)2 * block_count; } + const signed char* pB1 = pB2 - k0; + const float* pB_descales1 = pB_descales2 - block_start; for (; jj < max_jj; jj++) { - float32x4_t _fsum0 = vdupq_n_f32(0.f); - float32x4_t _fsum1 = vdupq_n_f32(0.f); - float32x4_t _fsum2 = vdupq_n_f32(0.f); - float32x4_t _fsum3 = vdupq_n_f32(0.f); + const signed char* pB = pB1; + const float* pB_descales = pB_descales1; + float32x4_t _fsum0 = k0 == 0 ? vdupq_n_f32(0.f) : vsetq_lane_f32(outptr[0], vdupq_n_f32(0.f), 0); + float32x4_t _fsum1 = k0 == 0 ? vdupq_n_f32(0.f) : vsetq_lane_f32(outptr[1], vdupq_n_f32(0.f), 0); + float32x4_t _fsum2 = k0 == 0 ? vdupq_n_f32(0.f) : vsetq_lane_f32(outptr[2], vdupq_n_f32(0.f), 0); + float32x4_t _fsum3 = k0 == 0 ? vdupq_n_f32(0.f) : vsetq_lane_f32(outptr[3], vdupq_n_f32(0.f), 0); const signed char* pA = pA0_block; const float* pA_descales = pA_descales0_block; @@ -2370,13 +2904,14 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de int32x4_t _sum3 = vdupq_n_s32(0); const int max_kk = std::min(K - k, block_size); int kk = 0; +#if __ARM_FEATURE_DOTPROD #if __ARM_FEATURE_MATMUL_INT8 for (; kk + 7 < max_kk; kk += 8) { - const int8x16_t _a01 = vld1q_s8(pA); - const int8x16_t _a23 = vld1q_s8(pA + 16); - const int8x16_t _b0 = vcombine_s8(vreinterpret_s8_s32(vld1_dup_s32((const int*)pB)), vdup_n_s8(0)); - const int8x16_t _b1 = vcombine_s8(vreinterpret_s8_s32(vld1_dup_s32((const int*)(pB + 4))), vdup_n_s8(0)); + int8x16_t _a01 = vld1q_s8(pA); + int8x16_t _a23 = vld1q_s8(pA + 16); + int8x16_t _b0 = vcombine_s8(vreinterpret_s8_s32(vld1_dup_s32((const int*)pB)), vdup_n_s8(0)); + int8x16_t _b1 = vcombine_s8(vreinterpret_s8_s32(vld1_dup_s32((const int*)(pB + 4))), vdup_n_s8(0)); _sum0 = vdotq_laneq_s32(_sum0, _b0, _a01, 0); _sum0 = vdotq_laneq_s32(_sum0, _b1, _a01, 1); _sum1 = vdotq_laneq_s32(_sum1, _b0, _a01, 2); @@ -2389,11 +2924,10 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB += 8; } #endif // __ARM_FEATURE_MATMUL_INT8 -#if __ARM_FEATURE_DOTPROD for (; kk + 3 < max_kk; kk += 4) { - const int8x16_t _b0 = vcombine_s8(vreinterpret_s8_s32(vld1_dup_s32((const int*)pB)), vdup_n_s8(0)); - const int8x16_t _a = vld1q_s8(pA); + int8x16_t _b0 = vcombine_s8(vreinterpret_s8_s32(vld1_dup_s32((const int*)pB)), vdup_n_s8(0)); + int8x16_t _a = vld1q_s8(pA); _sum0 = vdotq_laneq_s32(_sum0, _b0, _a, 0); _sum1 = vdotq_laneq_s32(_sum1, _b0, _a, 1); _sum2 = vdotq_laneq_s32(_sum2, _b0, _a, 2); @@ -2401,11 +2935,80 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pA += 16; pB += 4; } +#else // __ARM_FEATURE_DOTPROD +#if NCNN_GNU_INLINE_ASM && !__aarch64__ + { + int nn = (max_kk - kk) >> 2; + const int remain = nn; + asm volatile( + "cmp %2, #0 \n" + "beq 1f \n" + "0: \n" + "vld1.s8 {d0-d1}, [%0]! \n" + "vld1.32 {d2[]}, [%1]! \n" + "vext.8 d3, d2, d2, #2 \n" + "vdup.16 q2, d0[0] \n" + "vmull.s8 q3, d2, d4 \n" + "vdup.16 q2, d1[0] \n" + "vmlal.s8 q3, d3, d4 \n" + "vpadal.s16 %q3, q3 \n" + "vdup.16 q2, d0[1] \n" + "vmull.s8 q3, d2, d4 \n" + "vdup.16 q2, d1[1] \n" + "vmlal.s8 q3, d3, d4 \n" + "vpadal.s16 %q4, q3 \n" + "vdup.16 q2, d0[2] \n" + "vmull.s8 q3, d2, d4 \n" + "vdup.16 q2, d1[2] \n" + "vmlal.s8 q3, d3, d4 \n" + "vpadal.s16 %q5, q3 \n" + "vdup.16 q2, d0[3] \n" + "vmull.s8 q3, d2, d4 \n" + "vdup.16 q2, d1[3] \n" + "vmlal.s8 q3, d3, d4 \n" + "subs %2, %2, #1 \n" + "vpadal.s16 %q6, q3 \n" + "bne 0b \n" + "1: \n" + : "+r"(pA), "+r"(pB), "+r"(nn), "+w"(_sum0), "+w"(_sum1), "+w"(_sum2), "+w"(_sum3) + : + : "cc", "memory", "q0", "q1", "q2", "q3"); + kk += remain * 4; + } +#else // NCNN_GNU_INLINE_ASM && !__aarch64__ + for (; kk + 3 < max_kk; kk += 4) + { + int16x8_t _a = vreinterpretq_s16_s8(vld1q_s8(pA)); + int16x4_t _a01 = vget_low_s16(_a); + int16x4_t _a23 = vget_high_s16(_a); +#if NCNN_GNU_INLINE_ASM && __aarch64__ + asm volatile("" + : + : "w"(_a)); +#endif + int8x8_t _b01 = vreinterpret_s8_s32(vld1_dup_s32((const int*)pB)); + int8x8_t _b23 = vext_s8(_b01, _b01, 2); + int16x8_t _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a01, 0))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a23, 0))); + _sum0 = vpadalq_s16(_sum0, _s); + _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a01, 1))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a23, 1))); + _sum1 = vpadalq_s16(_sum1, _s); + _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a01, 2))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a23, 2))); + _sum2 = vpadalq_s16(_sum2, _s); + _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a01, 3))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a23, 3))); + _sum3 = vpadalq_s16(_sum3, _s); + pA += 16; + pB += 4; + } +#endif // NCNN_GNU_INLINE_ASM && !__aarch64__ #endif // __ARM_FEATURE_DOTPROD for (; kk + 1 < max_kk; kk += 2) { - const int8x8_t _b0 = vreinterpret_s8_s16(vld1_dup_s16((const short*)(pB))); - const int16x4_t _a = vreinterpret_s16_s8(vld1_s8(pA)); + int8x8_t _b0 = vreinterpret_s8_s16(vld1_dup_s16((const short*)(pB))); + int16x4_t _a = vreinterpret_s16_s8(vld1_s8(pA)); _sum0 = vaddq_s32(_sum0, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 0))))); _sum1 = vaddq_s32(_sum1, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 1))))); _sum2 = vaddq_s32(_sum2, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 2))))); @@ -2415,12 +3018,12 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } if (kk < max_kk) { - const int8x8_t _b = vset_lane_s8(pB[0], vdup_n_s8(0), 0); - const int8x8_t _a = vreinterpret_s8_s32(vld1_lane_s32((const int*)pA, vdup_n_s32(0), 0)); - const int16x8_t _p0 = vmull_s8(_b, vdup_lane_s8(_a, 0)); - const int16x8_t _p1 = vmull_s8(_b, vdup_lane_s8(_a, 1)); - const int16x8_t _p2 = vmull_s8(_b, vdup_lane_s8(_a, 2)); - const int16x8_t _p3 = vmull_s8(_b, vdup_lane_s8(_a, 3)); + int8x8_t _b = vset_lane_s8(pB[0], vdup_n_s8(0), 0); + int8x8_t _a = vreinterpret_s8_s32(vld1_lane_s32((const int*)pA, vdup_n_s32(0), 0)); + int16x8_t _p0 = vmull_s8(_b, vdup_lane_s8(_a, 0)); + int16x8_t _p1 = vmull_s8(_b, vdup_lane_s8(_a, 1)); + int16x8_t _p2 = vmull_s8(_b, vdup_lane_s8(_a, 2)); + int16x8_t _p3 = vmull_s8(_b, vdup_lane_s8(_a, 3)); _sum0 = vaddq_s32(_sum0, vmovl_s16(vget_low_s16(_p0))); _sum1 = vaddq_s32(_sum1, vmovl_s16(vget_low_s16(_p1))); _sum2 = vaddq_s32(_sum2, vmovl_s16(vget_low_s16(_p2))); @@ -2429,8 +3032,8 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB += 1; } - const float32x4_t _bd0 = vdupq_n_f32(pB_descales[0]); - const float32x4_t _ad = vld1q_f32(pA_descales); + float32x4_t _bd0 = vdupq_n_f32(pB_descales[0]); + float32x4_t _ad = vld1q_f32(pA_descales); _fsum0 = vmlaq_f32(_fsum0, vcvtq_f32_s32(_sum0), vmulq_n_f32(_bd0, vgetq_lane_f32(_ad, 0))); _fsum1 = vmlaq_f32(_fsum1, vcvtq_f32_s32(_sum1), vmulq_n_f32(_bd0, vgetq_lane_f32(_ad, 1))); _fsum2 = vmlaq_f32(_fsum2, vcvtq_f32_s32(_sum2), vmulq_n_f32(_bd0, vgetq_lane_f32(_ad, 2))); @@ -2448,27 +3051,28 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de outptr++; vst1q_lane_f32(outptr, _fsum3, 0); outptr++; + pB1 += full_K; + pB_descales1 += block_count; } } for (; ii + 1 < max_ii; ii += 2) { const signed char* pA0_block = pAT + ii * A_hstep; const float* pA_descales0_block = pAT_descales + ii * A_descales_hstep; - const signed char* pB = pBT; - const float* pB_descales = pBT_descales; - int jj = 0; + const signed char* pB8 = pBT + 4 * k0; + const float* pB_descales8 = pBT_descales + 4 * block_start; #if __aarch64__ for (; jj + 7 < max_jj; jj += 8) { - const signed char* pB0 = pB; - const signed char* pB1 = pB + 4 * K; - const float* pB_descales0 = pB_descales; - const float* pB_descales1 = pB_descales + 4 * ((K + block_size - 1) / block_size); - float32x4_t _fsum0 = vdupq_n_f32(0.f); - float32x4_t _fsum1 = vdupq_n_f32(0.f); - float32x4_t _fsum2 = vdupq_n_f32(0.f); - float32x4_t _fsum3 = vdupq_n_f32(0.f); + const signed char* pB0 = pB8; + const signed char* pB1 = pB8 + (size_t)4 * full_K; + const float* pB_descales0 = pB_descales8; + const float* pB_descales1 = pB_descales8 + (size_t)4 * block_count; + float32x4_t _fsum0 = k0 == 0 ? vdupq_n_f32(0.f) : vld1q_f32(outptr); + float32x4_t _fsum1 = k0 == 0 ? vdupq_n_f32(0.f) : vld1q_f32(outptr + 4); + float32x4_t _fsum2 = k0 == 0 ? vdupq_n_f32(0.f) : vld1q_f32(outptr + 8); + float32x4_t _fsum3 = k0 == 0 ? vdupq_n_f32(0.f) : vld1q_f32(outptr + 12); const signed char* pA = pA0_block; const float* pA_descales = pA_descales0_block; @@ -2480,6 +3084,7 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de int32x4_t _sum3 = vdupq_n_s32(0); const int max_kk = std::min(K - k, block_size); int kk = 0; +#if __ARM_FEATURE_DOTPROD #if __ARM_FEATURE_MATMUL_INT8 int32x4_t _msum0 = vdupq_n_s32(0); int32x4_t _msum1 = vdupq_n_s32(0); @@ -2487,11 +3092,11 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de int32x4_t _msum3 = vdupq_n_s32(0); for (; kk + 7 < max_kk; kk += 8) { - const int8x16_t _b0 = vld1q_s8(pB0); - const int8x16_t _b1 = vld1q_s8(pB0 + 16); - const int8x16_t _b2 = vld1q_s8(pB1); - const int8x16_t _b3 = vld1q_s8(pB1 + 16); - const int8x16_t _a0 = vld1q_s8(pA); + int8x16_t _b0 = vld1q_s8(pB0); + int8x16_t _b1 = vld1q_s8(pB0 + 16); + int8x16_t _b2 = vld1q_s8(pB1); + int8x16_t _b3 = vld1q_s8(pB1 + 16); + int8x16_t _a0 = vld1q_s8(pA); _msum0 = vmmlaq_s32(_msum0, _a0, _b0); _msum1 = vmmlaq_s32(_msum1, _a0, _b1); _msum2 = vmmlaq_s32(_msum2, _a0, _b2); @@ -2505,12 +3110,11 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de _sum2 = vcombine_s32(vget_high_s32(_msum0), vget_high_s32(_msum1)); _sum3 = vcombine_s32(vget_high_s32(_msum2), vget_high_s32(_msum3)); #endif // __ARM_FEATURE_MATMUL_INT8 -#if __ARM_FEATURE_DOTPROD for (; kk + 3 < max_kk; kk += 4) { - const int8x16_t _b0 = vld1q_s8(pB0); - const int8x16_t _b1 = vld1q_s8(pB1); - const int8x16_t _a = vcombine_s8(vld1_s8(pA), vdup_n_s8(0)); + int8x16_t _b0 = vld1q_s8(pB0); + int8x16_t _b1 = vld1q_s8(pB1); + int8x16_t _a = vcombine_s8(vld1_s8(pA), vdup_n_s8(0)); _sum0 = vdotq_laneq_s32(_sum0, _b0, _a, 0); _sum1 = vdotq_laneq_s32(_sum1, _b1, _a, 0); _sum2 = vdotq_laneq_s32(_sum2, _b0, _a, 1); @@ -2519,13 +3123,44 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB0 += 16; pB1 += 16; } +#else // __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk; kk += 4) + { + int16x4_t _a = vreinterpret_s16_s8(vld1_s8(pA)); +#if NCNN_GNU_INLINE_ASM && __aarch64__ + asm volatile("" + : + : "w"(_a)); +#endif + int8x16_t _b0 = vld1q_s8(pB0); + int8x16_t _b1 = vld1q_s8(pB1); + int8x8_t _b001 = vget_low_s8(_b0); + int8x8_t _b023 = vget_high_s8(_b0); + int8x8_t _b101 = vget_low_s8(_b1); + int8x8_t _b123 = vget_high_s8(_b1); + int16x8_t _s = vmull_s8(_b001, vreinterpret_s8_s16(vdup_lane_s16(_a, 0))); + _s = vmlal_s8(_s, _b023, vreinterpret_s8_s16(vdup_lane_s16(_a, 2))); + _sum0 = vpadalq_s16(_sum0, _s); + _s = vmull_s8(_b101, vreinterpret_s8_s16(vdup_lane_s16(_a, 0))); + _s = vmlal_s8(_s, _b123, vreinterpret_s8_s16(vdup_lane_s16(_a, 2))); + _sum1 = vpadalq_s16(_sum1, _s); + _s = vmull_s8(_b001, vreinterpret_s8_s16(vdup_lane_s16(_a, 1))); + _s = vmlal_s8(_s, _b023, vreinterpret_s8_s16(vdup_lane_s16(_a, 3))); + _sum2 = vpadalq_s16(_sum2, _s); + _s = vmull_s8(_b101, vreinterpret_s8_s16(vdup_lane_s16(_a, 1))); + _s = vmlal_s8(_s, _b123, vreinterpret_s8_s16(vdup_lane_s16(_a, 3))); + _sum3 = vpadalq_s16(_sum3, _s); + pA += 8; + pB0 += 16; + pB1 += 16; + } #endif // __ARM_FEATURE_DOTPROD for (; kk + 1 < max_kk; kk += 2) { - const int8x16_t _b = vcombine_s8(vld1_s8(pB0), vld1_s8(pB1)); - const int8x8_t _b0 = vget_low_s8(_b); - const int8x8_t _b1 = vget_high_s8(_b); - const int16x4_t _a = vreinterpret_s16_s32(vld1_lane_s32((const int*)pA, vdup_n_s32(0), 0)); + int8x16_t _b = vcombine_s8(vld1_s8(pB0), vld1_s8(pB1)); + int8x8_t _b0 = vget_low_s8(_b); + int8x8_t _b1 = vget_high_s8(_b); + int16x4_t _a = vreinterpret_s16_s32(vld1_lane_s32((const int*)pA, vdup_n_s32(0), 0)); _sum0 = vaddq_s32(_sum0, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 0))))); _sum1 = vaddq_s32(_sum1, vpaddlq_s16(vmull_s8(_b1, vreinterpret_s8_s16(vdup_lane_s16(_a, 0))))); _sum2 = vaddq_s32(_sum2, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 1))))); @@ -2536,10 +3171,10 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } if (kk < max_kk) { - const int8x8_t _b = vreinterpret_s8_s32(vld1_lane_s32((const int*)pB1, vld1_dup_s32((const int*)pB0), 1)); - const int8x8_t _a = vreinterpret_s8_s16(vld1_lane_s16((const short*)pA, vdup_n_s16(0), 0)); - const int16x8_t _p0 = vmull_s8(_b, vdup_lane_s8(_a, 0)); - const int16x8_t _p1 = vmull_s8(_b, vdup_lane_s8(_a, 1)); + int8x8_t _b = vreinterpret_s8_s32(vld1_lane_s32((const int*)pB1, vld1_dup_s32((const int*)pB0), 1)); + int8x8_t _a = vreinterpret_s8_s16(vld1_lane_s16((const short*)pA, vdup_n_s16(0), 0)); + int16x8_t _p0 = vmull_s8(_b, vdup_lane_s8(_a, 0)); + int16x8_t _p1 = vmull_s8(_b, vdup_lane_s8(_a, 1)); _sum0 = vaddq_s32(_sum0, vmovl_s16(vget_low_s16(_p0))); _sum1 = vaddq_s32(_sum1, vmovl_s16(vget_high_s16(_p0))); _sum2 = vaddq_s32(_sum2, vmovl_s16(vget_low_s16(_p1))); @@ -2549,9 +3184,9 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB1 += 4; } - const float32x4_t _bd0 = vld1q_f32(pB_descales0); - const float32x4_t _bd1 = vld1q_f32(pB_descales1); - const float32x2_t _ad = vld1_f32(pA_descales); + float32x4_t _bd0 = vld1q_f32(pB_descales0); + float32x4_t _bd1 = vld1q_f32(pB_descales1); + float32x2_t _ad = vld1_f32(pA_descales); _fsum0 = vmlaq_f32(_fsum0, vcvtq_f32_s32(_sum0), vmulq_lane_f32(_bd0, _ad, 0)); _fsum1 = vmlaq_f32(_fsum1, vcvtq_f32_s32(_sum1), vmulq_lane_f32(_bd1, _ad, 0)); _fsum2 = vmlaq_f32(_fsum2, vcvtq_f32_s32(_sum2), vmulq_lane_f32(_bd0, _ad, 1)); @@ -2562,9 +3197,6 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB_descales1 += 4; } - pB = pB1; - pB_descales = pB_descales1; - vst1q_f32(outptr, _fsum0); outptr += 4; vst1q_f32(outptr, _fsum1); @@ -2573,12 +3205,18 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de outptr += 4; vst1q_f32(outptr, _fsum3); outptr += 4; + pB8 += (size_t)8 * full_K; + pB_descales8 += (size_t)8 * block_count; } #endif // __aarch64__ + const signed char* pB4 = pB8; + const float* pB_descales4 = pB_descales8; for (; jj + 3 < max_jj; jj += 4) { - float32x4_t _fsum0 = vdupq_n_f32(0.f); - float32x4_t _fsum1 = vdupq_n_f32(0.f); + const signed char* pB = pB4; + const float* pB_descales = pB_descales4; + float32x4_t _fsum0 = k0 == 0 ? vdupq_n_f32(0.f) : vld1q_f32(outptr); + float32x4_t _fsum1 = k0 == 0 ? vdupq_n_f32(0.f) : vld1q_f32(outptr + 4); const signed char* pA = pA0_block; const float* pA_descales = pA_descales0_block; @@ -2588,14 +3226,15 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de int32x4_t _sum1 = vdupq_n_s32(0); const int max_kk = std::min(K - k, block_size); int kk = 0; +#if __ARM_FEATURE_DOTPROD #if __ARM_FEATURE_MATMUL_INT8 int32x4_t _msum0 = vdupq_n_s32(0); int32x4_t _msum1 = vdupq_n_s32(0); for (; kk + 7 < max_kk; kk += 8) { - const int8x16_t _b0 = vld1q_s8(pB + 0); - const int8x16_t _b1 = vld1q_s8(pB + 16); - const int8x16_t _a0 = vld1q_s8(pA); + int8x16_t _b0 = vld1q_s8(pB + 0); + int8x16_t _b1 = vld1q_s8(pB + 16); + int8x16_t _a0 = vld1q_s8(pA); _msum0 = vmmlaq_s32(_msum0, _a0, _b0); _msum1 = vmmlaq_s32(_msum1, _a0, _b1); pA += 16; @@ -2604,21 +3243,41 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de _sum0 = vcombine_s32(vget_low_s32(_msum0), vget_low_s32(_msum1)); _sum1 = vcombine_s32(vget_high_s32(_msum0), vget_high_s32(_msum1)); #endif // __ARM_FEATURE_MATMUL_INT8 -#if __ARM_FEATURE_DOTPROD for (; kk + 3 < max_kk; kk += 4) { - const int8x16_t _b0 = vld1q_s8(pB); - const int8x16_t _a = vcombine_s8(vld1_s8(pA), vdup_n_s8(0)); + int8x16_t _b0 = vld1q_s8(pB); + int8x16_t _a = vcombine_s8(vld1_s8(pA), vdup_n_s8(0)); _sum0 = vdotq_laneq_s32(_sum0, _b0, _a, 0); _sum1 = vdotq_laneq_s32(_sum1, _b0, _a, 1); pA += 8; pB += 16; } +#else // __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk; kk += 4) + { + int16x4_t _a = vreinterpret_s16_s8(vld1_s8(pA)); +#if NCNN_GNU_INLINE_ASM && __aarch64__ + asm volatile("" + : + : "w"(_a)); +#endif + int8x16_t _b = vld1q_s8(pB); + int8x8_t _b01 = vget_low_s8(_b); + int8x8_t _b23 = vget_high_s8(_b); + int16x8_t _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a, 0))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a, 2))); + _sum0 = vpadalq_s16(_sum0, _s); + _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a, 1))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a, 3))); + _sum1 = vpadalq_s16(_sum1, _s); + pA += 8; + pB += 16; + } #endif // __ARM_FEATURE_DOTPROD for (; kk + 1 < max_kk; kk += 2) { - const int8x8_t _b0 = vld1_s8(pB); - const int16x4_t _a = vreinterpret_s16_s32(vld1_lane_s32((const int*)pA, vdup_n_s32(0), 0)); + int8x8_t _b0 = vld1_s8(pB); + int16x4_t _a = vreinterpret_s16_s32(vld1_lane_s32((const int*)pA, vdup_n_s32(0), 0)); _sum0 = vaddq_s32(_sum0, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 0))))); _sum1 = vaddq_s32(_sum1, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 1))))); pA += 4; @@ -2626,18 +3285,18 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } if (kk < max_kk) { - const int8x8_t _b = vreinterpret_s8_s32(vld1_dup_s32((const int*)(pB))); - const int8x8_t _a = vreinterpret_s8_s16(vld1_lane_s16((const short*)pA, vdup_n_s16(0), 0)); - const int16x8_t _p0 = vmull_s8(_b, vdup_lane_s8(_a, 0)); - const int16x8_t _p1 = vmull_s8(_b, vdup_lane_s8(_a, 1)); + int8x8_t _b = vreinterpret_s8_s32(vld1_dup_s32((const int*)(pB))); + int8x8_t _a = vreinterpret_s8_s16(vld1_lane_s16((const short*)pA, vdup_n_s16(0), 0)); + int16x8_t _p0 = vmull_s8(_b, vdup_lane_s8(_a, 0)); + int16x8_t _p1 = vmull_s8(_b, vdup_lane_s8(_a, 1)); _sum0 = vaddq_s32(_sum0, vmovl_s16(vget_low_s16(_p0))); _sum1 = vaddq_s32(_sum1, vmovl_s16(vget_low_s16(_p1))); pA += 2; pB += 4; } - const float32x4_t _bd0 = vld1q_f32(pB_descales); - const float32x2_t _ad = vld1_f32(pA_descales); + float32x4_t _bd0 = vld1q_f32(pB_descales); + float32x2_t _ad = vld1_f32(pA_descales); _fsum0 = vmlaq_f32(_fsum0, vcvtq_f32_s32(_sum0), vmulq_lane_f32(_bd0, _ad, 0)); _fsum1 = vmlaq_f32(_fsum1, vcvtq_f32_s32(_sum1), vmulq_lane_f32(_bd0, _ad, 1)); @@ -2649,11 +3308,17 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de outptr += 4; vst1q_f32(outptr, _fsum1); outptr += 4; + pB4 += (size_t)4 * full_K; + pB_descales4 += (size_t)4 * block_count; } + const signed char* pB2 = pB4 - 2 * k0; + const float* pB_descales2 = pB_descales4 - 2 * block_start; for (; jj + 1 < max_jj; jj += 2) { - float32x4_t _fsum0 = vdupq_n_f32(0.f); - float32x4_t _fsum1 = vdupq_n_f32(0.f); + const signed char* pB = pB2; + const float* pB_descales = pB_descales2; + float32x4_t _fsum0 = k0 == 0 ? vdupq_n_f32(0.f) : vcombine_f32(vld1_f32(outptr), vdup_n_f32(0.f)); + float32x4_t _fsum1 = k0 == 0 ? vdupq_n_f32(0.f) : vcombine_f32(vld1_f32(outptr + 2), vdup_n_f32(0.f)); const signed char* pA = pA0_block; const float* pA_descales = pA_descales0_block; @@ -2663,12 +3328,13 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de int32x4_t _sum1 = vdupq_n_s32(0); const int max_kk = std::min(K - k, block_size); int kk = 0; +#if __ARM_FEATURE_DOTPROD #if __ARM_FEATURE_MATMUL_INT8 int32x4_t _msum0 = vdupq_n_s32(0); for (; kk + 7 < max_kk; kk += 8) { - const int8x16_t _b0 = vld1q_s8(pB + 0); - const int8x16_t _a0 = vld1q_s8(pA); + int8x16_t _b0 = vld1q_s8(pB + 0); + int8x16_t _a0 = vld1q_s8(pA); _msum0 = vmmlaq_s32(_msum0, _a0, _b0); pA += 16; pB += 16; @@ -2676,21 +3342,42 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de _sum0 = vcombine_s32(vget_low_s32(_msum0), vdup_n_s32(0)); _sum1 = vcombine_s32(vget_high_s32(_msum0), vdup_n_s32(0)); #endif // __ARM_FEATURE_MATMUL_INT8 -#if __ARM_FEATURE_DOTPROD for (; kk + 3 < max_kk; kk += 4) { - const int8x16_t _b0 = vcombine_s8(vld1_s8(pB), vdup_n_s8(0)); - const int8x16_t _a = vcombine_s8(vld1_s8(pA), vdup_n_s8(0)); + int8x16_t _b0 = vcombine_s8(vld1_s8(pB), vdup_n_s8(0)); + int8x16_t _a = vcombine_s8(vld1_s8(pA), vdup_n_s8(0)); _sum0 = vdotq_laneq_s32(_sum0, _b0, _a, 0); _sum1 = vdotq_laneq_s32(_sum1, _b0, _a, 1); pA += 8; pB += 8; } +#else // __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk; kk += 4) + { + int16x4_t _a = vreinterpret_s16_s8(vld1_s8(pA)); +#if NCNN_GNU_INLINE_ASM && __aarch64__ + asm volatile("" + : + : "w"(_a)); +#endif + int8x8_t _b = vld1_s8(pB); + int32x2_t _b02 = vreinterpret_s32_s8(_b); + int8x8_t _b01 = vreinterpret_s8_s32(vdup_lane_s32(_b02, 0)); + int8x8_t _b23 = vreinterpret_s8_s32(vdup_lane_s32(_b02, 1)); + int16x8_t _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a, 0))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a, 2))); + _sum0 = vpadalq_s16(_sum0, _s); + _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a, 1))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a, 3))); + _sum1 = vpadalq_s16(_sum1, _s); + pA += 8; + pB += 8; + } #endif // __ARM_FEATURE_DOTPROD for (; kk + 1 < max_kk; kk += 2) { - const int8x8_t _b0 = vreinterpret_s8_s32(vld1_dup_s32((const int*)(pB))); - const int16x4_t _a = vreinterpret_s16_s32(vld1_lane_s32((const int*)pA, vdup_n_s32(0), 0)); + int8x8_t _b0 = vreinterpret_s8_s32(vld1_dup_s32((const int*)(pB))); + int16x4_t _a = vreinterpret_s16_s32(vld1_lane_s32((const int*)pA, vdup_n_s32(0), 0)); _sum0 = vaddq_s32(_sum0, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 0))))); _sum1 = vaddq_s32(_sum1, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 1))))); pA += 4; @@ -2698,18 +3385,18 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } if (kk < max_kk) { - const int8x8_t _b = vreinterpret_s8_s16(vld1_dup_s16((const short*)(pB))); - const int8x8_t _a = vreinterpret_s8_s16(vld1_lane_s16((const short*)pA, vdup_n_s16(0), 0)); - const int16x8_t _p0 = vmull_s8(_b, vdup_lane_s8(_a, 0)); - const int16x8_t _p1 = vmull_s8(_b, vdup_lane_s8(_a, 1)); + int8x8_t _b = vreinterpret_s8_s16(vld1_dup_s16((const short*)(pB))); + int8x8_t _a = vreinterpret_s8_s16(vld1_lane_s16((const short*)pA, vdup_n_s16(0), 0)); + int16x8_t _p0 = vmull_s8(_b, vdup_lane_s8(_a, 0)); + int16x8_t _p1 = vmull_s8(_b, vdup_lane_s8(_a, 1)); _sum0 = vaddq_s32(_sum0, vmovl_s16(vget_low_s16(_p0))); _sum1 = vaddq_s32(_sum1, vmovl_s16(vget_low_s16(_p1))); pA += 2; pB += 2; } - const float32x4_t _bd0 = vcombine_f32(vld1_f32(pB_descales), vdup_n_f32(0.f)); - const float32x2_t _ad = vld1_f32(pA_descales); + float32x4_t _bd0 = vcombine_f32(vld1_f32(pB_descales), vdup_n_f32(0.f)); + float32x2_t _ad = vld1_f32(pA_descales); _fsum0 = vmlaq_f32(_fsum0, vcvtq_f32_s32(_sum0), vmulq_lane_f32(_bd0, _ad, 0)); _fsum1 = vmlaq_f32(_fsum1, vcvtq_f32_s32(_sum1), vmulq_lane_f32(_bd0, _ad, 1)); @@ -2721,11 +3408,17 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de outptr += 2; vst1_f32(outptr, vget_low_f32(_fsum1)); outptr += 2; + pB2 += (size_t)2 * full_K; + pB_descales2 += (size_t)2 * block_count; } + const signed char* pB1 = pB2 - k0; + const float* pB_descales1 = pB_descales2 - block_start; for (; jj < max_jj; jj++) { - float32x4_t _fsum0 = vdupq_n_f32(0.f); - float32x4_t _fsum1 = vdupq_n_f32(0.f); + const signed char* pB = pB1; + const float* pB_descales = pB_descales1; + float32x4_t _fsum0 = k0 == 0 ? vdupq_n_f32(0.f) : vsetq_lane_f32(outptr[0], vdupq_n_f32(0.f), 0); + float32x4_t _fsum1 = k0 == 0 ? vdupq_n_f32(0.f) : vsetq_lane_f32(outptr[1], vdupq_n_f32(0.f), 0); const signed char* pA = pA0_block; const float* pA_descales = pA_descales0_block; @@ -2735,12 +3428,13 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de int32x4_t _sum1 = vdupq_n_s32(0); const int max_kk = std::min(K - k, block_size); int kk = 0; +#if __ARM_FEATURE_DOTPROD #if __ARM_FEATURE_MATMUL_INT8 for (; kk + 7 < max_kk; kk += 8) { - const int8x16_t _a = vld1q_s8(pA); - const int8x16_t _b0 = vcombine_s8(vreinterpret_s8_s32(vld1_dup_s32((const int*)pB)), vdup_n_s8(0)); - const int8x16_t _b1 = vcombine_s8(vreinterpret_s8_s32(vld1_dup_s32((const int*)(pB + 4))), vdup_n_s8(0)); + int8x16_t _a = vld1q_s8(pA); + int8x16_t _b0 = vcombine_s8(vreinterpret_s8_s32(vld1_dup_s32((const int*)pB)), vdup_n_s8(0)); + int8x16_t _b1 = vcombine_s8(vreinterpret_s8_s32(vld1_dup_s32((const int*)(pB + 4))), vdup_n_s8(0)); _sum0 = vdotq_laneq_s32(_sum0, _b0, _a, 0); _sum0 = vdotq_laneq_s32(_sum0, _b1, _a, 1); _sum1 = vdotq_laneq_s32(_sum1, _b0, _a, 2); @@ -2749,21 +3443,40 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB += 8; } #endif // __ARM_FEATURE_MATMUL_INT8 -#if __ARM_FEATURE_DOTPROD for (; kk + 3 < max_kk; kk += 4) { - const int8x16_t _b0 = vcombine_s8(vreinterpret_s8_s32(vld1_dup_s32((const int*)pB)), vdup_n_s8(0)); - const int8x16_t _a = vcombine_s8(vld1_s8(pA), vdup_n_s8(0)); + int8x16_t _b0 = vcombine_s8(vreinterpret_s8_s32(vld1_dup_s32((const int*)pB)), vdup_n_s8(0)); + int8x16_t _a = vcombine_s8(vld1_s8(pA), vdup_n_s8(0)); _sum0 = vdotq_laneq_s32(_sum0, _b0, _a, 0); _sum1 = vdotq_laneq_s32(_sum1, _b0, _a, 1); pA += 8; pB += 4; } +#else // __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk; kk += 4) + { + int16x4_t _a = vreinterpret_s16_s8(vld1_s8(pA)); +#if NCNN_GNU_INLINE_ASM && __aarch64__ + asm volatile("" + : + : "w"(_a)); +#endif + int8x8_t _b01 = vreinterpret_s8_s32(vld1_dup_s32((const int*)pB)); + int8x8_t _b23 = vext_s8(_b01, _b01, 2); + int16x8_t _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a, 0))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a, 2))); + _sum0 = vpadalq_s16(_sum0, _s); + _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a, 1))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a, 3))); + _sum1 = vpadalq_s16(_sum1, _s); + pA += 8; + pB += 4; + } #endif // __ARM_FEATURE_DOTPROD for (; kk + 1 < max_kk; kk += 2) { - const int8x8_t _b0 = vreinterpret_s8_s16(vld1_dup_s16((const short*)(pB))); - const int16x4_t _a = vreinterpret_s16_s32(vld1_lane_s32((const int*)pA, vdup_n_s32(0), 0)); + int8x8_t _b0 = vreinterpret_s8_s16(vld1_dup_s16((const short*)(pB))); + int16x4_t _a = vreinterpret_s16_s32(vld1_lane_s32((const int*)pA, vdup_n_s32(0), 0)); _sum0 = vaddq_s32(_sum0, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 0))))); _sum1 = vaddq_s32(_sum1, vpaddlq_s16(vmull_s8(_b0, vreinterpret_s8_s16(vdup_lane_s16(_a, 1))))); pA += 4; @@ -2771,18 +3484,18 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } if (kk < max_kk) { - const int8x8_t _b = vset_lane_s8(pB[0], vdup_n_s8(0), 0); - const int8x8_t _a = vreinterpret_s8_s16(vld1_lane_s16((const short*)pA, vdup_n_s16(0), 0)); - const int16x8_t _p0 = vmull_s8(_b, vdup_lane_s8(_a, 0)); - const int16x8_t _p1 = vmull_s8(_b, vdup_lane_s8(_a, 1)); + int8x8_t _b = vset_lane_s8(pB[0], vdup_n_s8(0), 0); + int8x8_t _a = vreinterpret_s8_s16(vld1_lane_s16((const short*)pA, vdup_n_s16(0), 0)); + int16x8_t _p0 = vmull_s8(_b, vdup_lane_s8(_a, 0)); + int16x8_t _p1 = vmull_s8(_b, vdup_lane_s8(_a, 1)); _sum0 = vaddq_s32(_sum0, vmovl_s16(vget_low_s16(_p0))); _sum1 = vaddq_s32(_sum1, vmovl_s16(vget_low_s16(_p1))); pA += 2; pB += 1; } - const float32x4_t _bd0 = vdupq_n_f32(pB_descales[0]); - const float32x2_t _ad = vld1_f32(pA_descales); + float32x4_t _bd0 = vdupq_n_f32(pB_descales[0]); + float32x2_t _ad = vld1_f32(pA_descales); _fsum0 = vmlaq_f32(_fsum0, vcvtq_f32_s32(_sum0), vmulq_lane_f32(_bd0, _ad, 0)); _fsum1 = vmlaq_f32(_fsum1, vcvtq_f32_s32(_sum1), vmulq_lane_f32(_bd0, _ad, 1)); @@ -2794,25 +3507,26 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de outptr++; vst1q_lane_f32(outptr, _fsum1, 0); outptr++; + pB1 += full_K; + pB_descales1 += block_count; } } for (; ii < max_ii; ii++) { const signed char* pA0_block = pAT + ii * A_hstep; const float* pA_descales0_block = pAT_descales + ii * A_descales_hstep; - const signed char* pB = pBT; - const float* pB_descales = pBT_descales; - int jj = 0; + const signed char* pB8 = pBT + 4 * k0; + const float* pB_descales8 = pBT_descales + 4 * block_start; #if __aarch64__ for (; jj + 7 < max_jj; jj += 8) { - const signed char* pB0 = pB; - const signed char* pB1 = pB + 4 * K; - const float* pB_descales0 = pB_descales; - const float* pB_descales1 = pB_descales + 4 * ((K + block_size - 1) / block_size); - float32x4_t _fsum0 = vdupq_n_f32(0.f); - float32x4_t _fsum1 = vdupq_n_f32(0.f); + const signed char* pB0 = pB8; + const signed char* pB1 = pB8 + (size_t)4 * full_K; + const float* pB_descales0 = pB_descales8; + const float* pB_descales1 = pB_descales8 + (size_t)4 * block_count; + float32x4_t _fsum0 = k0 == 0 ? vdupq_n_f32(0.f) : vld1q_f32(outptr); + float32x4_t _fsum1 = k0 == 0 ? vdupq_n_f32(0.f) : vld1q_f32(outptr + 4); const signed char* pA0 = pA0_block; const float* pA_descales0 = pA_descales0_block; @@ -2822,6 +3536,7 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de int32x4_t _sum1 = vdupq_n_s32(0); const int max_kk = std::min(K - k, block_size); int kk = 0; +#if __ARM_FEATURE_DOTPROD #if __ARM_FEATURE_MATMUL_INT8 int32x4_t _msum0 = vdupq_n_s32(0); int32x4_t _msum1 = vdupq_n_s32(0); @@ -2829,11 +3544,11 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de int32x4_t _msum3 = vdupq_n_s32(0); for (; kk + 7 < max_kk; kk += 8) { - const int8x16_t _b0 = vld1q_s8(pB0); - const int8x16_t _b1 = vld1q_s8(pB0 + 16); - const int8x16_t _b2 = vld1q_s8(pB1); - const int8x16_t _b3 = vld1q_s8(pB1 + 16); - const int8x16_t _a0 = vcombine_s8(vld1_s8(pA0 + kk), vdup_n_s8(0)); + int8x16_t _b0 = vld1q_s8(pB0); + int8x16_t _b1 = vld1q_s8(pB0 + 16); + int8x16_t _b2 = vld1q_s8(pB1); + int8x16_t _b3 = vld1q_s8(pB1 + 16); + int8x16_t _a0 = vcombine_s8(vld1_s8(pA0 + kk), vdup_n_s8(0)); _msum0 = vmmlaq_s32(_msum0, _a0, _b0); _msum1 = vmmlaq_s32(_msum1, _a0, _b1); _msum2 = vmmlaq_s32(_msum2, _a0, _b2); @@ -2844,24 +3559,38 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de _sum0 = vcombine_s32(vget_low_s32(_msum0), vget_low_s32(_msum1)); _sum1 = vcombine_s32(vget_low_s32(_msum2), vget_low_s32(_msum3)); #endif // __ARM_FEATURE_MATMUL_INT8 -#if __ARM_FEATURE_DOTPROD for (; kk + 3 < max_kk; kk += 4) { - const int8x16_t _b0 = vld1q_s8(pB0); - const int8x16_t _b1 = vld1q_s8(pB1); - const int8x16_t _a0 = vreinterpretq_s8_s32(vdupq_lane_s32(vld1_lane_s32((const int*)(pA0 + kk), vdup_n_s32(0), 0), 0)); + int8x16_t _b0 = vld1q_s8(pB0); + int8x16_t _b1 = vld1q_s8(pB1); + int8x16_t _a0 = vreinterpretq_s8_s32(vdupq_lane_s32(vld1_lane_s32((const int*)(pA0 + kk), vdup_n_s32(0), 0), 0)); _sum0 = vdotq_s32(_sum0, _b0, _a0); _sum1 = vdotq_s32(_sum1, _b1, _a0); pB0 += 16; pB1 += 16; } +#else // __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk; kk += 4) + { + int16x4_t _a = vreinterpret_s16_s32(vld1_lane_s32((const int*)(pA0 + kk), vdup_n_s32(0), 0)); + int8x16_t _b0 = vld1q_s8(pB0); + int8x16_t _b1 = vld1q_s8(pB1); + int16x8_t _s = vmull_s8(vget_low_s8(_b0), vreinterpret_s8_s16(vdup_lane_s16(_a, 0))); + _s = vmlal_s8(_s, vget_high_s8(_b0), vreinterpret_s8_s16(vdup_lane_s16(_a, 1))); + _sum0 = vpadalq_s16(_sum0, _s); + _s = vmull_s8(vget_low_s8(_b1), vreinterpret_s8_s16(vdup_lane_s16(_a, 0))); + _s = vmlal_s8(_s, vget_high_s8(_b1), vreinterpret_s8_s16(vdup_lane_s16(_a, 1))); + _sum1 = vpadalq_s16(_sum1, _s); + pB0 += 16; + pB1 += 16; + } #endif // __ARM_FEATURE_DOTPROD for (; kk + 1 < max_kk; kk += 2) { - const int8x16_t _b = vcombine_s8(vld1_s8(pB0), vld1_s8(pB1)); - const int8x8_t _b0 = vget_low_s8(_b); - const int8x8_t _b1 = vget_high_s8(_b); - const int8x8_t _a0 = vreinterpret_s8_s16(vdup_lane_s16(vld1_lane_s16((const short*)(pA0 + kk), vdup_n_s16(0), 0), 0)); + int8x16_t _b = vcombine_s8(vld1_s8(pB0), vld1_s8(pB1)); + int8x8_t _b0 = vget_low_s8(_b); + int8x8_t _b1 = vget_high_s8(_b); + int8x8_t _a0 = vreinterpret_s8_s16(vdup_lane_s16(vld1_lane_s16((const short*)(pA0 + kk), vdup_n_s16(0), 0), 0)); _sum0 = vaddq_s32(_sum0, vpaddlq_s16(vmull_s8(_b0, _a0))); _sum1 = vaddq_s32(_sum1, vpaddlq_s16(vmull_s8(_b1, _a0))); pB0 += 8; @@ -2869,17 +3598,17 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } if (kk < max_kk) { - const int8x8_t _b = vreinterpret_s8_s32(vld1_lane_s32((const int*)pB1, vld1_dup_s32((const int*)pB0), 1)); - const int8x8_t _a0 = vld1_lane_s8(pA0 + kk, vdup_n_s8(0), 0); - const int16x8_t _p0 = vmull_s8(_b, vdup_lane_s8(_a0, 0)); + int8x8_t _b = vreinterpret_s8_s32(vld1_lane_s32((const int*)pB1, vld1_dup_s32((const int*)pB0), 1)); + int8x8_t _a0 = vld1_lane_s8(pA0 + kk, vdup_n_s8(0), 0); + int16x8_t _p0 = vmull_s8(_b, vdup_lane_s8(_a0, 0)); _sum0 = vaddq_s32(_sum0, vmovl_s16(vget_low_s16(_p0))); _sum1 = vaddq_s32(_sum1, vmovl_s16(vget_high_s16(_p0))); pB0 += 4; pB1 += 4; } - const float32x4_t _bd0 = vld1q_f32(pB_descales0); - const float32x4_t _bd1 = vld1q_f32(pB_descales1); + float32x4_t _bd0 = vld1q_f32(pB_descales0); + float32x4_t _bd1 = vld1q_f32(pB_descales1); const float _ad0 = pA_descales0[0]; _fsum0 = vmlaq_f32(_fsum0, vcvtq_f32_s32(_sum0), vmulq_n_f32(_bd0, _ad0)); _fsum1 = vmlaq_f32(_fsum1, vcvtq_f32_s32(_sum1), vmulq_n_f32(_bd1, _ad0)); @@ -2890,18 +3619,21 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB_descales1 += 4; } - pB = pB1; - pB_descales = pB_descales1; - vst1q_f32(outptr, _fsum0); outptr += 4; vst1q_f32(outptr, _fsum1); outptr += 4; + pB8 += (size_t)8 * full_K; + pB_descales8 += (size_t)8 * block_count; } #endif // __aarch64__ + const signed char* pB4 = pB8; + const float* pB_descales4 = pB_descales8; for (; jj + 3 < max_jj; jj += 4) { - float32x4_t _fsum0 = vdupq_n_f32(0.f); + const signed char* pB = pB4; + const float* pB_descales = pB_descales4; + float32x4_t _fsum0 = k0 == 0 ? vdupq_n_f32(0.f) : vld1q_f32(outptr); const signed char* pA0 = pA0_block; const float* pA_descales0 = pA_descales0_block; @@ -2910,46 +3642,56 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de int32x4_t _sum0 = vdupq_n_s32(0); const int max_kk = std::min(K - k, block_size); int kk = 0; +#if __ARM_FEATURE_DOTPROD #if __ARM_FEATURE_MATMUL_INT8 int32x4_t _msum0 = vdupq_n_s32(0); int32x4_t _msum1 = vdupq_n_s32(0); for (; kk + 7 < max_kk; kk += 8) { - const int8x16_t _b0 = vld1q_s8(pB + 0); - const int8x16_t _b1 = vld1q_s8(pB + 16); - const int8x16_t _a0 = vcombine_s8(vld1_s8(pA0 + kk), vdup_n_s8(0)); + int8x16_t _b0 = vld1q_s8(pB + 0); + int8x16_t _b1 = vld1q_s8(pB + 16); + int8x16_t _a0 = vcombine_s8(vld1_s8(pA0 + kk), vdup_n_s8(0)); _msum0 = vmmlaq_s32(_msum0, _a0, _b0); _msum1 = vmmlaq_s32(_msum1, _a0, _b1); pB += 32; } _sum0 = vcombine_s32(vget_low_s32(_msum0), vget_low_s32(_msum1)); #endif // __ARM_FEATURE_MATMUL_INT8 -#if __ARM_FEATURE_DOTPROD for (; kk + 3 < max_kk; kk += 4) { - const int8x16_t _b0 = vld1q_s8(pB); - const int8x16_t _a0 = vreinterpretq_s8_s32(vdupq_lane_s32(vld1_lane_s32((const int*)(pA0 + kk), vdup_n_s32(0), 0), 0)); + int8x16_t _b0 = vld1q_s8(pB); + int8x16_t _a0 = vreinterpretq_s8_s32(vdupq_lane_s32(vld1_lane_s32((const int*)(pA0 + kk), vdup_n_s32(0), 0), 0)); _sum0 = vdotq_s32(_sum0, _b0, _a0); pB += 16; } -#endif // __ARM_FEATURE_DOTPROD - for (; kk + 1 < max_kk; kk += 2) +#else // __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk; kk += 4) + { + int16x4_t _a = vreinterpret_s16_s32(vld1_lane_s32((const int*)(pA0 + kk), vdup_n_s32(0), 0)); + int8x16_t _b = vld1q_s8(pB); + int16x8_t _s = vmull_s8(vget_low_s8(_b), vreinterpret_s8_s16(vdup_lane_s16(_a, 0))); + _s = vmlal_s8(_s, vget_high_s8(_b), vreinterpret_s8_s16(vdup_lane_s16(_a, 1))); + _sum0 = vpadalq_s16(_sum0, _s); + pB += 16; + } +#endif // __ARM_FEATURE_DOTPROD + for (; kk + 1 < max_kk; kk += 2) { - const int8x8_t _b0 = vld1_s8(pB); - const int8x8_t _a0 = vreinterpret_s8_s16(vdup_lane_s16(vld1_lane_s16((const short*)(pA0 + kk), vdup_n_s16(0), 0), 0)); + int8x8_t _b0 = vld1_s8(pB); + int8x8_t _a0 = vreinterpret_s8_s16(vdup_lane_s16(vld1_lane_s16((const short*)(pA0 + kk), vdup_n_s16(0), 0), 0)); _sum0 = vaddq_s32(_sum0, vpaddlq_s16(vmull_s8(_b0, _a0))); pB += 8; } if (kk < max_kk) { - const int8x8_t _b = vreinterpret_s8_s32(vld1_dup_s32((const int*)(pB))); - const int8x8_t _a0 = vld1_lane_s8(pA0 + kk, vdup_n_s8(0), 0); - const int16x8_t _p0 = vmull_s8(_b, vdup_lane_s8(_a0, 0)); + int8x8_t _b = vreinterpret_s8_s32(vld1_dup_s32((const int*)(pB))); + int8x8_t _a0 = vld1_lane_s8(pA0 + kk, vdup_n_s8(0), 0); + int16x8_t _p0 = vmull_s8(_b, vdup_lane_s8(_a0, 0)); _sum0 = vaddq_s32(_sum0, vmovl_s16(vget_low_s16(_p0))); pB += 4; } - const float32x4_t _bd0 = vld1q_f32(pB_descales); + float32x4_t _bd0 = vld1q_f32(pB_descales); const float _ad0 = pA_descales0[0]; _fsum0 = vmlaq_f32(_fsum0, vcvtq_f32_s32(_sum0), vmulq_n_f32(_bd0, _ad0)); @@ -2960,10 +3702,16 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de vst1q_f32(outptr, _fsum0); outptr += 4; + pB4 += (size_t)4 * full_K; + pB_descales4 += (size_t)4 * block_count; } + const signed char* pB2 = pB4 - 2 * k0; + const float* pB_descales2 = pB_descales4 - 2 * block_start; for (; jj + 1 < max_jj; jj += 2) { - float32x4_t _fsum0 = vdupq_n_f32(0.f); + const signed char* pB = pB2; + const float* pB_descales = pB_descales2; + float32x4_t _fsum0 = k0 == 0 ? vdupq_n_f32(0.f) : vcombine_f32(vld1_f32(outptr), vdup_n_f32(0.f)); const signed char* pA0 = pA0_block; const float* pA_descales0 = pA_descales0_block; @@ -2972,43 +3720,56 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de int32x4_t _sum0 = vdupq_n_s32(0); const int max_kk = std::min(K - k, block_size); int kk = 0; +#if __ARM_FEATURE_DOTPROD #if __ARM_FEATURE_MATMUL_INT8 int32x4_t _msum0 = vdupq_n_s32(0); for (; kk + 7 < max_kk; kk += 8) { - const int8x16_t _b0 = vld1q_s8(pB + 0); - const int8x16_t _a0 = vcombine_s8(vld1_s8(pA0 + kk), vdup_n_s8(0)); + int8x16_t _b0 = vld1q_s8(pB + 0); + int8x16_t _a0 = vcombine_s8(vld1_s8(pA0 + kk), vdup_n_s8(0)); _msum0 = vmmlaq_s32(_msum0, _a0, _b0); pB += 16; } _sum0 = vcombine_s32(vget_low_s32(_msum0), vdup_n_s32(0)); #endif // __ARM_FEATURE_MATMUL_INT8 -#if __ARM_FEATURE_DOTPROD for (; kk + 3 < max_kk; kk += 4) { - const int8x16_t _b0 = vcombine_s8(vld1_s8(pB), vdup_n_s8(0)); - const int8x16_t _a0 = vreinterpretq_s8_s32(vdupq_lane_s32(vld1_lane_s32((const int*)(pA0 + kk), vdup_n_s32(0), 0), 0)); + int8x16_t _b0 = vcombine_s8(vld1_s8(pB), vdup_n_s8(0)); + int8x16_t _a0 = vreinterpretq_s8_s32(vdupq_lane_s32(vld1_lane_s32((const int*)(pA0 + kk), vdup_n_s32(0), 0), 0)); _sum0 = vdotq_s32(_sum0, _b0, _a0); pB += 8; } +#else // __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk; kk += 4) + { + int16x4_t _a = vreinterpret_s16_s32(vld1_lane_s32((const int*)(pA0 + kk), vdup_n_s32(0), 0)); + int8x8_t _b = vld1_s8(pB); + int32x2_t _b02 = vreinterpret_s32_s8(_b); + int8x8_t _b01 = vreinterpret_s8_s32(vdup_lane_s32(_b02, 0)); + int8x8_t _b23 = vreinterpret_s8_s32(vdup_lane_s32(_b02, 1)); + int16x8_t _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a, 0))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a, 1))); + _sum0 = vpadalq_s16(_sum0, _s); + pB += 8; + } #endif // __ARM_FEATURE_DOTPROD for (; kk + 1 < max_kk; kk += 2) { - const int8x8_t _b0 = vreinterpret_s8_s32(vld1_dup_s32((const int*)(pB))); - const int8x8_t _a0 = vreinterpret_s8_s16(vdup_lane_s16(vld1_lane_s16((const short*)(pA0 + kk), vdup_n_s16(0), 0), 0)); + int8x8_t _b0 = vreinterpret_s8_s32(vld1_dup_s32((const int*)(pB))); + int8x8_t _a0 = vreinterpret_s8_s16(vdup_lane_s16(vld1_lane_s16((const short*)(pA0 + kk), vdup_n_s16(0), 0), 0)); _sum0 = vaddq_s32(_sum0, vpaddlq_s16(vmull_s8(_b0, _a0))); pB += 4; } if (kk < max_kk) { - const int8x8_t _b = vreinterpret_s8_s16(vld1_dup_s16((const short*)(pB))); - const int8x8_t _a0 = vld1_lane_s8(pA0 + kk, vdup_n_s8(0), 0); - const int16x8_t _p0 = vmull_s8(_b, vdup_lane_s8(_a0, 0)); + int8x8_t _b = vreinterpret_s8_s16(vld1_dup_s16((const short*)(pB))); + int8x8_t _a0 = vld1_lane_s8(pA0 + kk, vdup_n_s8(0), 0); + int16x8_t _p0 = vmull_s8(_b, vdup_lane_s8(_a0, 0)); _sum0 = vaddq_s32(_sum0, vmovl_s16(vget_low_s16(_p0))); pB += 2; } - const float32x4_t _bd0 = vcombine_f32(vld1_f32(pB_descales), vdup_n_f32(0.f)); + float32x4_t _bd0 = vcombine_f32(vld1_f32(pB_descales), vdup_n_f32(0.f)); const float _ad0 = pA_descales0[0]; _fsum0 = vmlaq_f32(_fsum0, vcvtq_f32_s32(_sum0), vmulq_n_f32(_bd0, _ad0)); @@ -3019,10 +3780,16 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de vst1_f32(outptr, vget_low_f32(_fsum0)); outptr += 2; + pB2 += (size_t)2 * full_K; + pB_descales2 += (size_t)2 * block_count; } + const signed char* pB1 = pB2 - k0; + const float* pB_descales1 = pB_descales2 - block_start; for (; jj < max_jj; jj++) { - float32x4_t _fsum0 = vdupq_n_f32(0.f); + const signed char* pB = pB1; + const float* pB_descales = pB_descales1; + float32x4_t _fsum0 = k0 == 0 ? vdupq_n_f32(0.f) : vsetq_lane_f32(outptr[0], vdupq_n_f32(0.f), 0); const signed char* pA0 = pA0_block; const float* pA_descales0 = pA_descales0_block; @@ -3034,29 +3801,40 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #if __ARM_FEATURE_DOTPROD for (; kk + 3 < max_kk; kk += 4) { - const int8x16_t _b0 = vcombine_s8(vreinterpret_s8_s32(vld1_dup_s32((const int*)pB)), vdup_n_s8(0)); - const int8x16_t _a0 = vreinterpretq_s8_s32(vdupq_lane_s32(vld1_lane_s32((const int*)(pA0 + kk), vdup_n_s32(0), 0), 0)); + int8x16_t _b0 = vcombine_s8(vreinterpret_s8_s32(vld1_dup_s32((const int*)pB)), vdup_n_s8(0)); + int8x16_t _a0 = vreinterpretq_s8_s32(vdupq_lane_s32(vld1_lane_s32((const int*)(pA0 + kk), vdup_n_s32(0), 0), 0)); _sum0 = vdotq_s32(_sum0, _b0, _a0); pB += 4; } +#else // __ARM_FEATURE_DOTPROD + for (; kk + 3 < max_kk; kk += 4) + { + int16x4_t _a = vreinterpret_s16_s32(vld1_lane_s32((const int*)(pA0 + kk), vdup_n_s32(0), 0)); + int8x8_t _b01 = vreinterpret_s8_s32(vld1_dup_s32((const int*)pB)); + int8x8_t _b23 = vext_s8(_b01, _b01, 2); + int16x8_t _s = vmull_s8(_b01, vreinterpret_s8_s16(vdup_lane_s16(_a, 0))); + _s = vmlal_s8(_s, _b23, vreinterpret_s8_s16(vdup_lane_s16(_a, 1))); + _sum0 = vpadalq_s16(_sum0, _s); + pB += 4; + } #endif // __ARM_FEATURE_DOTPROD for (; kk + 1 < max_kk; kk += 2) { - const int8x8_t _b0 = vreinterpret_s8_s16(vld1_dup_s16((const short*)(pB))); - const int8x8_t _a0 = vreinterpret_s8_s16(vdup_lane_s16(vld1_lane_s16((const short*)(pA0 + kk), vdup_n_s16(0), 0), 0)); + int8x8_t _b0 = vreinterpret_s8_s16(vld1_dup_s16((const short*)(pB))); + int8x8_t _a0 = vreinterpret_s8_s16(vdup_lane_s16(vld1_lane_s16((const short*)(pA0 + kk), vdup_n_s16(0), 0), 0)); _sum0 = vaddq_s32(_sum0, vpaddlq_s16(vmull_s8(_b0, _a0))); pB += 2; } if (kk < max_kk) { - const int8x8_t _b = vset_lane_s8(pB[0], vdup_n_s8(0), 0); - const int8x8_t _a0 = vld1_lane_s8(pA0 + kk, vdup_n_s8(0), 0); - const int16x8_t _p0 = vmull_s8(_b, vdup_lane_s8(_a0, 0)); + int8x8_t _b = vset_lane_s8(pB[0], vdup_n_s8(0), 0); + int8x8_t _a0 = vld1_lane_s8(pA0 + kk, vdup_n_s8(0), 0); + int16x8_t _p0 = vmull_s8(_b, vdup_lane_s8(_a0, 0)); _sum0 = vaddq_s32(_sum0, vmovl_s16(vget_low_s16(_p0))); pB += 1; } - const float32x4_t _bd0 = vdupq_n_f32(pB_descales[0]); + float32x4_t _bd0 = vdupq_n_f32(pB_descales[0]); const float _ad0 = pA_descales0[0]; _fsum0 = vmlaq_f32(_fsum0, vcvtq_f32_s32(_sum0), vmulq_n_f32(_bd0, _ad0)); @@ -3067,6 +3845,8 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de vst1q_lane_f32(outptr, _fsum0, 0); outptr++; + pB1 += full_K; + pB_descales1 += block_count; } } #else @@ -3075,16 +3855,17 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de { const signed char* pA_block = pAT + ii * A_hstep; const float* pA_descales_block = pAT_descales + ii * A_descales_hstep; - const signed char* pB = pBT; - const float* pB_descales = pBT_descales; - int jj = 0; + const signed char* pB2 = pBT + 2 * k0; + const float* pB_descales2 = pBT_descales + 2 * block_start; for (; jj + 1 < max_jj; jj += 2) { - float fsum00 = 0.f; - float fsum01 = 0.f; - float fsum10 = 0.f; - float fsum11 = 0.f; + const signed char* pB = pB2; + const float* pB_descales = pB_descales2; + float fsum00 = k0 == 0 ? 0.f : outptr[0]; + float fsum01 = k0 == 0 ? 0.f : outptr[1]; + float fsum10 = k0 == 0 ? 0.f : outptr[2]; + float fsum11 = k0 == 0 ? 0.f : outptr[3]; const signed char* pA = pA_block; const float* pA_descales = pA_descales_block; @@ -3201,11 +3982,17 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de *outptr++ = fsum01; *outptr++ = fsum10; *outptr++ = fsum11; + pB2 += (size_t)2 * full_K; + pB_descales2 += (size_t)2 * block_count; } + const signed char* pB1 = pB2 - k0; + const float* pB_descales1 = pB_descales2 - block_start; for (; jj < max_jj; jj++) { - float fsum00 = 0.f; - float fsum10 = 0.f; + const signed char* pB = pB1; + const float* pB_descales = pB_descales1; + float fsum00 = k0 == 0 ? 0.f : outptr[0]; + float fsum10 = k0 == 0 ? 0.f : outptr[1]; const signed char* pA = pA_block; const float* pA_descales = pA_descales_block; @@ -3234,6 +4021,8 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de *outptr++ = fsum00; *outptr++ = fsum10; + pB1 += full_K; + pB_descales1 += block_count; } } #else @@ -3243,16 +4032,17 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de const signed char* pA1_block = pA0_block + A_hstep; const float* pA_descales0_block = pAT_descales + ii * A_descales_hstep; const float* pA_descales1_block = pA_descales0_block + A_descales_hstep; - const signed char* pB = pBT; - const float* pB_descales = pBT_descales; - int jj = 0; + const signed char* pB2 = pBT + 2 * k0; + const float* pB_descales2 = pBT_descales + 2 * block_start; for (; jj + 1 < max_jj; jj += 2) { - float fsum00 = 0.f; - float fsum01 = 0.f; - float fsum10 = 0.f; - float fsum11 = 0.f; + const signed char* pB = pB2; + const float* pB_descales = pB_descales2; + float fsum00 = k0 == 0 ? 0.f : outptr[0]; + float fsum01 = k0 == 0 ? 0.f : outptr[1]; + float fsum10 = k0 == 0 ? 0.f : outptr[2]; + float fsum11 = k0 == 0 ? 0.f : outptr[3]; const signed char* pA0 = pA0_block; const signed char* pA1 = pA1_block; @@ -3313,11 +4103,17 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de outptr++; outptr[0] = fsum11; outptr++; + pB2 += (size_t)2 * full_K; + pB_descales2 += (size_t)2 * block_count; } + const signed char* pB1 = pB2 - k0; + const float* pB_descales1 = pB_descales2 - block_start; for (; jj < max_jj; jj++) { - float fsum00 = 0.f; - float fsum10 = 0.f; + const signed char* pB = pB1; + const float* pB_descales = pB_descales1; + float fsum00 = k0 == 0 ? 0.f : outptr[0]; + float fsum10 = k0 == 0 ? 0.f : outptr[1]; const signed char* pA0 = pA0_block; const signed char* pA1 = pA1_block; @@ -3362,6 +4158,8 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de outptr++; outptr[0] = fsum10; outptr++; + pB1 += full_K; + pB_descales1 += block_count; } } #endif // __ARM_FEATURE_SIMD32 && NCNN_GNU_INLINE_ASM @@ -3369,14 +4167,15 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de { const signed char* pA0_block = pAT + ii * A_hstep; const float* pA_descales0_block = pAT_descales + ii * A_descales_hstep; - const signed char* pB = pBT; - const float* pB_descales = pBT_descales; - int jj = 0; + const signed char* pB2 = pBT + 2 * k0; + const float* pB_descales2 = pBT_descales + 2 * block_start; for (; jj + 1 < max_jj; jj += 2) { - float fsum00 = 0.f; - float fsum01 = 0.f; + const signed char* pB = pB2; + const float* pB_descales = pB_descales2; + float fsum00 = k0 == 0 ? 0.f : outptr[0]; + float fsum01 = k0 == 0 ? 0.f : outptr[1]; const signed char* pA0 = pA0_block; const float* pA_descales0 = pA_descales0_block; @@ -3425,10 +4224,16 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de outptr++; outptr[0] = fsum01; outptr++; + pB2 += (size_t)2 * full_K; + pB_descales2 += (size_t)2 * block_count; } + const signed char* pB1 = pB2 - k0; + const float* pB_descales1 = pB_descales2 - block_start; for (; jj < max_jj; jj++) { - float fsum00 = 0.f; + const signed char* pB = pB1; + const float* pB_descales = pB_descales1; + float fsum00 = k0 == 0 ? 0.f : outptr[0]; const signed char* pA0 = pA0_block; const float* pA_descales0 = pA_descales0_block; @@ -3462,12 +4267,14 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de outptr[0] = fsum00; outptr++; + pB1 += full_K; + pB_descales1 += block_count; } } #endif // __ARM_NEON } -static void get_optimal_tile_mnk_wq_int8(int M, int N, int K, int constant_TILE_M, int constant_TILE_N, int constant_TILE_K, int& TILE_M, int& TILE_N, int& TILE_K, int nT) +static void get_optimal_tile_mnk_wq_int8(int M, int N, int K, int block_size, int constant_TILE_M, int constant_TILE_N, int constant_TILE_K, int& TILE_M, int& TILE_N, int& TILE_K, int nT) { // resolve optimal tile size from cache size const size_t l2_cache_size = get_cpu_level2_cache_size(); @@ -3475,7 +4282,16 @@ static void get_optimal_tile_mnk_wq_int8(int M, int N, int K, int constant_TILE_ if (nT == 0) nT = get_physical_big_cpu_count(); - const int tile_size = std::max(1, (int)((float)l2_cache_size / 2 / sizeof(signed char) / std::max(1, K))); + int tile_size = (int)sqrtf((float)l2_cache_size / (2 * sizeof(signed char) + sizeof(float))); + + TILE_K = std::max(8, tile_size / 8 * 8); + if (TILE_K < block_size) + TILE_K = block_size; + else + TILE_K = TILE_K / block_size * block_size; + TILE_K = std::min(TILE_K, K); + + tile_size = std::max(1, (int)((float)l2_cache_size / 2 / sizeof(signed char) / std::max(1, TILE_K))); #if __aarch64__ TILE_M = M >= nT * 8 ? 8 : M >= nT * 4 ? 4 : M >= nT * 2 ? 2 : 1; @@ -3487,8 +4303,6 @@ static void get_optimal_tile_mnk_wq_int8(int M, int N, int K, int constant_TILE_ TILE_M = M >= nT * 2 ? 2 : 1; TILE_N = std::max(2, tile_size / 2 * 2); #endif - TILE_K = K; - if (N > 0) { const int nn_N = (N + TILE_N - 1) / TILE_N; @@ -3524,6 +4338,15 @@ static void get_optimal_tile_mnk_wq_int8(int M, int N, int K, int constant_TILE_ #endif } + if (constant_TILE_K > 0) + { + if (constant_TILE_K < block_size) + TILE_K = block_size; + else + TILE_K = constant_TILE_K / block_size * block_size; + TILE_K = std::min(TILE_K, K); + } + // one driver M tile follows the natural producer slab #if __aarch64__ TILE_M = std::min(TILE_M, 8); @@ -3533,7 +4356,6 @@ static void get_optimal_tile_mnk_wq_int8(int M, int N, int K, int constant_TILE_ TILE_M = std::min(TILE_M, 2); #endif - (void)constant_TILE_K; } static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta) @@ -3625,7 +4447,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b } if (broadcast_type_C == 4) { - const float32x4_t _c = beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta); + float32x4_t _c = beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta); _out0 = vaddq_f32(_out0, _c); _out1 = vaddq_f32(_out1, _c); _out2 = vaddq_f32(_out2, _c); @@ -3703,7 +4525,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b } if (broadcast_type_C == 4) { - const float32x2_t _c = beta == 1.f ? vld1_f32(pC) : vmul_n_f32(vld1_f32(pC), beta); + float32x2_t _c = beta == 1.f ? vld1_f32(pC) : vmul_n_f32(vld1_f32(pC), beta); _out0 = vadd_f32(_out0, _c); _out1 = vadd_f32(_out1, _c); _out2 = vadd_f32(_out2, _c); @@ -3871,7 +4693,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b { if (broadcast_type_C == 0) { - const float32x4_t _c = vdupq_n_f32(c0); + float32x4_t _c = vdupq_n_f32(c0); _out0 = vaddq_f32(_out0, _c); _out1 = vaddq_f32(_out1, _c); _out2 = vaddq_f32(_out2, _c); @@ -3902,19 +4724,21 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _out5 = vaddq_f32(_out5, (beta == 1.f ? vld1q_f32(pC + c_hstep * 2 + 4) : vmulq_n_f32(vld1q_f32(pC + c_hstep * 2 + 4), beta))); _out6 = vaddq_f32(_out6, (beta == 1.f ? vld1q_f32(pC + c_hstep * 3) : vmulq_n_f32(vld1q_f32(pC + c_hstep * 3), beta))); _out7 = vaddq_f32(_out7, (beta == 1.f ? vld1q_f32(pC + c_hstep * 3 + 4) : vmulq_n_f32(vld1q_f32(pC + c_hstep * 3 + 4), beta))); + pC += 8; } if (broadcast_type_C == 4) { - const float32x4_t _cc0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); + float32x4_t _cc0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); _out0 = vaddq_f32(_out0, _cc0); _out2 = vaddq_f32(_out2, _cc0); _out4 = vaddq_f32(_out4, _cc0); _out6 = vaddq_f32(_out6, _cc0); - const float32x4_t _cc1 = (beta == 1.f ? vld1q_f32(pC + 4) : vmulq_n_f32(vld1q_f32(pC + 4), beta)); + float32x4_t _cc1 = (beta == 1.f ? vld1q_f32(pC + 4) : vmulq_n_f32(vld1q_f32(pC + 4), beta)); _out1 = vaddq_f32(_out1, _cc1); _out3 = vaddq_f32(_out3, _cc1); _out5 = vaddq_f32(_out5, _cc1); _out7 = vaddq_f32(_out7, _cc1); + pC += 8; } } @@ -3943,8 +4767,6 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b outptr1 += 8; outptr2 += 8; outptr3 += 8; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 8; pp += 32; } #endif // __aarch64__ @@ -3959,7 +4781,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b { if (broadcast_type_C == 0) { - const float32x4_t _c = vdupq_n_f32(c0); + float32x4_t _c = vdupq_n_f32(c0); _out0 = vaddq_f32(_out0, _c); _out1 = vaddq_f32(_out1, _c); _out2 = vaddq_f32(_out2, _c); @@ -3978,14 +4800,16 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _out1 = vaddq_f32(_out1, (beta == 1.f ? vld1q_f32(pC + c_hstep) : vmulq_n_f32(vld1q_f32(pC + c_hstep), beta))); _out2 = vaddq_f32(_out2, (beta == 1.f ? vld1q_f32(pC + c_hstep * 2) : vmulq_n_f32(vld1q_f32(pC + c_hstep * 2), beta))); _out3 = vaddq_f32(_out3, (beta == 1.f ? vld1q_f32(pC + c_hstep * 3) : vmulq_n_f32(vld1q_f32(pC + c_hstep * 3), beta))); + pC += 4; } if (broadcast_type_C == 4) { - const float32x4_t _cc0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); + float32x4_t _cc0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); _out0 = vaddq_f32(_out0, _cc0); _out1 = vaddq_f32(_out1, _cc0); _out2 = vaddq_f32(_out2, _cc0); _out3 = vaddq_f32(_out3, _cc0); + pC += 4; } } @@ -4006,8 +4830,6 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b outptr1 += 4; outptr2 += 4; outptr3 += 4; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 4; pp += 16; } for (; jj + 1 < max_jj; jj += 2) @@ -4019,32 +4841,34 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b { if (broadcast_type_C == 0) { - const float32x4_t _c = vdupq_n_f32(c0); + float32x4_t _c = vdupq_n_f32(c0); _out0 = vaddq_f32(_out0, _c); _out1 = vaddq_f32(_out1, _c); } if (broadcast_type_C == 1 || broadcast_type_C == 2) { - const float32x4_t _c01 = vcombine_f32(vdup_n_f32(c0), vdup_n_f32(c1)); - const float32x4_t _c23 = vcombine_f32(vdup_n_f32(c2), vdup_n_f32(c3)); + float32x4_t _c01 = vcombine_f32(vdup_n_f32(c0), vdup_n_f32(c1)); + float32x4_t _c23 = vcombine_f32(vdup_n_f32(c2), vdup_n_f32(c3)); _out0 = vaddq_f32(_out0, _c01); _out1 = vaddq_f32(_out1, _c23); } if (broadcast_type_C == 3) { - const float32x4_t _c01 = vcombine_f32(vld1_f32(pC), vld1_f32(pC + c_hstep)); - const float32x4_t _c23 = vcombine_f32(vld1_f32(pC + c_hstep * 2), vld1_f32(pC + c_hstep * 3)); + float32x4_t _c01 = vcombine_f32(vld1_f32(pC), vld1_f32(pC + c_hstep)); + float32x4_t _c23 = vcombine_f32(vld1_f32(pC + c_hstep * 2), vld1_f32(pC + c_hstep * 3)); _out0 = beta == 1.f ? vaddq_f32(_out0, _c01) : vmlaq_n_f32(_out0, _c01, beta); _out1 = beta == 1.f ? vaddq_f32(_out1, _c23) : vmlaq_n_f32(_out1, _c23, beta); + pC += 2; } if (broadcast_type_C == 4) { float32x2_t _c = vld1_f32(pC); if (beta != 1.f) _c = vmul_n_f32(_c, beta); - const float32x4_t _cc0 = vcombine_f32(_c, _c); + float32x4_t _cc0 = vcombine_f32(_c, _c); _out0 = vaddq_f32(_out0, _cc0); _out1 = vaddq_f32(_out1, _cc0); + pC += 2; } } @@ -4063,8 +4887,6 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b outptr1 += 2; outptr2 += 2; outptr3 += 2; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 2; pp += 8; } for (; jj < max_jj; jj += 1) @@ -4092,11 +4914,13 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _c = vsetq_lane_f32(pC[c_hstep * 2], _c, 2); _c = vsetq_lane_f32(pC[c_hstep * 3], _c, 3); _out0 = beta == 1.f ? vaddq_f32(_out0, _c) : vmlaq_n_f32(_out0, _c, beta); + pC++; } if (broadcast_type_C == 4) { - const float32x4_t _cc0 = vdupq_n_f32(beta == 1.f ? pC[0] : pC[0] * beta); + float32x4_t _cc0 = vdupq_n_f32(beta == 1.f ? pC[0] : pC[0] * beta); _out0 = vaddq_f32(_out0, _cc0); + pC++; } } @@ -4114,18 +4938,19 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b outptr1++; outptr2++; outptr3++; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC++; pp += 4; } } +#endif // __ARM_NEON for (; ii + 1 < max_ii; ii += 2) { pC = (const float*)C; float* outptr0 = (float*)top_blob + (size_t)(i + ii) * out_hstep + j; float* outptr1 = outptr0 + out_hstep; +#if __ARM_NEON float32x4_t _c0; float32x4_t _c1; +#endif float c0 = 0.f; float c1 = 0.f; if (pC) @@ -4135,18 +4960,23 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b c0 = pC[0]; if (beta != 1.f) c0 *= beta; + c1 = c0; +#if __ARM_NEON _c0 = vdupq_n_f32(c0); +#endif } if (broadcast_type_C == 1 || broadcast_type_C == 2) { c0 = pC[i + ii]; if (beta != 1.f) c0 *= beta; - _c0 = vdupq_n_f32(c0); c1 = pC[i + ii + 1]; if (beta != 1.f) c1 *= beta; +#if __ARM_NEON + _c0 = vdupq_n_f32(c0); _c1 = vdupq_n_f32(c1); +#endif } if (broadcast_type_C == 3) pC = (const float*)C + (i + ii) * c_hstep + j; @@ -4155,6 +4985,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b } int jj = 0; +#if __ARM_NEON #if __aarch64__ for (; jj + 7 < max_jj; jj += 8) { @@ -4185,15 +5016,17 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _out1 = vaddq_f32(_out1, (beta == 1.f ? vld1q_f32(pC + 4) : vmulq_n_f32(vld1q_f32(pC + 4), beta))); _out2 = vaddq_f32(_out2, (beta == 1.f ? vld1q_f32(pC + c_hstep) : vmulq_n_f32(vld1q_f32(pC + c_hstep), beta))); _out3 = vaddq_f32(_out3, (beta == 1.f ? vld1q_f32(pC + c_hstep + 4) : vmulq_n_f32(vld1q_f32(pC + c_hstep + 4), beta))); + pC += 8; } if (broadcast_type_C == 4) { - const float32x4_t _c0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); + float32x4_t _c0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); _out0 = vaddq_f32(_out0, _c0); _out2 = vaddq_f32(_out2, _c0); - const float32x4_t _c1 = (beta == 1.f ? vld1q_f32(pC + 4) : vmulq_n_f32(vld1q_f32(pC + 4), beta)); + float32x4_t _c1 = (beta == 1.f ? vld1q_f32(pC + 4) : vmulq_n_f32(vld1q_f32(pC + 4), beta)); _out1 = vaddq_f32(_out1, _c1); _out3 = vaddq_f32(_out3, _c1); + pC += 8; } } @@ -4212,8 +5045,6 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b outptr0 += 8; outptr1 += 8; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 8; pp += 16; } #endif // __aarch64__ @@ -4238,12 +5069,14 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b { _out0 = vaddq_f32(_out0, (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta))); _out1 = vaddq_f32(_out1, (beta == 1.f ? vld1q_f32(pC + c_hstep) : vmulq_n_f32(vld1q_f32(pC + c_hstep), beta))); + pC += 4; } if (broadcast_type_C == 4) { - const float32x4_t _c0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); + float32x4_t _c0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); _out0 = vaddq_f32(_out0, _c0); _out1 = vaddq_f32(_out1, _c0); + pC += 4; } } @@ -4258,8 +5091,6 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b outptr0 += 4; outptr1 += 4; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 4; pp += 8; } for (; jj + 1 < max_jj; jj += 2) @@ -4271,7 +5102,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b { if (broadcast_type_C == 0) { - const float32x2_t _c = vdup_n_f32(c0); + float32x2_t _c = vdup_n_f32(c0); _out0 = vadd_f32(_out0, _c); _out1 = vadd_f32(_out1, _c); } @@ -4282,16 +5113,18 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b } if (broadcast_type_C == 3) { - const float32x2_t _c0 = vld1_f32(pC); - const float32x2_t _c1 = vld1_f32(pC + c_hstep); + float32x2_t _c0 = vld1_f32(pC); + float32x2_t _c1 = vld1_f32(pC + c_hstep); _out0 = beta == 1.f ? vadd_f32(_out0, _c0) : vmla_n_f32(_out0, _c0, beta); _out1 = beta == 1.f ? vadd_f32(_out1, _c1) : vmla_n_f32(_out1, _c1, beta); + pC += 2; } if (broadcast_type_C == 4) { - const float32x2_t _c0 = beta == 1.f ? vld1_f32(pC) : vmul_n_f32(vld1_f32(pC), beta); + float32x2_t _c0 = beta == 1.f ? vld1_f32(pC) : vmul_n_f32(vld1_f32(pC), beta); _out0 = vadd_f32(_out0, _c0); _out1 = vadd_f32(_out1, _c0); + pC += 2; } } @@ -4306,8 +5139,6 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b outptr0 += 2; outptr1 += 2; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 2; pp += 4; } for (; jj < max_jj; jj += 1) @@ -4331,10 +5162,12 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b float32x2_t _c = vdup_n_f32(pC[0]); _c = vset_lane_f32(pC[c_hstep], _c, 1); _out0 = beta == 1.f ? vadd_f32(_out0, _c) : vmla_n_f32(_out0, _c, beta); + pC++; } if (broadcast_type_C == 4) { _out0 = vadd_f32(_out0, vdup_n_f32(beta == 1.f ? pC[0] : pC[0] * beta)); + pC++; } } @@ -4348,8 +5181,109 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b outptr0++; outptr1++; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC++; + pp += 2; + } +#endif // __ARM_NEON + for (; jj + 1 < max_jj; jj += 2) + { + float out00 = pp[0]; + float out01 = pp[1]; + float out10 = pp[2]; + float out11 = pp[3]; + + if (pC) + { + if (broadcast_type_C == 0) + { + out00 += c0; + out01 += c0; + out10 += c0; + out11 += c0; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + out00 += c0; + out01 += c0; + out10 += c1; + out11 += c1; + } + if (broadcast_type_C == 3) + { + out00 += beta == 1.f ? pC[0] : pC[0] * beta; + out01 += beta == 1.f ? pC[1] : pC[1] * beta; + out10 += beta == 1.f ? pC[c_hstep] : pC[c_hstep] * beta; + out11 += beta == 1.f ? pC[c_hstep + 1] : pC[c_hstep + 1] * beta; + pC += 2; + } + if (broadcast_type_C == 4) + { + out00 += beta == 1.f ? pC[0] : pC[0] * beta; + out01 += beta == 1.f ? pC[1] : pC[1] * beta; + out10 += beta == 1.f ? pC[0] : pC[0] * beta; + out11 += beta == 1.f ? pC[1] : pC[1] * beta; + pC += 2; + } + } + + if (alpha != 1.f) + { + out00 *= alpha; + out01 *= alpha; + out10 *= alpha; + out11 *= alpha; + } + + outptr0[0] = out00; + outptr0[1] = out01; + outptr1[0] = out10; + outptr1[1] = out11; + + outptr0 += 2; + outptr1 += 2; + pp += 4; + } + for (; jj < max_jj; jj += 1) + { + float out00 = pp[0]; + float out10 = pp[1]; + + if (pC) + { + if (broadcast_type_C == 0) + { + out00 += c0; + out10 += c0; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + out00 += c0; + out10 += c1; + } + if (broadcast_type_C == 3) + { + out00 += beta == 1.f ? pC[0] : pC[0] * beta; + out10 += beta == 1.f ? pC[c_hstep] : pC[c_hstep] * beta; + pC++; + } + if (broadcast_type_C == 4) + { + out00 += beta == 1.f ? pC[0] : pC[0] * beta; + out10 += beta == 1.f ? pC[0] : pC[0] * beta; + pC++; + } + } + + if (alpha != 1.f) + { + out00 *= alpha; + out10 *= alpha; + } + + outptr0[0] = out00; + outptr1[0] = out10; + + outptr0++; + outptr1++; pp += 2; } } @@ -4357,7 +5291,9 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b { pC = (const float*)C; float* outptr0 = (float*)top_blob + (size_t)(i + ii) * out_hstep + j; +#if __ARM_NEON float32x4_t _c0; +#endif float c0 = 0.f; if (pC) { @@ -4366,14 +5302,18 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b c0 = pC[0]; if (beta != 1.f) c0 *= beta; +#if __ARM_NEON _c0 = vdupq_n_f32(c0); +#endif } if (broadcast_type_C == 1 || broadcast_type_C == 2) { c0 = pC[i + ii]; if (beta != 1.f) c0 *= beta; +#if __ARM_NEON _c0 = vdupq_n_f32(c0); +#endif } if (broadcast_type_C == 3) pC = (const float*)C + (i + ii) * c_hstep + j; @@ -4382,7 +5322,54 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b } int jj = 0; +#if __ARM_NEON #if __aarch64__ + for (; jj + 15 < max_jj; jj += 16) + { + float32x4_t _out0 = vld1q_f32(pp); + float32x4_t _out1 = vld1q_f32(pp + 4); + float32x4_t _out2 = vld1q_f32(pp + 8); + float32x4_t _out3 = vld1q_f32(pp + 12); + + if (pC) + { + if (broadcast_type_C <= 2) + { + _out0 = vaddq_f32(_out0, _c0); + _out1 = vaddq_f32(_out1, _c0); + _out2 = vaddq_f32(_out2, _c0); + _out3 = vaddq_f32(_out3, _c0); + } + if (broadcast_type_C == 3 || broadcast_type_C == 4) + { + float32x4_t _c0 = beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta); + float32x4_t _c1 = beta == 1.f ? vld1q_f32(pC + 4) : vmulq_n_f32(vld1q_f32(pC + 4), beta); + float32x4_t _c2 = beta == 1.f ? vld1q_f32(pC + 8) : vmulq_n_f32(vld1q_f32(pC + 8), beta); + float32x4_t _c3 = beta == 1.f ? vld1q_f32(pC + 12) : vmulq_n_f32(vld1q_f32(pC + 12), beta); + _out0 = vaddq_f32(_out0, _c0); + _out1 = vaddq_f32(_out1, _c1); + _out2 = vaddq_f32(_out2, _c2); + _out3 = vaddq_f32(_out3, _c3); + pC += 16; + } + } + + if (alpha != 1.f) + { + _out0 = vmulq_n_f32(_out0, alpha); + _out1 = vmulq_n_f32(_out1, alpha); + _out2 = vmulq_n_f32(_out2, alpha); + _out3 = vmulq_n_f32(_out3, alpha); + } + + vst1q_f32(outptr0, _out0); + vst1q_f32(outptr0 + 4, _out1); + vst1q_f32(outptr0 + 8, _out2); + vst1q_f32(outptr0 + 12, _out3); + + outptr0 += 16; + pp += 16; + } for (; jj + 7 < max_jj; jj += 8) { float32x4_t _out0 = vld1q_f32(pp + 0); @@ -4404,13 +5391,15 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b { _out0 = vaddq_f32(_out0, (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta))); _out1 = vaddq_f32(_out1, (beta == 1.f ? vld1q_f32(pC + 4) : vmulq_n_f32(vld1q_f32(pC + 4), beta))); + pC += 8; } if (broadcast_type_C == 4) { - const float32x4_t _c0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); + float32x4_t _c0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); _out0 = vaddq_f32(_out0, _c0); - const float32x4_t _c1 = (beta == 1.f ? vld1q_f32(pC + 4) : vmulq_n_f32(vld1q_f32(pC + 4), beta)); + float32x4_t _c1 = (beta == 1.f ? vld1q_f32(pC + 4) : vmulq_n_f32(vld1q_f32(pC + 4), beta)); _out1 = vaddq_f32(_out1, _c1); + pC += 8; } } @@ -4424,8 +5413,6 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b vst1q_f32(outptr0 + 4, _out1); outptr0 += 8; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 8; pp += 8; } #endif // __aarch64__ @@ -4446,11 +5433,13 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b if (broadcast_type_C == 3) { _out0 = vaddq_f32(_out0, (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta))); + pC += 4; } if (broadcast_type_C == 4) { - const float32x4_t _c0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); + float32x4_t _c0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); _out0 = vaddq_f32(_out0, _c0); + pC += 4; } } @@ -4462,8 +5451,6 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b vst1q_f32(outptr0, _out0); outptr0 += 4; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 4; pp += 4; } for (; jj + 1 < max_jj; jj += 2) @@ -4482,13 +5469,15 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b } if (broadcast_type_C == 3) { - const float32x2_t _c0 = vld1_f32(pC); + float32x2_t _c0 = vld1_f32(pC); _out0 = beta == 1.f ? vadd_f32(_out0, _c0) : vmla_n_f32(_out0, _c0, beta); + pC += 2; } if (broadcast_type_C == 4) { - const float32x2_t _c0 = beta == 1.f ? vld1_f32(pC) : vmul_n_f32(vld1_f32(pC), beta); + float32x2_t _c0 = beta == 1.f ? vld1_f32(pC) : vmul_n_f32(vld1_f32(pC), beta); _out0 = vadd_f32(_out0, _c0); + pC += 2; } } @@ -4500,292 +5489,89 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b vst1_f32(outptr0, _out0); outptr0 += 2; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 2; + pp += 2; + } +#endif // __ARM_NEON + for (; jj + 1 < max_jj; jj += 2) + { + float out00 = pp[0]; + float out01 = pp[1]; + + if (pC) + { + if (broadcast_type_C == 0) + { + out00 += c0; + out01 += c0; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + out00 += c0; + out01 += c0; + } + if (broadcast_type_C == 3) + { + out00 += beta == 1.f ? pC[0] : pC[0] * beta; + out01 += beta == 1.f ? pC[1] : pC[1] * beta; + pC += 2; + } + if (broadcast_type_C == 4) + { + out00 += beta == 1.f ? pC[0] : pC[0] * beta; + out01 += beta == 1.f ? pC[1] : pC[1] * beta; + pC += 2; + } + } + + if (alpha != 1.f) + { + out00 *= alpha; + out01 *= alpha; + } + + outptr0[0] = out00; + outptr0[1] = out01; + + outptr0 += 2; pp += 2; } for (; jj < max_jj; jj += 1) { - float out0 = pp[0]; + float out00 = pp[0]; if (pC) { if (broadcast_type_C == 0) { - out0 += c0; + out00 += c0; } if (broadcast_type_C == 1 || broadcast_type_C == 2) { - out0 += c0; + out00 += c0; } if (broadcast_type_C == 3) { - out0 += beta == 1.f ? pC[0] : pC[0] * beta; + out00 += beta == 1.f ? pC[0] : pC[0] * beta; + pC++; } if (broadcast_type_C == 4) { - out0 += beta == 1.f ? pC[0] : pC[0] * beta; + out00 += beta == 1.f ? pC[0] : pC[0] * beta; + pC++; } } if (alpha != 1.f) { - out0 *= alpha; + out00 *= alpha; } - outptr0[0] = out0; + outptr0[0] = out00; outptr0++; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC++; pp += 1; } } -#else - for (; ii + 1 < max_ii; ii += 2) - { - pC = (const float*)C; - float* outptr0 = (float*)top_blob + (size_t)(i + ii) * out_hstep + j; - float* outptr1 = outptr0 + out_hstep; - float c0 = 0.f; - float c1 = 0.f; - if (pC) - { - if (broadcast_type_C == 0) - { - float c = pC[0]; - if (beta != 1.f) - c *= beta; - c0 = c; - c1 = c; - } - if (broadcast_type_C == 1 || broadcast_type_C == 2) - { - c0 = pC[i + ii]; - if (beta != 1.f) - c0 *= beta; - c1 = pC[i + ii + 1]; - if (beta != 1.f) - c1 *= beta; - } - if (broadcast_type_C == 3) - pC = (const float*)C + (i + ii) * c_hstep + j; - if (broadcast_type_C == 4) - pC = (const float*)C + j; - } - - int jj = 0; - for (; jj + 1 < max_jj; jj += 2) - { - float out00 = pp[0]; - float out01 = pp[1]; - float out10 = pp[2]; - float out11 = pp[3]; - - if (pC) - { - if (broadcast_type_C == 0) - { - out00 += c0; - out01 += c0; - out10 += c0; - out11 += c0; - } - if (broadcast_type_C == 1 || broadcast_type_C == 2) - { - out00 += c0; - out01 += c0; - out10 += c1; - out11 += c1; - } - if (broadcast_type_C == 3) - { - out00 += beta == 1.f ? pC[0] : pC[0] * beta; - out01 += beta == 1.f ? pC[1] : pC[1] * beta; - out10 += beta == 1.f ? pC[c_hstep] : pC[c_hstep] * beta; - out11 += beta == 1.f ? pC[c_hstep + 1] : pC[c_hstep + 1] * beta; - } - if (broadcast_type_C == 4) - { - out00 += beta == 1.f ? pC[0] : pC[0] * beta; - out01 += beta == 1.f ? pC[1] : pC[1] * beta; - out10 += beta == 1.f ? pC[0] : pC[0] * beta; - out11 += beta == 1.f ? pC[1] : pC[1] * beta; - } - } - - if (alpha != 1.f) - { - out00 *= alpha; - out01 *= alpha; - out10 *= alpha; - out11 *= alpha; - } - - outptr0[0] = out00; - outptr0[1] = out01; - outptr1[0] = out10; - outptr1[1] = out11; - - outptr0 += 2; - outptr1 += 2; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 2; - pp += 4; - } - for (; jj < max_jj; jj += 1) - { - float out00 = pp[0]; - float out10 = pp[1]; - - if (pC) - { - if (broadcast_type_C == 0) - { - out00 += c0; - out10 += c0; - } - if (broadcast_type_C == 1 || broadcast_type_C == 2) - { - out00 += c0; - out10 += c1; - } - if (broadcast_type_C == 3) - { - out00 += beta == 1.f ? pC[0] : pC[0] * beta; - out10 += beta == 1.f ? pC[c_hstep] : pC[c_hstep] * beta; - } - if (broadcast_type_C == 4) - { - out00 += beta == 1.f ? pC[0] : pC[0] * beta; - out10 += beta == 1.f ? pC[0] : pC[0] * beta; - } - } - - if (alpha != 1.f) - { - out00 *= alpha; - out10 *= alpha; - } - - outptr0[0] = out00; - outptr1[0] = out10; - - outptr0++; - outptr1++; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC++; - pp += 2; - } - } - for (; ii < max_ii; ii++) - { - pC = (const float*)C; - float* outptr0 = (float*)top_blob + (size_t)(i + ii) * out_hstep + j; - float c0 = 0.f; - if (pC) - { - if (broadcast_type_C == 0) - { - float c = pC[0]; - if (beta != 1.f) - c *= beta; - c0 = c; - } - if (broadcast_type_C == 1 || broadcast_type_C == 2) - { - c0 = pC[i + ii]; - if (beta != 1.f) - c0 *= beta; - } - if (broadcast_type_C == 3) - pC = (const float*)C + (i + ii) * c_hstep + j; - if (broadcast_type_C == 4) - pC = (const float*)C + j; - } - - int jj = 0; - for (; jj + 1 < max_jj; jj += 2) - { - float out00 = pp[0]; - float out01 = pp[1]; - - if (pC) - { - if (broadcast_type_C == 0) - { - out00 += c0; - out01 += c0; - } - if (broadcast_type_C == 1 || broadcast_type_C == 2) - { - out00 += c0; - out01 += c0; - } - if (broadcast_type_C == 3) - { - out00 += beta == 1.f ? pC[0] : pC[0] * beta; - out01 += beta == 1.f ? pC[1] : pC[1] * beta; - } - if (broadcast_type_C == 4) - { - out00 += beta == 1.f ? pC[0] : pC[0] * beta; - out01 += beta == 1.f ? pC[1] : pC[1] * beta; - } - } - - if (alpha != 1.f) - { - out00 *= alpha; - out01 *= alpha; - } - - outptr0[0] = out00; - outptr0[1] = out01; - - outptr0 += 2; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 2; - pp += 2; - } - for (; jj < max_jj; jj += 1) - { - float out00 = pp[0]; - - if (pC) - { - if (broadcast_type_C == 0) - { - out00 += c0; - } - if (broadcast_type_C == 1 || broadcast_type_C == 2) - { - out00 += c0; - } - if (broadcast_type_C == 3) - { - out00 += beta == 1.f ? pC[0] : pC[0] * beta; - } - if (broadcast_type_C == 4) - { - out00 += beta == 1.f ? pC[0] : pC[0] * beta; - } - } - - if (alpha != 1.f) - { - out00 *= alpha; - } - - outptr0[0] = out00; - - outptr0++; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC++; - pp += 1; - } - } -#endif // __ARM_NEON } static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta) @@ -4875,7 +5661,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, } if (broadcast_type_C == 4) { - const float32x4_t _c = beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta); + float32x4_t _c = beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta); _out0 = vaddq_f32(_out0, _c); _out1 = vaddq_f32(_out1, _c); _out2 = vaddq_f32(_out2, _c); @@ -4951,7 +5737,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, } if (broadcast_type_C == 4) { - const float32x2_t _c = beta == 1.f ? vld1_f32(pC) : vmul_n_f32(vld1_f32(pC), beta); + float32x2_t _c = beta == 1.f ? vld1_f32(pC) : vmul_n_f32(vld1_f32(pC), beta); _out0 = vadd_f32(_out0, _c); _out1 = vadd_f32(_out1, _c); _out2 = vadd_f32(_out2, _c); @@ -4992,8 +5778,8 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, { if (broadcast_type_C <= 2) { - const float32x4_t _c0 = {c0, c1, c2, c3}; - const float32x4_t _c1 = {c4, c5, c6, c7}; + float32x4_t _c0 = {c0, c1, c2, c3}; + float32x4_t _c1 = {c4, c5, c6, c7}; _out0 = vaddq_f32(_out0, _c0); _out1 = vaddq_f32(_out1, _c1); } @@ -5007,7 +5793,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, } if (broadcast_type_C == 4) { - const float32x4_t _c = vdupq_n_f32(beta == 1.f ? pC[0] : pC[0] * beta); + float32x4_t _c = vdupq_n_f32(beta == 1.f ? pC[0] : pC[0] * beta); _out0 = vaddq_f32(_out0, _c); _out1 = vaddq_f32(_out1, _c); pC++; @@ -5076,19 +5862,21 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _out5 = vaddq_f32(_out5, (beta == 1.f ? vld1q_f32(pC + c_hstep * 2 + 4) : vmulq_n_f32(vld1q_f32(pC + c_hstep * 2 + 4), beta))); _out6 = vaddq_f32(_out6, (beta == 1.f ? vld1q_f32(pC + c_hstep * 3) : vmulq_n_f32(vld1q_f32(pC + c_hstep * 3), beta))); _out7 = vaddq_f32(_out7, (beta == 1.f ? vld1q_f32(pC + c_hstep * 3 + 4) : vmulq_n_f32(vld1q_f32(pC + c_hstep * 3 + 4), beta))); + pC += 8; } if (broadcast_type_C == 4) { - const float32x4_t _cc0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); + float32x4_t _cc0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); _out0 = vaddq_f32(_out0, _cc0); _out2 = vaddq_f32(_out2, _cc0); _out4 = vaddq_f32(_out4, _cc0); _out6 = vaddq_f32(_out6, _cc0); - const float32x4_t _cc1 = (beta == 1.f ? vld1q_f32(pC + 4) : vmulq_n_f32(vld1q_f32(pC + 4), beta)); + float32x4_t _cc1 = (beta == 1.f ? vld1q_f32(pC + 4) : vmulq_n_f32(vld1q_f32(pC + 4), beta)); _out1 = vaddq_f32(_out1, _cc1); _out3 = vaddq_f32(_out3, _cc1); _out5 = vaddq_f32(_out5, _cc1); _out7 = vaddq_f32(_out7, _cc1); + pC += 8; } } @@ -5109,10 +5897,10 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, } if (broadcast_type_C == 1 || broadcast_type_C == 2) { - const float32x4_t _c0 = vdupq_lane_f32(vget_low_f32(_c0123), 0); - const float32x4_t _c1 = vdupq_lane_f32(vget_low_f32(_c0123), 1); - const float32x4_t _c2 = vdupq_lane_f32(vget_high_f32(_c0123), 0); - const float32x4_t _c3 = vdupq_lane_f32(vget_high_f32(_c0123), 1); + float32x4_t _c0 = vdupq_lane_f32(vget_low_f32(_c0123), 0); + float32x4_t _c1 = vdupq_lane_f32(vget_low_f32(_c0123), 1); + float32x4_t _c2 = vdupq_lane_f32(vget_high_f32(_c0123), 0); + float32x4_t _c3 = vdupq_lane_f32(vget_high_f32(_c0123), 1); _out0 = vaddq_f32(_out0, _c0); _out1 = vaddq_f32(_out1, _c0); _out2 = vaddq_f32(_out2, _c1); @@ -5189,8 +5977,6 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, } outptr0 += out_hstep * 8; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 8; pp += 32; } #endif // __aarch64__ @@ -5209,14 +5995,16 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _out1 = vaddq_f32(_out1, (beta == 1.f ? vld1q_f32(pC + c_hstep) : vmulq_n_f32(vld1q_f32(pC + c_hstep), beta))); _out2 = vaddq_f32(_out2, (beta == 1.f ? vld1q_f32(pC + c_hstep * 2) : vmulq_n_f32(vld1q_f32(pC + c_hstep * 2), beta))); _out3 = vaddq_f32(_out3, (beta == 1.f ? vld1q_f32(pC + c_hstep * 3) : vmulq_n_f32(vld1q_f32(pC + c_hstep * 3), beta))); + pC += 4; } if (broadcast_type_C == 4) { - const float32x4_t _cc0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); + float32x4_t _cc0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); _out0 = vaddq_f32(_out0, _cc0); _out1 = vaddq_f32(_out1, _cc0); _out2 = vaddq_f32(_out2, _cc0); _out3 = vaddq_f32(_out3, _cc0); + pC += 4; } } @@ -5282,8 +6070,6 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, } outptr0 += out_hstep * 4; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 4; pp += 16; } for (; jj + 1 < max_jj; jj += 2) @@ -5295,19 +6081,21 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, { if (broadcast_type_C == 3) { - const float32x4_t _c01 = vcombine_f32(vld1_f32(pC), vld1_f32(pC + c_hstep)); - const float32x4_t _c23 = vcombine_f32(vld1_f32(pC + c_hstep * 2), vld1_f32(pC + c_hstep * 3)); + float32x4_t _c01 = vcombine_f32(vld1_f32(pC), vld1_f32(pC + c_hstep)); + float32x4_t _c23 = vcombine_f32(vld1_f32(pC + c_hstep * 2), vld1_f32(pC + c_hstep * 3)); _out0 = beta == 1.f ? vaddq_f32(_out0, _c01) : vmlaq_n_f32(_out0, _c01, beta); _out1 = beta == 1.f ? vaddq_f32(_out1, _c23) : vmlaq_n_f32(_out1, _c23, beta); + pC += 2; } if (broadcast_type_C == 4) { float32x2_t _c = vld1_f32(pC); if (beta != 1.f) _c = vmul_n_f32(_c, beta); - const float32x4_t _cc0 = vcombine_f32(_c, _c); + float32x4_t _cc0 = vcombine_f32(_c, _c); _out0 = vaddq_f32(_out0, _cc0); _out1 = vaddq_f32(_out1, _cc0); + pC += 2; } } @@ -5362,8 +6150,6 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, } outptr0 += out_hstep * 2; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 2; pp += 8; } for (; jj < max_jj; jj += 1) @@ -5387,11 +6173,13 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _c = vsetq_lane_f32(pC[c_hstep * 2], _c, 2); _c = vsetq_lane_f32(pC[c_hstep * 3], _c, 3); _out0 = beta == 1.f ? vaddq_f32(_out0, _c) : vmlaq_n_f32(_out0, _c, beta); + pC++; } if (broadcast_type_C == 4) { - const float32x4_t _cc0 = vdupq_n_f32(beta == 1.f ? pC[0] : pC[0] * beta); + float32x4_t _cc0 = vdupq_n_f32(beta == 1.f ? pC[0] : pC[0] * beta); _out0 = vaddq_f32(_out0, _cc0); + pC++; } } @@ -5403,16 +6191,19 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, vst1q_f32(outptr0, _out0); outptr0 += out_hstep; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC++; pp += 4; } } +#endif // __ARM_NEON for (; ii + 1 < max_ii; ii += 2) { pC = (const float*)C; float* outptr0 = (float*)top_blob + (size_t)j * out_hstep + i + ii; +#if __ARM_NEON float32x4_t _c01 = vdupq_n_f32(0.f); +#endif + float c0 = 0.f; + float c1 = 0.f; if (pC) { if (broadcast_type_C == 0) @@ -5420,15 +6211,27 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, float c = pC[0]; if (beta != 1.f) c *= beta; + c0 = c; + c1 = c; +#if __ARM_NEON _c01 = vdupq_n_f32(c); +#endif } if (broadcast_type_C == 1 || broadcast_type_C == 2) { - float32x2_t _c = vld1_f32(pC + i + ii); + c0 = pC[i + ii]; if (beta != 1.f) - _c = vmul_n_f32(_c, beta); - _c01 = vcombine_f32(_c, _c); - } + c0 *= beta; + c1 = pC[i + ii + 1]; + if (beta != 1.f) + c1 *= beta; +#if __ARM_NEON + float32x2_t _c = vld1_f32(pC + i + ii); + if (beta != 1.f) + _c = vmul_n_f32(_c, beta); + _c01 = vcombine_f32(_c, _c); +#endif + } if (broadcast_type_C == 3) pC = (const float*)C + (i + ii) * c_hstep + j; if (broadcast_type_C == 4) @@ -5436,6 +6239,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, } int jj = 0; +#if __ARM_NEON #if __aarch64__ for (; jj + 7 < max_jj; jj += 8) { @@ -5452,15 +6256,17 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _out1 = vaddq_f32(_out1, (beta == 1.f ? vld1q_f32(pC + 4) : vmulq_n_f32(vld1q_f32(pC + 4), beta))); _out2 = vaddq_f32(_out2, (beta == 1.f ? vld1q_f32(pC + c_hstep) : vmulq_n_f32(vld1q_f32(pC + c_hstep), beta))); _out3 = vaddq_f32(_out3, (beta == 1.f ? vld1q_f32(pC + c_hstep + 4) : vmulq_n_f32(vld1q_f32(pC + c_hstep + 4), beta))); + pC += 8; } if (broadcast_type_C == 4) { - const float32x4_t _c0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); + float32x4_t _c0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); _out0 = vaddq_f32(_out0, _c0); _out2 = vaddq_f32(_out2, _c0); - const float32x4_t _c1 = (beta == 1.f ? vld1q_f32(pC + 4) : vmulq_n_f32(vld1q_f32(pC + 4), beta)); + float32x4_t _c1 = (beta == 1.f ? vld1q_f32(pC + 4) : vmulq_n_f32(vld1q_f32(pC + 4), beta)); _out1 = vaddq_f32(_out1, _c1); _out3 = vaddq_f32(_out3, _c1); + pC += 8; } } @@ -5477,8 +6283,8 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, } if (broadcast_type_C == 1 || broadcast_type_C == 2) { - const float32x4_t _c0 = vdupq_lane_f32(vget_low_f32(_c01), 0); - const float32x4_t _c1 = vdupq_lane_f32(vget_low_f32(_c01), 1); + float32x4_t _c0 = vdupq_lane_f32(vget_low_f32(_c01), 0); + float32x4_t _c1 = vdupq_lane_f32(vget_low_f32(_c01), 1); _out0 = vaddq_f32(_out0, _c0); _out1 = vaddq_f32(_out1, _c0); _out2 = vaddq_f32(_out2, _c1); @@ -5568,8 +6374,6 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, } outptr0 += out_hstep * 8; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 8; pp += 16; } #endif // __aarch64__ @@ -5584,12 +6388,14 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, { _out0 = vaddq_f32(_out0, (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta))); _out1 = vaddq_f32(_out1, (beta == 1.f ? vld1q_f32(pC + c_hstep) : vmulq_n_f32(vld1q_f32(pC + c_hstep), beta))); + pC += 4; } if (broadcast_type_C == 4) { - const float32x4_t _c0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); + float32x4_t _c0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); _out0 = vaddq_f32(_out0, _c0); _out1 = vaddq_f32(_out1, _c0); + pC += 4; } } @@ -5664,8 +6470,6 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, } outptr0 += out_hstep * 4; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 4; pp += 8; } for (; jj + 1 < max_jj; jj += 2) @@ -5677,16 +6481,18 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, { if (broadcast_type_C == 3) { - const float32x2_t _c0 = vld1_f32(pC); - const float32x2_t _c1 = vld1_f32(pC + c_hstep); + float32x2_t _c0 = vld1_f32(pC); + float32x2_t _c1 = vld1_f32(pC + c_hstep); _out0 = beta == 1.f ? vadd_f32(_out0, _c0) : vmla_n_f32(_out0, _c0, beta); _out1 = beta == 1.f ? vadd_f32(_out1, _c1) : vmla_n_f32(_out1, _c1, beta); + pC += 2; } if (broadcast_type_C == 4) { - const float32x2_t _c0 = beta == 1.f ? vld1_f32(pC) : vmul_n_f32(vld1_f32(pC), beta); + float32x2_t _c0 = beta == 1.f ? vld1_f32(pC) : vmul_n_f32(vld1_f32(pC), beta); _out0 = vadd_f32(_out0, _c0); _out1 = vadd_f32(_out1, _c0); + pC += 2; } } @@ -5743,8 +6549,6 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, } outptr0 += out_hstep * 2; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 2; pp += 4; } for (; jj < max_jj; jj += 1) @@ -5766,10 +6570,12 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, float32x2_t _c = vdup_n_f32(pC[0]); _c = vset_lane_f32(pC[c_hstep], _c, 1); _out0 = beta == 1.f ? vadd_f32(_out0, _c) : vmla_n_f32(_out0, _c, beta); + pC++; } if (broadcast_type_C == 4) { _out0 = vadd_f32(_out0, vdup_n_f32(beta == 1.f ? pC[0] : pC[0] * beta)); + pC++; } } @@ -5781,8 +6587,107 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, vst1_f32(outptr0, _out0); outptr0 += out_hstep; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC++; + pp += 2; + } +#endif // __ARM_NEON + for (; jj + 1 < max_jj; jj += 2) + { + float out00 = pp[0]; + float out01 = pp[1]; + float out10 = pp[2]; + float out11 = pp[3]; + + if (pC) + { + if (broadcast_type_C == 0) + { + out00 += c0; + out01 += c0; + out10 += c0; + out11 += c0; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + out00 += c0; + out01 += c0; + out10 += c1; + out11 += c1; + } + if (broadcast_type_C == 3) + { + out00 += beta == 1.f ? pC[0] : pC[0] * beta; + out01 += beta == 1.f ? pC[1] : pC[1] * beta; + out10 += beta == 1.f ? pC[c_hstep] : pC[c_hstep] * beta; + out11 += beta == 1.f ? pC[c_hstep + 1] : pC[c_hstep + 1] * beta; + pC += 2; + } + if (broadcast_type_C == 4) + { + out00 += beta == 1.f ? pC[0] : pC[0] * beta; + out01 += beta == 1.f ? pC[1] : pC[1] * beta; + out10 += beta == 1.f ? pC[0] : pC[0] * beta; + out11 += beta == 1.f ? pC[1] : pC[1] * beta; + pC += 2; + } + } + + if (alpha != 1.f) + { + out00 *= alpha; + out01 *= alpha; + out10 *= alpha; + out11 *= alpha; + } + + outptr0[0] = out00; + outptr0[out_hstep] = out01; + outptr0[1] = out10; + outptr0[out_hstep + 1] = out11; + + outptr0 += out_hstep * 2; + pp += 4; + } + for (; jj < max_jj; jj += 1) + { + float out00 = pp[0]; + float out10 = pp[1]; + + if (pC) + { + if (broadcast_type_C == 0) + { + out00 += c0; + out10 += c0; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + out00 += c0; + out10 += c1; + } + if (broadcast_type_C == 3) + { + out00 += beta == 1.f ? pC[0] : pC[0] * beta; + out10 += beta == 1.f ? pC[c_hstep] : pC[c_hstep] * beta; + pC++; + } + if (broadcast_type_C == 4) + { + out00 += beta == 1.f ? pC[0] : pC[0] * beta; + out10 += beta == 1.f ? pC[0] : pC[0] * beta; + pC++; + } + } + + if (alpha != 1.f) + { + out00 *= alpha; + out10 *= alpha; + } + + outptr0[0] = out00; + outptr0[1] = out10; + + outptr0 += out_hstep; pp += 2; } } @@ -5790,7 +6695,9 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, { pC = (const float*)C; float* outptr0 = (float*)top_blob + (size_t)j * out_hstep + i + ii; +#if __ARM_NEON float32x4_t _c0; +#endif float c0 = 0.f; if (pC) { @@ -5799,14 +6706,18 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, c0 = pC[0]; if (beta != 1.f) c0 *= beta; +#if __ARM_NEON _c0 = vdupq_n_f32(c0); +#endif } if (broadcast_type_C == 1 || broadcast_type_C == 2) { c0 = pC[i + ii]; if (beta != 1.f) c0 *= beta; +#if __ARM_NEON _c0 = vdupq_n_f32(c0); +#endif } if (broadcast_type_C == 3) pC = (const float*)C + (i + ii) * c_hstep + j; @@ -5815,7 +6726,76 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, } int jj = 0; +#if __ARM_NEON #if __aarch64__ + for (; jj + 15 < max_jj; jj += 16) + { + float32x4_t _out0 = vld1q_f32(pp); + float32x4_t _out1 = vld1q_f32(pp + 4); + float32x4_t _out2 = vld1q_f32(pp + 8); + float32x4_t _out3 = vld1q_f32(pp + 12); + + if (pC) + { + if (broadcast_type_C <= 2) + { + _out0 = vaddq_f32(_out0, _c0); + _out1 = vaddq_f32(_out1, _c0); + _out2 = vaddq_f32(_out2, _c0); + _out3 = vaddq_f32(_out3, _c0); + } + if (broadcast_type_C == 3 || broadcast_type_C == 4) + { + float32x4_t _c0 = beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta); + float32x4_t _c1 = beta == 1.f ? vld1q_f32(pC + 4) : vmulq_n_f32(vld1q_f32(pC + 4), beta); + float32x4_t _c2 = beta == 1.f ? vld1q_f32(pC + 8) : vmulq_n_f32(vld1q_f32(pC + 8), beta); + float32x4_t _c3 = beta == 1.f ? vld1q_f32(pC + 12) : vmulq_n_f32(vld1q_f32(pC + 12), beta); + _out0 = vaddq_f32(_out0, _c0); + _out1 = vaddq_f32(_out1, _c1); + _out2 = vaddq_f32(_out2, _c2); + _out3 = vaddq_f32(_out3, _c3); + pC += 16; + } + } + + if (alpha != 1.f) + { + _out0 = vmulq_n_f32(_out0, alpha); + _out1 = vmulq_n_f32(_out1, alpha); + _out2 = vmulq_n_f32(_out2, alpha); + _out3 = vmulq_n_f32(_out3, alpha); + } + + if (out_hstep == 1) + { + vst1q_f32(outptr0, _out0); + vst1q_f32(outptr0 + 4, _out1); + vst1q_f32(outptr0 + 8, _out2); + vst1q_f32(outptr0 + 12, _out3); + } + else + { + vst1q_lane_f32(outptr0, _out0, 0); + vst1q_lane_f32(outptr0 + out_hstep, _out0, 1); + vst1q_lane_f32(outptr0 + out_hstep * 2, _out0, 2); + vst1q_lane_f32(outptr0 + out_hstep * 3, _out0, 3); + vst1q_lane_f32(outptr0 + out_hstep * 4, _out1, 0); + vst1q_lane_f32(outptr0 + out_hstep * 5, _out1, 1); + vst1q_lane_f32(outptr0 + out_hstep * 6, _out1, 2); + vst1q_lane_f32(outptr0 + out_hstep * 7, _out1, 3); + vst1q_lane_f32(outptr0 + out_hstep * 8, _out2, 0); + vst1q_lane_f32(outptr0 + out_hstep * 9, _out2, 1); + vst1q_lane_f32(outptr0 + out_hstep * 10, _out2, 2); + vst1q_lane_f32(outptr0 + out_hstep * 11, _out2, 3); + vst1q_lane_f32(outptr0 + out_hstep * 12, _out3, 0); + vst1q_lane_f32(outptr0 + out_hstep * 13, _out3, 1); + vst1q_lane_f32(outptr0 + out_hstep * 14, _out3, 2); + vst1q_lane_f32(outptr0 + out_hstep * 15, _out3, 3); + } + + outptr0 += out_hstep * 16; + pp += 16; + } for (; jj + 7 < max_jj; jj += 8) { float32x4_t _out0 = vld1q_f32(pp + 0); @@ -5837,13 +6817,15 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, { _out0 = vaddq_f32(_out0, (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta))); _out1 = vaddq_f32(_out1, (beta == 1.f ? vld1q_f32(pC + 4) : vmulq_n_f32(vld1q_f32(pC + 4), beta))); + pC += 8; } if (broadcast_type_C == 4) { - const float32x4_t _c0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); + float32x4_t _c0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); _out0 = vaddq_f32(_out0, _c0); - const float32x4_t _c1 = (beta == 1.f ? vld1q_f32(pC + 4) : vmulq_n_f32(vld1q_f32(pC + 4), beta)); + float32x4_t _c1 = (beta == 1.f ? vld1q_f32(pC + 4) : vmulq_n_f32(vld1q_f32(pC + 4), beta)); _out1 = vaddq_f32(_out1, _c1); + pC += 8; } } @@ -5871,8 +6853,6 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, } outptr0 += out_hstep * 8; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 8; pp += 8; } #endif // __aarch64__ @@ -5893,11 +6873,13 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, if (broadcast_type_C == 3) { _out0 = vaddq_f32(_out0, (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta))); + pC += 4; } if (broadcast_type_C == 4) { - const float32x4_t _c0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); + float32x4_t _c0 = (beta == 1.f ? vld1q_f32(pC) : vmulq_n_f32(vld1q_f32(pC), beta)); _out0 = vaddq_f32(_out0, _c0); + pC += 4; } } @@ -5919,8 +6901,6 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, } outptr0 += out_hstep * 4; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 4; pp += 4; } for (; jj + 1 < max_jj; jj += 2) @@ -5939,13 +6919,15 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, } if (broadcast_type_C == 3) { - const float32x2_t _c0 = vld1_f32(pC); + float32x2_t _c0 = vld1_f32(pC); _out0 = beta == 1.f ? vadd_f32(_out0, _c0) : vmla_n_f32(_out0, _c0, beta); + pC += 2; } if (broadcast_type_C == 4) { - const float32x2_t _c0 = beta == 1.f ? vld1_f32(pC) : vmul_n_f32(vld1_f32(pC), beta); + float32x2_t _c0 = beta == 1.f ? vld1_f32(pC) : vmul_n_f32(vld1_f32(pC), beta); _out0 = vadd_f32(_out0, _c0); + pC += 2; } } @@ -5965,208 +6947,9 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, } outptr0 += out_hstep * 2; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 2; pp += 2; } - for (; jj < max_jj; jj += 1) - { - float out0 = pp[0]; - - if (pC) - { - if (broadcast_type_C == 0) - { - out0 += c0; - } - if (broadcast_type_C == 1 || broadcast_type_C == 2) - { - out0 += c0; - } - if (broadcast_type_C == 3) - { - out0 += beta == 1.f ? pC[0] : pC[0] * beta; - } - if (broadcast_type_C == 4) - { - out0 += beta == 1.f ? pC[0] : pC[0] * beta; - } - } - - if (alpha != 1.f) - { - out0 *= alpha; - } - - outptr0[0] = out0; - - outptr0 += out_hstep; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC++; - pp += 1; - } - } -#else - for (; ii + 1 < max_ii; ii += 2) - { - pC = (const float*)C; - float* outptr0 = (float*)top_blob + (size_t)j * out_hstep + i + ii; - float c0 = 0.f; - float c1 = 0.f; - if (pC) - { - if (broadcast_type_C == 0) - { - float c = pC[0]; - if (beta != 1.f) - c *= beta; - c0 = c; - c1 = c; - } - if (broadcast_type_C == 1 || broadcast_type_C == 2) - { - c0 = pC[i + ii]; - if (beta != 1.f) - c0 *= beta; - c1 = pC[i + ii + 1]; - if (beta != 1.f) - c1 *= beta; - } - if (broadcast_type_C == 3) - pC = (const float*)C + (i + ii) * c_hstep + j; - if (broadcast_type_C == 4) - pC = (const float*)C + j; - } - - int jj = 0; - for (; jj + 1 < max_jj; jj += 2) - { - float out00 = pp[0]; - float out01 = pp[1]; - float out10 = pp[2]; - float out11 = pp[3]; - - if (pC) - { - if (broadcast_type_C == 0) - { - out00 += c0; - out01 += c0; - out10 += c0; - out11 += c0; - } - if (broadcast_type_C == 1 || broadcast_type_C == 2) - { - out00 += c0; - out01 += c0; - out10 += c1; - out11 += c1; - } - if (broadcast_type_C == 3) - { - out00 += beta == 1.f ? pC[0] : pC[0] * beta; - out01 += beta == 1.f ? pC[1] : pC[1] * beta; - out10 += beta == 1.f ? pC[c_hstep] : pC[c_hstep] * beta; - out11 += beta == 1.f ? pC[c_hstep + 1] : pC[c_hstep + 1] * beta; - } - if (broadcast_type_C == 4) - { - out00 += beta == 1.f ? pC[0] : pC[0] * beta; - out01 += beta == 1.f ? pC[1] : pC[1] * beta; - out10 += beta == 1.f ? pC[0] : pC[0] * beta; - out11 += beta == 1.f ? pC[1] : pC[1] * beta; - } - } - - if (alpha != 1.f) - { - out00 *= alpha; - out01 *= alpha; - out10 *= alpha; - out11 *= alpha; - } - - outptr0[0] = out00; - outptr0[out_hstep] = out01; - outptr0[1] = out10; - outptr0[out_hstep + 1] = out11; - - outptr0 += out_hstep * 2; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 2; - pp += 4; - } - for (; jj < max_jj; jj += 1) - { - float out00 = pp[0]; - float out10 = pp[1]; - - if (pC) - { - if (broadcast_type_C == 0) - { - out00 += c0; - out10 += c0; - } - if (broadcast_type_C == 1 || broadcast_type_C == 2) - { - out00 += c0; - out10 += c1; - } - if (broadcast_type_C == 3) - { - out00 += beta == 1.f ? pC[0] : pC[0] * beta; - out10 += beta == 1.f ? pC[c_hstep] : pC[c_hstep] * beta; - } - if (broadcast_type_C == 4) - { - out00 += beta == 1.f ? pC[0] : pC[0] * beta; - out10 += beta == 1.f ? pC[0] : pC[0] * beta; - } - } - - if (alpha != 1.f) - { - out00 *= alpha; - out10 *= alpha; - } - - outptr0[0] = out00; - outptr0[1] = out10; - - outptr0 += out_hstep; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC++; - pp += 2; - } - } - for (; ii < max_ii; ii++) - { - pC = (const float*)C; - float* outptr0 = (float*)top_blob + (size_t)j * out_hstep + i + ii; - float c0 = 0.f; - if (pC) - { - if (broadcast_type_C == 0) - { - float c = pC[0]; - if (beta != 1.f) - c *= beta; - c0 = c; - } - if (broadcast_type_C == 1 || broadcast_type_C == 2) - { - c0 = pC[i + ii]; - if (beta != 1.f) - c0 *= beta; - } - if (broadcast_type_C == 3) - pC = (const float*)C + (i + ii) * c_hstep + j; - if (broadcast_type_C == 4) - pC = (const float*)C + j; - } - - int jj = 0; +#endif // __ARM_NEON for (; jj + 1 < max_jj; jj += 2) { float out00 = pp[0]; @@ -6188,11 +6971,13 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, { out00 += beta == 1.f ? pC[0] : pC[0] * beta; out01 += beta == 1.f ? pC[1] : pC[1] * beta; + pC += 2; } if (broadcast_type_C == 4) { out00 += beta == 1.f ? pC[0] : pC[0] * beta; out01 += beta == 1.f ? pC[1] : pC[1] * beta; + pC += 2; } } @@ -6206,8 +6991,6 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, outptr0[out_hstep] = out01; outptr0 += out_hstep * 2; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 2; pp += 2; } for (; jj < max_jj; jj += 1) @@ -6227,10 +7010,12 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, if (broadcast_type_C == 3) { out00 += beta == 1.f ? pC[0] : pC[0] * beta; + pC++; } if (broadcast_type_C == 4) { out00 += beta == 1.f ? pC[0] : pC[0] * beta; + pC++; } } @@ -6242,10 +7027,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, outptr0[0] = out00; outptr0 += out_hstep; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC++; pp += 1; } } -#endif // __ARM_NEON } diff --git a/src/layer/gemm.cpp b/src/layer/gemm.cpp index a43d3e500eb1..603d560259af 100644 --- a/src/layer/gemm.cpp +++ b/src/layer/gemm.cpp @@ -94,8 +94,7 @@ static void weight_block_quantize_activation_row_int8(const Mat& A, int transA, continue; } - volatile double scale_fp64 = 127.0 / (double)absmax; - const float scale = (float)scale_fp64; + const float scale = 127.f / absmax; descale_ptr[g] = absmax / 127.f; for (int kk = 0; kk < max_kk; kk++) @@ -103,11 +102,7 @@ static void weight_block_quantize_activation_row_int8(const Mat& A, int transA, const int k = k0 + kk; float v = transA ? ((const float*)A)[k * A_hstep + i] : ptrA[k]; if (input_scale_ptr) - { v *= input_scale_ptr[k]; - volatile float v_ordered = v; - v = v_ordered; - } outptr[k] = weight_block_quantize_float2int8(v * scale); } } diff --git a/src/layer/loongarch/gemm_loongarch.cpp b/src/layer/loongarch/gemm_loongarch.cpp index aed643974e11..c6eb3e653ac3 100644 --- a/src/layer/loongarch/gemm_loongarch.cpp +++ b/src/layer/loongarch/gemm_loongarch.cpp @@ -7377,13 +7377,13 @@ static int gemm_BT_loongarch_wq_int8(const Mat& A, const Mat& packed_B, const Ma const Mat BT = packed_B.reshape(K, N); const Mat BT_descales = packed_B_descales.reshape(block_count, N); int TILE_M, TILE_N, TILE_K; - get_optimal_tile_mnk_wq_int8(M, N, K, constant_TILE_M, constant_TILE_N, constant_TILE_K, TILE_M, TILE_N, TILE_K, nT); + get_optimal_tile_mnk_wq_int8(M, N, K, block_size, constant_TILE_M, constant_TILE_N, constant_TILE_K, TILE_M, TILE_N, TILE_K, nT); - (void)TILE_K; const int mr = std::min(M, TILE_M); const int nr = std::min(N, TILE_N); const int nn_M = (M + TILE_M - 1) / TILE_M; const int nn_N = (N + TILE_N - 1) / TILE_N; + const int nn_K = (K + TILE_K - 1) / TILE_K; const float* input_scale_ptr = input_scales; Mat topT(nr * mr, 1, nT, (size_t)4u, opt.workspace_allocator); @@ -7397,19 +7397,27 @@ static int gemm_BT_loongarch_wq_int8(const Mat& A, const Mat& packed_B, const Ma if (AT.empty() || AT_descales.empty()) return -100; + const int nn_MK = nn_M * nn_K; #pragma omp parallel for num_threads(nT) - for (int ppi = 0; ppi < nn_M; ppi++) + for (int ppik = 0; ppik < nn_MK; ppik++) { + const int ppi = ppik / nn_K; + const int ppk = ppik % nn_K; const int i = ppi * TILE_M; + const int k = ppk * TILE_K; const int max_ii = std::min(M - i, TILE_M); + const int max_kk = std::min(K - k, TILE_K); + const int max_block_count = (max_kk + block_size - 1) / block_size; - Mat AT_tile = AT.channel(i / TILE_M); - Mat AT_descales_tile = AT_descales.channel(i / TILE_M); + Mat AT_channel = AT.channel(i / TILE_M); + Mat AT_descales_channel = AT_descales.channel(i / TILE_M); + Mat AT_tile(max_kk, max_ii, (signed char*)AT_channel + (size_t)k * mr, (size_t)1u); + Mat AT_descales_tile(max_block_count, max_ii, (float*)AT_descales_channel + (size_t)(k / block_size) * mr, (size_t)4u); if (transA) - transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); + transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, k, max_kk, block_size, input_scale_ptr); else - quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); + quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, k, max_kk, block_size, input_scale_ptr); } #pragma omp parallel for num_threads(nT) @@ -7421,12 +7429,20 @@ static int gemm_BT_loongarch_wq_int8(const Mat& A, const Mat& packed_B, const Ma const int j = ppj * TILE_N; const int max_ii = std::min(M - i, TILE_M); const int max_jj = std::min(N - j, TILE_N); - Mat AT_tile = AT.channel(i / TILE_M); - Mat AT_descales_tile = AT_descales.channel(i / TILE_M); Mat BT_tile = BT.row_range(j, max_jj); Mat BT_descales_tile = BT_descales.row_range(j, max_jj); Mat topT_tile = topT.channel(get_omp_thread_num()); - gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); + Mat AT_channel = AT.channel(i / TILE_M); + Mat AT_descales_channel = AT_descales.channel(i / TILE_M); + for (int k = 0; k < K; k += TILE_K) + { + const int max_kk = std::min(K - k, TILE_K); + const int max_block_count = (max_kk + block_size - 1) / block_size; + Mat AT_tile(max_kk, max_ii, (signed char*)AT_channel + (size_t)k * mr, (size_t)1u); + Mat AT_descales_tile(max_block_count, max_ii, (float*)AT_descales_channel + (size_t)(k / block_size) * mr, (size_t)4u); + + gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, k, max_kk, block_size); + } if (output_transpose) transpose_unpack_output_tile_wq_int8(topT_tile, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, M, alpha, beta); else @@ -7446,21 +7462,32 @@ static int gemm_BT_loongarch_wq_int8(const Mat& A, const Mat& packed_B, const Ma const int i = ppi * TILE_M; const int max_ii = std::min(M - i, TILE_M); - Mat AT_tile = ATX.channel(get_omp_thread_num()); - Mat AT_descales_tile = ATX_descales.channel(get_omp_thread_num()); + Mat AT_channel = ATX.channel(get_omp_thread_num()); + Mat AT_descales_channel = ATX_descales.channel(get_omp_thread_num()); Mat topT_tile = topT.channel(get_omp_thread_num()); - if (transA) - transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); - else - quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); - for (int j = 0; j < N; j += TILE_N) { const int max_jj = std::min(N - j, TILE_N); Mat BT_tile = BT.row_range(j, max_jj); Mat BT_descales_tile = BT_descales.row_range(j, max_jj); - gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); + for (int k = 0; k < K; k += TILE_K) + { + const int max_kk = std::min(K - k, TILE_K); + const int max_block_count = (max_kk + block_size - 1) / block_size; + Mat AT_tile(max_kk, max_ii, (signed char*)AT_channel + (size_t)k * mr, (size_t)1u); + Mat AT_descales_tile(max_block_count, max_ii, (float*)AT_descales_channel + (size_t)(k / block_size) * mr, (size_t)4u); + + if (j == 0) + { + if (transA) + transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, k, max_kk, block_size, input_scale_ptr); + else + quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, k, max_kk, block_size, input_scale_ptr); + } + + gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, k, max_kk, block_size); + } if (output_transpose) transpose_unpack_output_tile_wq_int8(topT_tile, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, M, alpha, beta); else diff --git a/src/layer/loongarch/gemm_wq_int8.h b/src/layer/loongarch/gemm_wq_int8.h index a2ca710c2472..abaa73acff06 100644 --- a/src/layer/loongarch/gemm_wq_int8.h +++ b/src/layer/loongarch/gemm_wq_int8.h @@ -15,13 +15,13 @@ static int pack_B_wq_int8(const Mat& B, const Mat& B_scales, Mat& BT, Mat& BT_de int panel_start = 0; int panel_count = 0; +#if __loongarch_sx #if __loongarch_asx const int nn8 = (N - panel_start) / 8; const int panel_start8 = panel_start; panel_start += nn8 * 8; panel_count += nn8; #endif -#if __loongarch_sx const int nn4 = (N - panel_start) / 4; const int panel_start4 = panel_start; panel_start += nn4 * 4; @@ -41,6 +41,7 @@ static int pack_B_wq_int8(const Mat& B, const Mat& B_scales, Mat& BT, Mat& BT_de int q = p; int j = 0; int nr = 1; +#if __loongarch_sx #if __loongarch_asx if (q < nn8) { @@ -51,7 +52,6 @@ static int pack_B_wq_int8(const Mat& B, const Mat& B_scales, Mat& BT, Mat& BT_de { q -= nn8; #endif -#if __loongarch_sx if (q < nn4) { j = panel_start4 + q * 4; @@ -74,50 +74,267 @@ static int pack_B_wq_int8(const Mat& B, const Mat& B_scales, Mat& BT, Mat& BT_de } #if __loongarch_sx } -#endif #if __loongarch_asx } +#endif #endif signed char* pp = (signed char*)BT_packed + j * K; float* pd = (float*)BT_packed_descales + j * block_count; - for (int g = 0; g < block_count; g++) +#if __loongarch_sx +#if __loongarch_asx + if (nr == 8) { - const int k0 = g * block_size; - const int max_kk = std::min(K - k0, block_size); - int kk = 0; - for (; kk + 3 < max_kk; kk += 4) + const signed char* p0 = B.row(j); + const signed char* p1 = B.row(j + 1); + const signed char* p2 = B.row(j + 2); + const signed char* p3 = B.row(j + 3); + const signed char* p4 = B.row(j + 4); + const signed char* p5 = B.row(j + 5); + const signed char* p6 = B.row(j + 6); + const signed char* p7 = B.row(j + 7); + const float* s0 = B_scales.row(j); + const float* s1 = B_scales.row(j + 1); + const float* s2 = B_scales.row(j + 2); + const float* s3 = B_scales.row(j + 3); + const float* s4 = B_scales.row(j + 4); + const float* s5 = B_scales.row(j + 5); + const float* s6 = B_scales.row(j + 6); + const float* s7 = B_scales.row(j + 7); + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __m128i _p0 = __lsx_vldrepl_w(p0, 0); + __m128i _p1 = __lsx_vldrepl_w(p1, 0); + __m128i _p2 = __lsx_vldrepl_w(p2, 0); + __m128i _p3 = __lsx_vldrepl_w(p3, 0); + __m128i _p4 = __lsx_vldrepl_w(p4, 0); + __m128i _p5 = __lsx_vldrepl_w(p5, 0); + __m128i _p6 = __lsx_vldrepl_w(p6, 0); + __m128i _p7 = __lsx_vldrepl_w(p7, 0); + __m128i _p01 = __lsx_vilvl_w(_p1, _p0); + __m128i _p23 = __lsx_vilvl_w(_p3, _p2); + __m128i _p45 = __lsx_vilvl_w(_p5, _p4); + __m128i _p67 = __lsx_vilvl_w(_p7, _p6); + __m128i _p0123 = __lsx_vilvl_d(_p23, _p01); + __m128i _p4567 = __lsx_vilvl_d(_p67, _p45); + __lasx_xvst(__lasx_concat_128(_p0123, _p4567), pp, 0); + pp += 32; + p0 += 4; + p1 += 4; + p2 += 4; + p3 += 4; + p4 += 4; + p5 += 4; + p6 += 4; + p7 += 4; + } + if (kk + 1 < max_kk) + { + pp[0] = p0[0]; + pp[1] = p0[1]; + pp[2] = p1[0]; + pp[3] = p1[1]; + pp[4] = p2[0]; + pp[5] = p2[1]; + pp[6] = p3[0]; + pp[7] = p3[1]; + pp[8] = p4[0]; + pp[9] = p4[1]; + pp[10] = p5[0]; + pp[11] = p5[1]; + pp[12] = p6[0]; + pp[13] = p6[1]; + pp[14] = p7[0]; + pp[15] = p7[1]; + pp += 16; + p0 += 2; + p1 += 2; + p2 += 2; + p3 += 2; + p4 += 2; + p5 += 2; + p6 += 2; + p7 += 2; + kk += 2; + } + if (kk < max_kk) + { + pp[0] = *p0++; + pp[1] = *p1++; + pp[2] = *p2++; + pp[3] = *p3++; + pp[4] = *p4++; + pp[5] = *p5++; + pp[6] = *p6++; + pp[7] = *p7++; + pp += 8; + } + + *pd++ = 1.f / *s0++; + *pd++ = 1.f / *s1++; + *pd++ = 1.f / *s2++; + *pd++ = 1.f / *s3++; + *pd++ = 1.f / *s4++; + *pd++ = 1.f / *s5++; + *pd++ = 1.f / *s6++; + *pd++ = 1.f / *s7++; + } + + continue; + } +#endif + if (nr == 4) + { + const signed char* p0 = B.row(j); + const signed char* p1 = B.row(j + 1); + const signed char* p2 = B.row(j + 2); + const signed char* p3 = B.row(j + 3); + const float* s0 = B_scales.row(j); + const float* s1 = B_scales.row(j + 1); + const float* s2 = B_scales.row(j + 2); + const float* s3 = B_scales.row(j + 3); + + for (int g = 0; g < block_count; g++) { - for (int jj = 0; jj < nr; jj++) + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __m128i _p0 = __lsx_vldrepl_w(p0, 0); + __m128i _p1 = __lsx_vldrepl_w(p1, 0); + __m128i _p2 = __lsx_vldrepl_w(p2, 0); + __m128i _p3 = __lsx_vldrepl_w(p3, 0); + __m128i _p01 = __lsx_vilvl_w(_p1, _p0); + __m128i _p23 = __lsx_vilvl_w(_p3, _p2); + __lsx_vst(__lsx_vilvl_d(_p23, _p01), pp, 0); + pp += 16; + p0 += 4; + p1 += 4; + p2 += 4; + p3 += 4; + } + if (kk + 1 < max_kk) { - const signed char* pB = B.row(j + jj) + k0 + kk; - pp[0] = pB[0]; - pp[1] = pB[1]; - pp[2] = pB[2]; - pp[3] = pB[3]; + pp[0] = p0[0]; + pp[1] = p0[1]; + pp[2] = p1[0]; + pp[3] = p1[1]; + pp[4] = p2[0]; + pp[5] = p2[1]; + pp[6] = p3[0]; + pp[7] = p3[1]; + pp += 8; + p0 += 2; + p1 += 2; + p2 += 2; + p3 += 2; + kk += 2; + } + if (kk < max_kk) + { + pp[0] = *p0++; + pp[1] = *p1++; + pp[2] = *p2++; + pp[3] = *p3++; pp += 4; } + + *pd++ = 1.f / *s0++; + *pd++ = 1.f / *s1++; + *pd++ = 1.f / *s2++; + *pd++ = 1.f / *s3++; } - if (kk + 1 < max_kk) + + continue; + } +#endif + if (nr == 2) + { + const signed char* p0 = B.row(j); + const signed char* p1 = B.row(j + 1); + const float* s0 = B_scales.row(j); + const float* s1 = B_scales.row(j + 1); + + for (int g = 0; g < block_count; g++) { - for (int jj = 0; jj < nr; jj++) + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + pp[0] = p0[0]; + pp[1] = p0[1]; + pp[2] = p0[2]; + pp[3] = p0[3]; + pp[4] = p1[0]; + pp[5] = p1[1]; + pp[6] = p1[2]; + pp[7] = p1[3]; + pp += 8; + p0 += 4; + p1 += 4; + } + if (kk + 1 < max_kk) { - const signed char* pB = B.row(j + jj) + k0 + kk; - pp[0] = pB[0]; - pp[1] = pB[1]; + pp[0] = p0[0]; + pp[1] = p0[1]; + pp[2] = p1[0]; + pp[3] = p1[1]; + pp += 4; + p0 += 2; + p1 += 2; + kk += 2; + } + if (kk < max_kk) + { + pp[0] = *p0++; + pp[1] = *p1++; pp += 2; } - kk += 2; + + *pd++ = 1.f / *s0++; + *pd++ = 1.f / *s1++; } - if (kk < max_kk) + + continue; + } + + const signed char* p0 = B.row(j); + const float* s0 = B_scales.row(j); + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) { - for (int jj = 0; jj < nr; jj++) - *pp++ = B.row(j + jj)[k0 + kk]; + pp[0] = p0[0]; + pp[1] = p0[1]; + pp[2] = p0[2]; + pp[3] = p0[3]; + pp += 4; + p0 += 4; } + if (kk + 1 < max_kk) + { + pp[0] = p0[0]; + pp[1] = p0[1]; + pp += 2; + p0 += 2; + kk += 2; + } + if (kk < max_kk) + *pp++ = *p0++; - for (int jj = 0; jj < nr; jj++) - pd[g * nr + jj] = 1.f / B_scales.row(j + jj)[g]; + *pd++ = 1.f / *s0++; } } @@ -126,20 +343,23 @@ static int pack_B_wq_int8(const Mat& B, const Mat& B_scales, Mat& BT, Mat& BT_de return 0; } -static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr) +static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int k, int max_kk, int block_size, const float* input_scale_ptr) { signed char* outptr = AT_tile; const int out_hstep = AT_tile.w; float* descales = AT_descales_tile; const int descales_hstep = AT_descales_tile.w; + const int K = max_kk; const int block_count = (K + block_size - 1) / block_size; const size_t A_hstep = A.dims == 3 ? A.cstep : (size_t)A.w; + const float* A_data = (const float*)A + k; + input_scale_ptr = input_scale_ptr ? input_scale_ptr + k : 0; int ii = 0; #if __loongarch_sx for (; ii + 7 < max_ii; ii += 8) { - const float* p0 = (const float*)A + (i + ii) * A_hstep; + const float* p0 = A_data + (i + ii) * A_hstep; const float* p1 = p0 + A_hstep; const float* p2 = p1 + A_hstep; const float* p3 = p2 + A_hstep; @@ -155,6 +375,15 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales { const int k0 = g * block_size; const int max_kk = std::min(K - k0, block_size); + const float* p0g = p0 + k0; + const float* p1g = p1 + k0; + const float* p2g = p2 + k0; + const float* p3g = p3 + k0; + const float* p4g = p4 + k0; + const float* p5g = p5 + k0; + const float* p6g = p6 + k0; + const float* p7g = p7 + k0; + const float* sg = input_scale_ptr ? input_scale_ptr + k0 : 0; const __m128i _abs_mask = __lsx_vreplgr2vr_w(0x7fffffff); __m128 _absmax0 = (__m128)__lsx_vldi(0); __m128 _absmax1 = (__m128)__lsx_vldi(0); @@ -164,20 +393,29 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales __m128 _absmax5 = (__m128)__lsx_vldi(0); __m128 _absmax6 = (__m128)__lsx_vldi(0); __m128 _absmax7 = (__m128)__lsx_vldi(0); + const float* p0a = p0g; + const float* p1a = p1g; + const float* p2a = p2g; + const float* p3a = p3g; + const float* p4a = p4g; + const float* p5a = p5g; + const float* p6a = p6g; + const float* p7a = p7g; + const float* psa = sg; int kk = 0; for (; kk + 3 < max_kk; kk += 4) { - __m128 _v0 = (__m128)__lsx_vld(p0 + k0 + kk, 0); - __m128 _v1 = (__m128)__lsx_vld(p1 + k0 + kk, 0); - __m128 _v2 = (__m128)__lsx_vld(p2 + k0 + kk, 0); - __m128 _v3 = (__m128)__lsx_vld(p3 + k0 + kk, 0); - __m128 _v4 = (__m128)__lsx_vld(p4 + k0 + kk, 0); - __m128 _v5 = (__m128)__lsx_vld(p5 + k0 + kk, 0); - __m128 _v6 = (__m128)__lsx_vld(p6 + k0 + kk, 0); - __m128 _v7 = (__m128)__lsx_vld(p7 + k0 + kk, 0); - if (input_scale_ptr) - { - const __m128 _s = (__m128)__lsx_vld(input_scale_ptr + k0 + kk, 0); + __m128 _v0 = (__m128)__lsx_vld(p0a, 0); + __m128 _v1 = (__m128)__lsx_vld(p1a, 0); + __m128 _v2 = (__m128)__lsx_vld(p2a, 0); + __m128 _v3 = (__m128)__lsx_vld(p3a, 0); + __m128 _v4 = (__m128)__lsx_vld(p4a, 0); + __m128 _v5 = (__m128)__lsx_vld(p5a, 0); + __m128 _v6 = (__m128)__lsx_vld(p6a, 0); + __m128 _v7 = (__m128)__lsx_vld(p7a, 0); + if (psa) + { + __m128 _s = (__m128)__lsx_vld(psa, 0); _v0 = __lsx_vfmul_s(_v0, _s); _v1 = __lsx_vfmul_s(_v1, _s); _v2 = __lsx_vfmul_s(_v2, _s); @@ -195,6 +433,16 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales _absmax5 = __lsx_vfmax_s(_absmax5, (__m128)__lsx_vand_v((__m128i)_v5, _abs_mask)); _absmax6 = __lsx_vfmax_s(_absmax6, (__m128)__lsx_vand_v((__m128i)_v6, _abs_mask)); _absmax7 = __lsx_vfmax_s(_absmax7, (__m128)__lsx_vand_v((__m128i)_v7, _abs_mask)); + p0a += 4; + p1a += 4; + p2a += 4; + p3a += 4; + p4a += 4; + p5a += 4; + p6a += 4; + p7a += 4; + if (psa) + psa += 4; } float absmax0 = __lsx_reduce_fmax_s(_absmax0); float absmax1 = __lsx_reduce_fmax_s(_absmax1); @@ -206,16 +454,15 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales float absmax7 = __lsx_reduce_fmax_s(_absmax7); for (; kk < max_kk; kk++) { - const int k = k0 + kk; - const float s = input_scale_ptr ? input_scale_ptr[k] : 1.f; - absmax0 = std::max(absmax0, fabsf(p0[k] * s)); - absmax1 = std::max(absmax1, fabsf(p1[k] * s)); - absmax2 = std::max(absmax2, fabsf(p2[k] * s)); - absmax3 = std::max(absmax3, fabsf(p3[k] * s)); - absmax4 = std::max(absmax4, fabsf(p4[k] * s)); - absmax5 = std::max(absmax5, fabsf(p5[k] * s)); - absmax6 = std::max(absmax6, fabsf(p6[k] * s)); - absmax7 = std::max(absmax7, fabsf(p7[k] * s)); + const float s = psa ? *psa++ : 1.f; + absmax0 = std::max(absmax0, fabsf(*p0a++ * s)); + absmax1 = std::max(absmax1, fabsf(*p1a++ * s)); + absmax2 = std::max(absmax2, fabsf(*p2a++ * s)); + absmax3 = std::max(absmax3, fabsf(*p3a++ * s)); + absmax4 = std::max(absmax4, fabsf(*p4a++ * s)); + absmax5 = std::max(absmax5, fabsf(*p5a++ * s)); + absmax6 = std::max(absmax6, fabsf(*p6a++ * s)); + absmax7 = std::max(absmax7, fabsf(*p7a++ * s)); } pd[0] = absmax0 / 127.f; @@ -228,45 +475,45 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales pd[7] = absmax7 / 127.f; pd += 8; - volatile double scale0_fp64 = absmax0 == 0.f ? 0.0 : 127.0 / (double)absmax0; - volatile double scale1_fp64 = absmax1 == 0.f ? 0.0 : 127.0 / (double)absmax1; - volatile double scale2_fp64 = absmax2 == 0.f ? 0.0 : 127.0 / (double)absmax2; - volatile double scale3_fp64 = absmax3 == 0.f ? 0.0 : 127.0 / (double)absmax3; - volatile double scale4_fp64 = absmax4 == 0.f ? 0.0 : 127.0 / (double)absmax4; - volatile double scale5_fp64 = absmax5 == 0.f ? 0.0 : 127.0 / (double)absmax5; - volatile double scale6_fp64 = absmax6 == 0.f ? 0.0 : 127.0 / (double)absmax6; - volatile double scale7_fp64 = absmax7 == 0.f ? 0.0 : 127.0 / (double)absmax7; - const float scale0 = (float)scale0_fp64; - const float scale1 = (float)scale1_fp64; - const float scale2 = (float)scale2_fp64; - const float scale3 = (float)scale3_fp64; - const float scale4 = (float)scale4_fp64; - const float scale5 = (float)scale5_fp64; - const float scale6 = (float)scale6_fp64; - const float scale7 = (float)scale7_fp64; - const __m128 _scale0 = __lsx_vreplfr2vr_s(scale0); - const __m128 _scale1 = __lsx_vreplfr2vr_s(scale1); - const __m128 _scale2 = __lsx_vreplfr2vr_s(scale2); - const __m128 _scale3 = __lsx_vreplfr2vr_s(scale3); - const __m128 _scale4 = __lsx_vreplfr2vr_s(scale4); - const __m128 _scale5 = __lsx_vreplfr2vr_s(scale5); - const __m128 _scale6 = __lsx_vreplfr2vr_s(scale6); - const __m128 _scale7 = __lsx_vreplfr2vr_s(scale7); + const float scale0 = absmax0 == 0.f ? 0.f : 127.f / absmax0; + const float scale1 = absmax1 == 0.f ? 0.f : 127.f / absmax1; + const float scale2 = absmax2 == 0.f ? 0.f : 127.f / absmax2; + const float scale3 = absmax3 == 0.f ? 0.f : 127.f / absmax3; + const float scale4 = absmax4 == 0.f ? 0.f : 127.f / absmax4; + const float scale5 = absmax5 == 0.f ? 0.f : 127.f / absmax5; + const float scale6 = absmax6 == 0.f ? 0.f : 127.f / absmax6; + const float scale7 = absmax7 == 0.f ? 0.f : 127.f / absmax7; + __m128 _scale0 = __lsx_vreplfr2vr_s(scale0); + __m128 _scale1 = __lsx_vreplfr2vr_s(scale1); + __m128 _scale2 = __lsx_vreplfr2vr_s(scale2); + __m128 _scale3 = __lsx_vreplfr2vr_s(scale3); + __m128 _scale4 = __lsx_vreplfr2vr_s(scale4); + __m128 _scale5 = __lsx_vreplfr2vr_s(scale5); + __m128 _scale6 = __lsx_vreplfr2vr_s(scale6); + __m128 _scale7 = __lsx_vreplfr2vr_s(scale7); + const float* p0q = p0g; + const float* p1q = p1g; + const float* p2q = p2g; + const float* p3q = p3g; + const float* p4q = p4g; + const float* p5q = p5g; + const float* p6q = p6g; + const float* p7q = p7g; + const float* psq = sg; kk = 0; for (; kk + 3 < max_kk; kk += 4) { - const int k = k0 + kk; - __m128 _v0 = (__m128)__lsx_vld(p0 + k, 0); - __m128 _v1 = (__m128)__lsx_vld(p1 + k, 0); - __m128 _v2 = (__m128)__lsx_vld(p2 + k, 0); - __m128 _v3 = (__m128)__lsx_vld(p3 + k, 0); - __m128 _v4 = (__m128)__lsx_vld(p4 + k, 0); - __m128 _v5 = (__m128)__lsx_vld(p5 + k, 0); - __m128 _v6 = (__m128)__lsx_vld(p6 + k, 0); - __m128 _v7 = (__m128)__lsx_vld(p7 + k, 0); - if (input_scale_ptr) - { - const __m128 _s = (__m128)__lsx_vld(input_scale_ptr + k, 0); + __m128 _v0 = (__m128)__lsx_vld(p0q, 0); + __m128 _v1 = (__m128)__lsx_vld(p1q, 0); + __m128 _v2 = (__m128)__lsx_vld(p2q, 0); + __m128 _v3 = (__m128)__lsx_vld(p3q, 0); + __m128 _v4 = (__m128)__lsx_vld(p4q, 0); + __m128 _v5 = (__m128)__lsx_vld(p5q, 0); + __m128 _v6 = (__m128)__lsx_vld(p6q, 0); + __m128 _v7 = (__m128)__lsx_vld(p7q, 0); + if (psq) + { + __m128 _s = (__m128)__lsx_vld(psq, 0); _v0 = __lsx_vfmul_s(_v0, _s); _v1 = __lsx_vfmul_s(_v1, _s); _v2 = __lsx_vfmul_s(_v2, _s); @@ -281,28 +528,35 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales *((int64_t*)(pp + 16)) = float2int8(__lsx_vfmul_s(_v4, _scale4), __lsx_vfmul_s(_v5, _scale5)); *((int64_t*)(pp + 24)) = float2int8(__lsx_vfmul_s(_v6, _scale6), __lsx_vfmul_s(_v7, _scale7)); pp += 32; + p0q += 4; + p1q += 4; + p2q += 4; + p3q += 4; + p4q += 4; + p5q += 4; + p6q += 4; + p7q += 4; + if (psq) + psq += 4; } for (; kk < max_kk; kk++) { - const int k = k0 + kk; - const float s = input_scale_ptr ? input_scale_ptr[k] : 1.f; - pp[0] = float2int8(p0[k] * s * scale0); - pp[1] = float2int8(p1[k] * s * scale1); - pp[2] = float2int8(p2[k] * s * scale2); - pp[3] = float2int8(p3[k] * s * scale3); - pp[4] = float2int8(p4[k] * s * scale4); - pp[5] = float2int8(p5[k] * s * scale5); - pp[6] = float2int8(p6[k] * s * scale6); - pp[7] = float2int8(p7[k] * s * scale7); + const float s = psq ? *psq++ : 1.f; + pp[0] = float2int8(*p0q++ * s * scale0); + pp[1] = float2int8(*p1q++ * s * scale1); + pp[2] = float2int8(*p2q++ * s * scale2); + pp[3] = float2int8(*p3q++ * s * scale3); + pp[4] = float2int8(*p4q++ * s * scale4); + pp[5] = float2int8(*p5q++ * s * scale5); + pp[6] = float2int8(*p6q++ * s * scale6); + pp[7] = float2int8(*p7q++ * s * scale7); pp += 8; } } } -#endif // __loongarch_sx -#if __loongarch_sx for (; ii + 3 < max_ii; ii += 4) { - const float* p0 = (const float*)A + (i + ii) * A_hstep; + const float* p0 = A_data + (i + ii) * A_hstep; const float* p1 = p0 + A_hstep; const float* p2 = p1 + A_hstep; const float* p3 = p2 + A_hstep; @@ -314,20 +568,30 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales { const int k0 = g * block_size; const int max_kk = std::min(K - k0, block_size); + const float* p0g = p0 + k0; + const float* p1g = p1 + k0; + const float* p2g = p2 + k0; + const float* p3g = p3 + k0; + const float* sg = input_scale_ptr ? input_scale_ptr + k0 : 0; __m128 _absmax0 = (__m128)__lsx_vldi(0); __m128 _absmax1 = (__m128)__lsx_vldi(0); __m128 _absmax2 = (__m128)__lsx_vldi(0); __m128 _absmax3 = (__m128)__lsx_vldi(0); + const float* p0a = p0g; + const float* p1a = p1g; + const float* p2a = p2g; + const float* p3a = p3g; + const float* psa = sg; int kk = 0; for (; kk + 3 < max_kk; kk += 4) { - __m128 _v0 = (__m128)__lsx_vld(p0 + k0 + kk, 0); - __m128 _v1 = (__m128)__lsx_vld(p1 + k0 + kk, 0); - __m128 _v2 = (__m128)__lsx_vld(p2 + k0 + kk, 0); - __m128 _v3 = (__m128)__lsx_vld(p3 + k0 + kk, 0); - if (input_scale_ptr) + __m128 _v0 = (__m128)__lsx_vld(p0a, 0); + __m128 _v1 = (__m128)__lsx_vld(p1a, 0); + __m128 _v2 = (__m128)__lsx_vld(p2a, 0); + __m128 _v3 = (__m128)__lsx_vld(p3a, 0); + if (psa) { - const __m128 _s = (__m128)__lsx_vld(input_scale_ptr + k0 + kk, 0); + __m128 _s = (__m128)__lsx_vld(psa, 0); _v0 = __lsx_vfmul_s(_v0, _s); _v1 = __lsx_vfmul_s(_v1, _s); _v2 = __lsx_vfmul_s(_v2, _s); @@ -337,6 +601,12 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales _absmax1 = __lsx_vfmax_s(_absmax1, (__m128)__lsx_vand_v((__m128i)_v1, _abs_mask)); _absmax2 = __lsx_vfmax_s(_absmax2, (__m128)__lsx_vand_v((__m128i)_v2, _abs_mask)); _absmax3 = __lsx_vfmax_s(_absmax3, (__m128)__lsx_vand_v((__m128i)_v3, _abs_mask)); + p0a += 4; + p1a += 4; + p2a += 4; + p3a += 4; + if (psa) + psa += 4; } float absmax0 = __lsx_reduce_fmax_s(_absmax0); float absmax1 = __lsx_reduce_fmax_s(_absmax1); @@ -344,40 +614,40 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales float absmax3 = __lsx_reduce_fmax_s(_absmax3); for (; kk < max_kk; kk++) { - const int k = k0 + kk; - const float s = input_scale_ptr ? input_scale_ptr[k] : 1.f; - absmax0 = std::max(absmax0, fabsf(p0[k] * s)); - absmax1 = std::max(absmax1, fabsf(p1[k] * s)); - absmax2 = std::max(absmax2, fabsf(p2[k] * s)); - absmax3 = std::max(absmax3, fabsf(p3[k] * s)); + const float s = psa ? *psa++ : 1.f; + absmax0 = std::max(absmax0, fabsf(*p0a++ * s)); + absmax1 = std::max(absmax1, fabsf(*p1a++ * s)); + absmax2 = std::max(absmax2, fabsf(*p2a++ * s)); + absmax3 = std::max(absmax3, fabsf(*p3a++ * s)); } pd[0] = absmax0 / 127.f; pd[1] = absmax1 / 127.f; pd[2] = absmax2 / 127.f; pd[3] = absmax3 / 127.f; pd += 4; - volatile double scale0_fp64 = absmax0 == 0.f ? 0.0 : 127.0 / (double)absmax0; - volatile double scale1_fp64 = absmax1 == 0.f ? 0.0 : 127.0 / (double)absmax1; - volatile double scale2_fp64 = absmax2 == 0.f ? 0.0 : 127.0 / (double)absmax2; - volatile double scale3_fp64 = absmax3 == 0.f ? 0.0 : 127.0 / (double)absmax3; - const float scale0 = (float)scale0_fp64; - const float scale1 = (float)scale1_fp64; - const float scale2 = (float)scale2_fp64; - const float scale3 = (float)scale3_fp64; - const __m128 _scale0 = __lsx_vreplfr2vr_s(scale0); - const __m128 _scale1 = __lsx_vreplfr2vr_s(scale1); - const __m128 _scale2 = __lsx_vreplfr2vr_s(scale2); - const __m128 _scale3 = __lsx_vreplfr2vr_s(scale3); + const float scale0 = absmax0 == 0.f ? 0.f : 127.f / absmax0; + const float scale1 = absmax1 == 0.f ? 0.f : 127.f / absmax1; + const float scale2 = absmax2 == 0.f ? 0.f : 127.f / absmax2; + const float scale3 = absmax3 == 0.f ? 0.f : 127.f / absmax3; + __m128 _scale0 = __lsx_vreplfr2vr_s(scale0); + __m128 _scale1 = __lsx_vreplfr2vr_s(scale1); + __m128 _scale2 = __lsx_vreplfr2vr_s(scale2); + __m128 _scale3 = __lsx_vreplfr2vr_s(scale3); + const float* p0q = p0g; + const float* p1q = p1g; + const float* p2q = p2g; + const float* p3q = p3g; + const float* psq = sg; kk = 0; for (; kk + 3 < max_kk; kk += 4) { - __m128 _v0 = (__m128)__lsx_vld(p0 + k0 + kk, 0); - __m128 _v1 = (__m128)__lsx_vld(p1 + k0 + kk, 0); - __m128 _v2 = (__m128)__lsx_vld(p2 + k0 + kk, 0); - __m128 _v3 = (__m128)__lsx_vld(p3 + k0 + kk, 0); - if (input_scale_ptr) + __m128 _v0 = (__m128)__lsx_vld(p0q, 0); + __m128 _v1 = (__m128)__lsx_vld(p1q, 0); + __m128 _v2 = (__m128)__lsx_vld(p2q, 0); + __m128 _v3 = (__m128)__lsx_vld(p3q, 0); + if (psq) { - const __m128 _s = (__m128)__lsx_vld(input_scale_ptr + k0 + kk, 0); + __m128 _s = (__m128)__lsx_vld(psq, 0); _v0 = __lsx_vfmul_s(_v0, _s); _v1 = __lsx_vfmul_s(_v1, _s); _v2 = __lsx_vfmul_s(_v2, _s); @@ -386,38 +656,48 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales *((int64_t*)pp) = float2int8(__lsx_vfmul_s(_v0, _scale0), __lsx_vfmul_s(_v1, _scale1)); *((int64_t*)(pp + 8)) = float2int8(__lsx_vfmul_s(_v2, _scale2), __lsx_vfmul_s(_v3, _scale3)); pp += 16; + p0q += 4; + p1q += 4; + p2q += 4; + p3q += 4; + if (psq) + psq += 4; } if (kk + 1 < max_kk) { - const int k = k0 + kk; - const float s0 = input_scale_ptr ? input_scale_ptr[k] : 1.f; - const float s1 = input_scale_ptr ? input_scale_ptr[k + 1] : 1.f; - pp[0] = float2int8(p0[k] * s0 * scale0); - pp[1] = float2int8(p0[k + 1] * s1 * scale0); - pp[2] = float2int8(p1[k] * s0 * scale1); - pp[3] = float2int8(p1[k + 1] * s1 * scale1); - pp[4] = float2int8(p2[k] * s0 * scale2); - pp[5] = float2int8(p2[k + 1] * s1 * scale2); - pp[6] = float2int8(p3[k] * s0 * scale3); - pp[7] = float2int8(p3[k + 1] * s1 * scale3); + const float s0 = psq ? psq[0] : 1.f; + const float s1 = psq ? psq[1] : 1.f; + pp[0] = float2int8(p0q[0] * s0 * scale0); + pp[1] = float2int8(p0q[1] * s1 * scale0); + pp[2] = float2int8(p1q[0] * s0 * scale1); + pp[3] = float2int8(p1q[1] * s1 * scale1); + pp[4] = float2int8(p2q[0] * s0 * scale2); + pp[5] = float2int8(p2q[1] * s1 * scale2); + pp[6] = float2int8(p3q[0] * s0 * scale3); + pp[7] = float2int8(p3q[1] * s1 * scale3); pp += 8; + p0q += 2; + p1q += 2; + p2q += 2; + p3q += 2; + if (psq) + psq += 2; kk += 2; } if (kk < max_kk) { - const int k = k0 + kk; - const float s = input_scale_ptr ? input_scale_ptr[k] : 1.f; - pp[0] = float2int8(p0[k] * s * scale0); - pp[1] = float2int8(p1[k] * s * scale1); - pp[2] = float2int8(p2[k] * s * scale2); - pp[3] = float2int8(p3[k] * s * scale3); + const float s = psq ? *psq : 1.f; + pp[0] = float2int8(*p0q * s * scale0); + pp[1] = float2int8(*p1q * s * scale1); + pp[2] = float2int8(*p2q * s * scale2); + pp[3] = float2int8(*p3q * s * scale3); pp += 4; } } } for (; ii + 1 < max_ii; ii += 2) { - const float* p0 = (const float*)A + (i + ii) * A_hstep; + const float* p0 = A_data + (i + ii) * A_hstep; const float* p1 = p0 + A_hstep; signed char* pp = outptr + ii * out_hstep; float* pd = descales + ii * descales_hstep; @@ -427,88 +707,186 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales { const int k0 = g * block_size; const int max_kk = std::min(K - k0, block_size); + const float* p0g = p0 + k0; + const float* p1g = p1 + k0; + const float* sg = input_scale_ptr ? input_scale_ptr + k0 : 0; __m128 _absmax0 = (__m128)__lsx_vldi(0); __m128 _absmax1 = (__m128)__lsx_vldi(0); + const float* p0a = p0g; + const float* p1a = p1g; + const float* psa = sg; int kk = 0; for (; kk + 3 < max_kk; kk += 4) { - __m128 _v0 = (__m128)__lsx_vld(p0 + k0 + kk, 0); - __m128 _v1 = (__m128)__lsx_vld(p1 + k0 + kk, 0); - if (input_scale_ptr) + __m128 _v0 = (__m128)__lsx_vld(p0a, 0); + __m128 _v1 = (__m128)__lsx_vld(p1a, 0); + if (psa) { - const __m128 _s = (__m128)__lsx_vld(input_scale_ptr + k0 + kk, 0); + __m128 _s = (__m128)__lsx_vld(psa, 0); _v0 = __lsx_vfmul_s(_v0, _s); _v1 = __lsx_vfmul_s(_v1, _s); } _absmax0 = __lsx_vfmax_s(_absmax0, (__m128)__lsx_vand_v((__m128i)_v0, _abs_mask)); _absmax1 = __lsx_vfmax_s(_absmax1, (__m128)__lsx_vand_v((__m128i)_v1, _abs_mask)); + p0a += 4; + p1a += 4; + if (psa) + psa += 4; } float absmax0 = __lsx_reduce_fmax_s(_absmax0); float absmax1 = __lsx_reduce_fmax_s(_absmax1); for (; kk < max_kk; kk++) { - const int k = k0 + kk; - const float s = input_scale_ptr ? input_scale_ptr[k] : 1.f; - absmax0 = std::max(absmax0, fabsf(p0[k] * s)); - absmax1 = std::max(absmax1, fabsf(p1[k] * s)); + const float s = psa ? *psa++ : 1.f; + absmax0 = std::max(absmax0, fabsf(*p0a++ * s)); + absmax1 = std::max(absmax1, fabsf(*p1a++ * s)); } pd[0] = absmax0 / 127.f; pd[1] = absmax1 / 127.f; pd += 2; - volatile double scale0_fp64 = absmax0 == 0.f ? 0.0 : 127.0 / (double)absmax0; - volatile double scale1_fp64 = absmax1 == 0.f ? 0.0 : 127.0 / (double)absmax1; - const float scale0 = (float)scale0_fp64; - const float scale1 = (float)scale1_fp64; - const __m128 _scale0 = __lsx_vreplfr2vr_s(scale0); - const __m128 _scale1 = __lsx_vreplfr2vr_s(scale1); + const float scale0 = absmax0 == 0.f ? 0.f : 127.f / absmax0; + const float scale1 = absmax1 == 0.f ? 0.f : 127.f / absmax1; + __m128 _scale0 = __lsx_vreplfr2vr_s(scale0); + __m128 _scale1 = __lsx_vreplfr2vr_s(scale1); + const float* p0q = p0g; + const float* p1q = p1g; + const float* psq = sg; kk = 0; for (; kk + 3 < max_kk; kk += 4) { - __m128 _v0 = (__m128)__lsx_vld(p0 + k0 + kk, 0); - __m128 _v1 = (__m128)__lsx_vld(p1 + k0 + kk, 0); - if (input_scale_ptr) + __m128 _v0 = (__m128)__lsx_vld(p0q, 0); + __m128 _v1 = (__m128)__lsx_vld(p1q, 0); + if (psq) { - const __m128 _s = (__m128)__lsx_vld(input_scale_ptr + k0 + kk, 0); + __m128 _s = (__m128)__lsx_vld(psq, 0); _v0 = __lsx_vfmul_s(_v0, _s); _v1 = __lsx_vfmul_s(_v1, _s); } *((int64_t*)pp) = float2int8(__lsx_vfmul_s(_v0, _scale0), __lsx_vfmul_s(_v1, _scale1)); pp += 8; + p0q += 4; + p1q += 4; + if (psq) + psq += 4; } if (kk + 1 < max_kk) { - const int k = k0 + kk; - const float s0 = input_scale_ptr ? input_scale_ptr[k] : 1.f; - const float s1 = input_scale_ptr ? input_scale_ptr[k + 1] : 1.f; - pp[0] = float2int8(p0[k] * s0 * scale0); - pp[1] = float2int8(p0[k + 1] * s1 * scale0); - pp[2] = float2int8(p1[k] * s0 * scale1); - pp[3] = float2int8(p1[k + 1] * s1 * scale1); + const float s0 = psq ? psq[0] : 1.f; + const float s1 = psq ? psq[1] : 1.f; + pp[0] = float2int8(p0q[0] * s0 * scale0); + pp[1] = float2int8(p0q[1] * s1 * scale0); + pp[2] = float2int8(p1q[0] * s0 * scale1); + pp[3] = float2int8(p1q[1] * s1 * scale1); pp += 4; + p0q += 2; + p1q += 2; + if (psq) + psq += 2; kk += 2; } if (kk < max_kk) { - const int k = k0 + kk; - const float s = input_scale_ptr ? input_scale_ptr[k] : 1.f; - pp[0] = float2int8(p0[k] * s * scale0); - pp[1] = float2int8(p1[k] * s * scale1); + const float s = psq ? *psq : 1.f; + pp[0] = float2int8(*p0q * s * scale0); + pp[1] = float2int8(*p1q * s * scale1); pp += 2; } } } #endif // __loongarch_sx + for (; ii + 1 < max_ii; ii += 2) + { + const float* p0 = A_data + (i + ii) * A_hstep; + const float* p1 = p0 + A_hstep; + signed char* pp = outptr + ii * out_hstep; + float* pd = descales + ii * descales_hstep; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + const float* p0g = p0 + k0; + const float* p1g = p1 + k0; + const float* sg = input_scale_ptr ? input_scale_ptr + k0 : 0; + float absmax0 = 0.f; + float absmax1 = 0.f; + const float* p0a = p0g; + const float* p1a = p1g; + const float* psa = sg; + for (int kk = 0; kk < max_kk; kk++) + { + const float s = psa ? *psa++ : 1.f; + absmax0 = std::max(absmax0, fabsf(*p0a++ * s)); + absmax1 = std::max(absmax1, fabsf(*p1a++ * s)); + } + pd[0] = absmax0 / 127.f; + pd[1] = absmax1 / 127.f; + pd += 2; + const float scale0 = absmax0 == 0.f ? 0.f : 127.f / absmax0; + const float scale1 = absmax1 == 0.f ? 0.f : 127.f / absmax1; + const float* p0q = p0g; + const float* p1q = p1g; + const float* psq = sg; + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + const float s0 = psq ? psq[0] : 1.f; + const float s1 = psq ? psq[1] : 1.f; + const float s2 = psq ? psq[2] : 1.f; + const float s3 = psq ? psq[3] : 1.f; + pp[0] = float2int8(p0q[0] * s0 * scale0); + pp[1] = float2int8(p0q[1] * s1 * scale0); + pp[2] = float2int8(p0q[2] * s2 * scale0); + pp[3] = float2int8(p0q[3] * s3 * scale0); + pp[4] = float2int8(p1q[0] * s0 * scale1); + pp[5] = float2int8(p1q[1] * s1 * scale1); + pp[6] = float2int8(p1q[2] * s2 * scale1); + pp[7] = float2int8(p1q[3] * s3 * scale1); + pp += 8; + p0q += 4; + p1q += 4; + if (psq) + psq += 4; + } + if (kk + 1 < max_kk) + { + const float s0 = psq ? psq[0] : 1.f; + const float s1 = psq ? psq[1] : 1.f; + pp[0] = float2int8(p0q[0] * s0 * scale0); + pp[1] = float2int8(p0q[1] * s1 * scale0); + pp[2] = float2int8(p1q[0] * s0 * scale1); + pp[3] = float2int8(p1q[1] * s1 * scale1); + pp += 4; + p0q += 2; + p1q += 2; + if (psq) + psq += 2; + kk += 2; + } + if (kk < max_kk) + { + const float s = psq ? *psq : 1.f; + pp[0] = float2int8(*p0q * s * scale0); + pp[1] = float2int8(*p1q * s * scale1); + pp += 2; + } + } + } for (; ii < max_ii; ii++) { - const float* ptrA = (const float*)A + (i + ii) * A_hstep; - signed char* outptr0 = outptr + ii * out_hstep; - float* descale_ptr = descales + ii * descales_hstep; + const float* p0 = A_data + (i + ii) * A_hstep; + signed char* pp = outptr + ii * out_hstep; + float* pd = descales + ii * descales_hstep; for (int g = 0; g < block_count; g++) { const int k0 = g * block_size; const int max_kk = std::min(K - k0, block_size); + const float* p0g = p0 + k0; + const float* sg = input_scale_ptr ? input_scale_ptr + k0 : 0; float absmax = 0.f; + const float* p0a = p0g; + const float* psa = sg; int kk = 0; #if __loongarch_sx const __m128i _abs_mask = __lsx_vreplgr2vr_w(0x7fffffff); @@ -517,99 +895,110 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales __m256 _absmax256 = (__m256)__lasx_xvldi(0); for (; kk + 7 < max_kk; kk += 8) { - __m256 _v = (__m256)__lasx_xvld(ptrA + k0 + kk, 0); - if (input_scale_ptr) - _v = __lasx_xvfmul_s(_v, (__m256)__lasx_xvld(input_scale_ptr + k0 + kk, 0)); + __m256 _v = (__m256)__lasx_xvld(p0a, 0); + if (psa) + _v = __lasx_xvfmul_s(_v, (__m256)__lasx_xvld(psa, 0)); _v = (__m256)__lasx_xvand_v((__m256i)_v, _abs_mask256); _absmax256 = __lasx_xvfmax_s(_absmax256, _v); + p0a += 8; + if (psa) + psa += 8; } absmax = __lasx_reduce_fmax_s(_absmax256); #endif __m128 _absmax128 = __lsx_vreplfr2vr_s(absmax); for (; kk + 3 < max_kk; kk += 4) { - __m128 _v = (__m128)__lsx_vld(ptrA + k0 + kk, 0); - if (input_scale_ptr) - _v = __lsx_vfmul_s(_v, (__m128)__lsx_vld(input_scale_ptr + k0 + kk, 0)); + __m128 _v = (__m128)__lsx_vld(p0a, 0); + if (psa) + _v = __lsx_vfmul_s(_v, (__m128)__lsx_vld(psa, 0)); _v = (__m128)__lsx_vand_v((__m128i)_v, _abs_mask); _absmax128 = __lsx_vfmax_s(_absmax128, _v); + p0a += 4; + if (psa) + psa += 4; } absmax = __lsx_reduce_fmax_s(_absmax128); #endif for (; kk < max_kk; kk++) { - const int k = k0 + kk; - float v = ptrA[k]; - if (input_scale_ptr) - v *= input_scale_ptr[k]; + float v = *p0a++; + if (psa) + v *= *psa++; absmax = std::max(absmax, fabsf(v)); } if (absmax == 0.f) { - descale_ptr[g] = 0.f; - for (int k = 0; k < max_kk; k++) - outptr0[k0 + k] = 0; + *pd++ = 0.f; + for (int kk = 0; kk < max_kk; kk++) + *pp++ = 0; continue; } - volatile double scale_fp64 = 127.0 / (double)absmax; - const float scale = (float)scale_fp64; - descale_ptr[g] = absmax / 127.f; + const float scale = 127.f / absmax; + *pd++ = absmax / 127.f; + const float* p0q = p0g; + const float* psq = sg; kk = 0; #if __loongarch_sx #if __loongarch_asx - const __m256 _scale256 = (__m256)__lasx_xvreplfr2vr_s(scale); + __m256 _scale256 = (__m256)__lasx_xvreplfr2vr_s(scale); for (; kk + 7 < max_kk; kk += 8) { - __m256 _v = (__m256)__lasx_xvld(ptrA + k0 + kk, 0); - if (input_scale_ptr) - _v = __lasx_xvfmul_s(_v, (__m256)__lasx_xvld(input_scale_ptr + k0 + kk, 0)); + __m256 _v = (__m256)__lasx_xvld(p0q, 0); + if (psq) + _v = __lasx_xvfmul_s(_v, (__m256)__lasx_xvld(psq, 0)); _v = __lasx_xvfmul_s(_v, _scale256); - __lsx_vstelm_d(__lasx_extract_128_lo(float2int8(_v)), outptr0 + k0 + kk, 0, 0); + __lsx_vstelm_d(__lasx_extract_128_lo(float2int8(_v)), pp, 0, 0); + pp += 8; + p0q += 8; + if (psq) + psq += 8; } #endif - const __m128 _scale128 = __lsx_vreplfr2vr_s(scale); + __m128 _scale128 = __lsx_vreplfr2vr_s(scale); for (; kk + 3 < max_kk; kk += 4) { - __m128 _v = (__m128)__lsx_vld(ptrA + k0 + kk, 0); - if (input_scale_ptr) - _v = __lsx_vfmul_s(_v, (__m128)__lsx_vld(input_scale_ptr + k0 + kk, 0)); + __m128 _v = (__m128)__lsx_vld(p0q, 0); + if (psq) + _v = __lsx_vfmul_s(_v, (__m128)__lsx_vld(psq, 0)); _v = __lsx_vfmul_s(_v, _scale128); - __lsx_vstelm_w(float2int8(_v), outptr0 + k0 + kk, 0, 0); + __lsx_vstelm_w(float2int8(_v), pp, 0, 0); + pp += 4; + p0q += 4; + if (psq) + psq += 4; } #endif for (; kk < max_kk; kk++) { - const int k = k0 + kk; - float v = ptrA[k]; - if (input_scale_ptr) - { - v *= input_scale_ptr[k]; - // preserve multiplication order for consistent rounding - asm volatile("" - : "+f"(v)); - } - outptr0[k] = float2int8(v * scale); + float v = *p0q++; + if (psq) + v *= *psq++; + *pp++ = float2int8(v * scale); } } } } -static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr) +static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int k, int max_kk, int block_size, const float* input_scale_ptr) { signed char* outptr = AT_tile; const int out_hstep = AT_tile.w; float* descales = AT_descales_tile; const int descales_hstep = AT_descales_tile.w; + const int K = max_kk; const int block_count = (K + block_size - 1) / block_size; const size_t A_hstep = A.dims == 3 ? A.cstep : (size_t)A.w; + const float* A_data = (const float*)A + (size_t)k * A_hstep; + input_scale_ptr = input_scale_ptr ? input_scale_ptr + k : 0; int ii = 0; #if __loongarch_sx for (; ii + 7 < max_ii; ii += 8) { - const float* ptrA = (const float*)A + i + ii; + const float* ptrA = A_data + i + ii; signed char* pp = outptr + ii * out_hstep; float* pd = descales + ii * descales_hstep; @@ -617,23 +1006,26 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A { const int k0 = g * block_size; const int max_kk = std::min(K - k0, block_size); + const float* p0g = ptrA + (size_t)k0 * A_hstep; + const float* sg = input_scale_ptr ? input_scale_ptr + k0 : 0; const __m128i _abs_mask = __lsx_vreplgr2vr_w(0x7fffffff); __m128 _absmax0 = (__m128)__lsx_vldi(0); __m128 _absmax1 = (__m128)__lsx_vldi(0); + const float* p0a = p0g; + const float* psa = sg; for (int kk = 0; kk < max_kk; kk++) { - const int k = k0 + kk; - const float* p = ptrA + (size_t)k * A_hstep; - __m128 _v0 = (__m128)__lsx_vld(p, 0); - __m128 _v1 = (__m128)__lsx_vld(p + 4, 0); - if (input_scale_ptr) + __m128 _v0 = (__m128)__lsx_vld(p0a, 0); + __m128 _v1 = (__m128)__lsx_vld(p0a + 4, 0); + if (psa) { - const __m128 _s = __lsx_vreplfr2vr_s(input_scale_ptr[k]); + __m128 _s = __lsx_vreplfr2vr_s(*psa++); _v0 = __lsx_vfmul_s(_v0, _s); _v1 = __lsx_vfmul_s(_v1, _s); } _absmax0 = __lsx_vfmax_s(_absmax0, (__m128)__lsx_vand_v((__m128i)_v0, _abs_mask)); _absmax1 = __lsx_vfmax_s(_absmax1, (__m128)__lsx_vand_v((__m128i)_v1, _abs_mask)); + p0a += A_hstep; } float absmax[8]; @@ -658,40 +1050,33 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A pd[7] = absmax7 / 127.f; pd += 8; - volatile double scale0_fp64 = absmax0 == 0.f ? 0.0 : 127.0 / (double)absmax0; - volatile double scale1_fp64 = absmax1 == 0.f ? 0.0 : 127.0 / (double)absmax1; - volatile double scale2_fp64 = absmax2 == 0.f ? 0.0 : 127.0 / (double)absmax2; - volatile double scale3_fp64 = absmax3 == 0.f ? 0.0 : 127.0 / (double)absmax3; - volatile double scale4_fp64 = absmax4 == 0.f ? 0.0 : 127.0 / (double)absmax4; - volatile double scale5_fp64 = absmax5 == 0.f ? 0.0 : 127.0 / (double)absmax5; - volatile double scale6_fp64 = absmax6 == 0.f ? 0.0 : 127.0 / (double)absmax6; - volatile double scale7_fp64 = absmax7 == 0.f ? 0.0 : 127.0 / (double)absmax7; - const float scale0 = (float)scale0_fp64; - const float scale1 = (float)scale1_fp64; - const float scale2 = (float)scale2_fp64; - const float scale3 = (float)scale3_fp64; - const float scale4 = (float)scale4_fp64; - const float scale5 = (float)scale5_fp64; - const float scale6 = (float)scale6_fp64; - const float scale7 = (float)scale7_fp64; - const __m128 _scale0 = __lsx_vreplfr2vr_s(scale0); - const __m128 _scale1 = __lsx_vreplfr2vr_s(scale1); - const __m128 _scale2 = __lsx_vreplfr2vr_s(scale2); - const __m128 _scale3 = __lsx_vreplfr2vr_s(scale3); - const __m128 _scale4 = __lsx_vreplfr2vr_s(scale4); - const __m128 _scale5 = __lsx_vreplfr2vr_s(scale5); - const __m128 _scale6 = __lsx_vreplfr2vr_s(scale6); - const __m128 _scale7 = __lsx_vreplfr2vr_s(scale7); + const float scale0 = absmax0 == 0.f ? 0.f : 127.f / absmax0; + const float scale1 = absmax1 == 0.f ? 0.f : 127.f / absmax1; + const float scale2 = absmax2 == 0.f ? 0.f : 127.f / absmax2; + const float scale3 = absmax3 == 0.f ? 0.f : 127.f / absmax3; + const float scale4 = absmax4 == 0.f ? 0.f : 127.f / absmax4; + const float scale5 = absmax5 == 0.f ? 0.f : 127.f / absmax5; + const float scale6 = absmax6 == 0.f ? 0.f : 127.f / absmax6; + const float scale7 = absmax7 == 0.f ? 0.f : 127.f / absmax7; + __m128 _scale0 = __lsx_vreplfr2vr_s(scale0); + __m128 _scale1 = __lsx_vreplfr2vr_s(scale1); + __m128 _scale2 = __lsx_vreplfr2vr_s(scale2); + __m128 _scale3 = __lsx_vreplfr2vr_s(scale3); + __m128 _scale4 = __lsx_vreplfr2vr_s(scale4); + __m128 _scale5 = __lsx_vreplfr2vr_s(scale5); + __m128 _scale6 = __lsx_vreplfr2vr_s(scale6); + __m128 _scale7 = __lsx_vreplfr2vr_s(scale7); const float scales0[4] = {scale0, scale1, scale2, scale3}; const float scales1[4] = {scale4, scale5, scale6, scale7}; - const __m128 _scales0 = (__m128)__lsx_vld(scales0, 0); - const __m128 _scales1 = (__m128)__lsx_vld(scales1, 0); + __m128 _scales0 = (__m128)__lsx_vld(scales0, 0); + __m128 _scales1 = (__m128)__lsx_vld(scales1, 0); + const float* p0q = p0g; + const float* psq = sg; int kk = 0; for (; kk + 3 < max_kk; kk += 4) { - const int k = k0 + kk; - const float* p0 = ptrA + (size_t)k * A_hstep; + const float* p0 = p0q; const float* p1 = p0 + A_hstep; const float* p2 = p1 + A_hstep; const float* p3 = p2 + A_hstep; @@ -703,16 +1088,20 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A __m128 _p5 = (__m128)__lsx_vld(p1 + 4, 0); __m128 _p6 = (__m128)__lsx_vld(p2 + 4, 0); __m128 _p7 = (__m128)__lsx_vld(p3 + 4, 0); - if (input_scale_ptr) - { - _p0 = __lsx_vfmul_s(_p0, __lsx_vreplfr2vr_s(input_scale_ptr[k])); - _p1 = __lsx_vfmul_s(_p1, __lsx_vreplfr2vr_s(input_scale_ptr[k + 1])); - _p2 = __lsx_vfmul_s(_p2, __lsx_vreplfr2vr_s(input_scale_ptr[k + 2])); - _p3 = __lsx_vfmul_s(_p3, __lsx_vreplfr2vr_s(input_scale_ptr[k + 3])); - _p4 = __lsx_vfmul_s(_p4, __lsx_vreplfr2vr_s(input_scale_ptr[k])); - _p5 = __lsx_vfmul_s(_p5, __lsx_vreplfr2vr_s(input_scale_ptr[k + 1])); - _p6 = __lsx_vfmul_s(_p6, __lsx_vreplfr2vr_s(input_scale_ptr[k + 2])); - _p7 = __lsx_vfmul_s(_p7, __lsx_vreplfr2vr_s(input_scale_ptr[k + 3])); + if (psq) + { + __m128 _s0 = __lsx_vreplfr2vr_s(psq[0]); + __m128 _s1 = __lsx_vreplfr2vr_s(psq[1]); + __m128 _s2 = __lsx_vreplfr2vr_s(psq[2]); + __m128 _s3 = __lsx_vreplfr2vr_s(psq[3]); + _p0 = __lsx_vfmul_s(_p0, _s0); + _p1 = __lsx_vfmul_s(_p1, _s1); + _p2 = __lsx_vfmul_s(_p2, _s2); + _p3 = __lsx_vfmul_s(_p3, _s3); + _p4 = __lsx_vfmul_s(_p4, _s0); + _p5 = __lsx_vfmul_s(_p5, _s1); + _p6 = __lsx_vfmul_s(_p6, _s2); + _p7 = __lsx_vfmul_s(_p7, _s3); } transpose4x4_ps(_p0, _p1, _p2, _p3); transpose4x4_ps(_p4, _p5, _p6, _p7); @@ -721,16 +1110,17 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A *((int64_t*)(pp + 16)) = float2int8(__lsx_vfmul_s(_p4, _scale4), __lsx_vfmul_s(_p5, _scale5)); *((int64_t*)(pp + 24)) = float2int8(__lsx_vfmul_s(_p6, _scale6), __lsx_vfmul_s(_p7, _scale7)); pp += 32; + p0q += (size_t)4 * A_hstep; + if (psq) + psq += 4; } for (; kk < max_kk; kk++) { - const int k = k0 + kk; - const float* p = ptrA + (size_t)k * A_hstep; - __m128 _p0 = (__m128)__lsx_vld(p, 0); - __m128 _p1 = (__m128)__lsx_vld(p + 4, 0); - if (input_scale_ptr) + __m128 _p0 = (__m128)__lsx_vld(p0q, 0); + __m128 _p1 = (__m128)__lsx_vld(p0q + 4, 0); + if (psq) { - const __m128 _s = __lsx_vreplfr2vr_s(input_scale_ptr[k]); + __m128 _s = __lsx_vreplfr2vr_s(*psq++); _p0 = __lsx_vfmul_s(_p0, _s); _p1 = __lsx_vfmul_s(_p1, _s); } @@ -738,12 +1128,13 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A _p1 = __lsx_vfmul_s(_p1, _scales1); *((int64_t*)pp) = float2int8(_p0, _p1); pp += 8; + p0q += A_hstep; } } } for (; ii + 3 < max_ii; ii += 4) { - const float* ptrA = (const float*)A + i + ii; + const float* ptrA = A_data + i + ii; signed char* pp = outptr + ii * out_hstep; float* pd = descales + ii * descales_hstep; @@ -751,17 +1142,21 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A { const int k0 = g * block_size; const int max_kk = std::min(K - k0, block_size); + const float* p0g = ptrA + (size_t)k0 * A_hstep; + const float* sg = input_scale_ptr ? input_scale_ptr + k0 : 0; const __m128i _abs_mask = __lsx_vreplgr2vr_w(0x7fffffff); __m128 _absmax = (__m128)__lsx_vldi(0); + const float* p0a = p0g; + const float* psa = sg; for (int kk = 0; kk < max_kk; kk++) { - const int k = k0 + kk; - __m128 _v = (__m128)__lsx_vld(ptrA + (size_t)k * A_hstep, 0); - if (input_scale_ptr) - _v = __lsx_vfmul_s(_v, __lsx_vreplfr2vr_s(input_scale_ptr[k])); + __m128 _v = (__m128)__lsx_vld(p0a, 0); + if (psa) + _v = __lsx_vfmul_s(_v, __lsx_vreplfr2vr_s(*psa++)); _v = (__m128)__lsx_vand_v((__m128i)_v, _abs_mask); _absmax = __lsx_vfmax_s(_absmax, _v); + p0a += A_hstep; } float absmax[4]; @@ -772,23 +1167,20 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A pd[3] = absmax[3] / 127.f; pd += 4; - volatile double scale0_fp64 = absmax[0] == 0.f ? 0.0 : 127.0 / (double)absmax[0]; - volatile double scale1_fp64 = absmax[1] == 0.f ? 0.0 : 127.0 / (double)absmax[1]; - volatile double scale2_fp64 = absmax[2] == 0.f ? 0.0 : 127.0 / (double)absmax[2]; - volatile double scale3_fp64 = absmax[3] == 0.f ? 0.0 : 127.0 / (double)absmax[3]; const float scales[4] = { - (float)scale0_fp64, - (float)scale1_fp64, - (float)scale2_fp64, - (float)scale3_fp64 + absmax[0] == 0.f ? 0.f : 127.f / absmax[0], + absmax[1] == 0.f ? 0.f : 127.f / absmax[1], + absmax[2] == 0.f ? 0.f : 127.f / absmax[2], + absmax[3] == 0.f ? 0.f : 127.f / absmax[3] }; - const __m128 _scale = (__m128)__lsx_vld(scales, 0); + __m128 _scale = (__m128)__lsx_vld(scales, 0); + const float* p0q = p0g; + const float* psq = sg; int kk = 0; for (; kk + 3 < max_kk; kk += 4) { - const int k = k0 + kk; - const float* p0 = ptrA + (size_t)k * A_hstep; + const float* p0 = p0q; const float* p1 = p0 + A_hstep; const float* p2 = p1 + A_hstep; const float* p3 = p2 + A_hstep; @@ -796,27 +1188,29 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A __m128 _v1 = (__m128)__lsx_vld(p1, 0); __m128 _v2 = (__m128)__lsx_vld(p2, 0); __m128 _v3 = (__m128)__lsx_vld(p3, 0); - if (input_scale_ptr) + if (psq) { - _v0 = __lsx_vfmul_s(_v0, __lsx_vreplfr2vr_s(input_scale_ptr[k])); - _v1 = __lsx_vfmul_s(_v1, __lsx_vreplfr2vr_s(input_scale_ptr[k + 1])); - _v2 = __lsx_vfmul_s(_v2, __lsx_vreplfr2vr_s(input_scale_ptr[k + 2])); - _v3 = __lsx_vfmul_s(_v3, __lsx_vreplfr2vr_s(input_scale_ptr[k + 3])); + _v0 = __lsx_vfmul_s(_v0, __lsx_vreplfr2vr_s(psq[0])); + _v1 = __lsx_vfmul_s(_v1, __lsx_vreplfr2vr_s(psq[1])); + _v2 = __lsx_vfmul_s(_v2, __lsx_vreplfr2vr_s(psq[2])); + _v3 = __lsx_vfmul_s(_v3, __lsx_vreplfr2vr_s(psq[3])); } transpose4x4_ps(_v0, _v1, _v2, _v3); *((int64_t*)pp) = float2int8(__lsx_vfmul_s(_v0, __lsx_vreplfr2vr_s(scales[0])), __lsx_vfmul_s(_v1, __lsx_vreplfr2vr_s(scales[1]))); *((int64_t*)(pp + 8)) = float2int8(__lsx_vfmul_s(_v2, __lsx_vreplfr2vr_s(scales[2])), __lsx_vfmul_s(_v3, __lsx_vreplfr2vr_s(scales[3]))); pp += 16; + p0q += (size_t)4 * A_hstep; + if (psq) + psq += 4; } if (kk + 1 < max_kk) { - const int k = k0 + kk; - __m128 _v0 = (__m128)__lsx_vld(ptrA + (size_t)k * A_hstep, 0); - __m128 _v1 = (__m128)__lsx_vld(ptrA + (size_t)(k + 1) * A_hstep, 0); - if (input_scale_ptr) + __m128 _v0 = (__m128)__lsx_vld(p0q, 0); + __m128 _v1 = (__m128)__lsx_vld(p0q + A_hstep, 0); + if (psq) { - _v0 = __lsx_vfmul_s(_v0, __lsx_vreplfr2vr_s(input_scale_ptr[k])); - _v1 = __lsx_vfmul_s(_v1, __lsx_vreplfr2vr_s(input_scale_ptr[k + 1])); + _v0 = __lsx_vfmul_s(_v0, __lsx_vreplfr2vr_s(psq[0])); + _v1 = __lsx_vfmul_s(_v1, __lsx_vreplfr2vr_s(psq[1])); } const int q0 = __lsx_vpickve2gr_w(float2int8(__lsx_vfmul_s(_v0, _scale)), 0); const int q1 = __lsx_vpickve2gr_w(float2int8(__lsx_vfmul_s(_v1, _scale)), 0); @@ -829,14 +1223,16 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A pp[6] = (signed char)(q0 >> 24); pp[7] = (signed char)(q1 >> 24); pp += 8; + p0q += (size_t)2 * A_hstep; + if (psq) + psq += 2; kk += 2; } if (kk < max_kk) { - const int k = k0 + kk; - __m128 _v = (__m128)__lsx_vld(ptrA + (size_t)k * A_hstep, 0); - if (input_scale_ptr) - _v = __lsx_vfmul_s(_v, __lsx_vreplfr2vr_s(input_scale_ptr[k])); + __m128 _v = (__m128)__lsx_vld(p0q, 0); + if (psq) + _v = __lsx_vfmul_s(_v, __lsx_vreplfr2vr_s(*psq)); const int q = __lsx_vpickve2gr_w(float2int8(__lsx_vfmul_s(_v, _scale)), 0); pp[0] = (signed char)q; pp[1] = (signed char)(q >> 8); @@ -848,7 +1244,7 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A } for (; ii + 1 < max_ii; ii += 2) { - const float* ptrA = (const float*)A + i + ii; + const float* ptrA = A_data + i + ii; signed char* pp = outptr + ii * out_hstep; float* pd = descales + ii * descales_hstep; const __m128i _abs_mask = __lsx_vreplgr2vr_w(0x7fffffff); @@ -857,36 +1253,39 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A { const int k0 = g * block_size; const int max_kk = std::min(K - k0, block_size); + const float* p0g = ptrA + (size_t)k0 * A_hstep; + const float* sg = input_scale_ptr ? input_scale_ptr + k0 : 0; __m128 _absmax = (__m128)__lsx_vldi(0); + const float* p0a = p0g; + const float* psa = sg; for (int kk = 0; kk < max_kk; kk++) { - const int k = k0 + kk; - __m128 _v = (__m128)__lsx_vldrepl_d(ptrA + (size_t)k * A_hstep, 0); - if (input_scale_ptr) - _v = __lsx_vfmul_s(_v, __lsx_vreplfr2vr_s(input_scale_ptr[k])); + __m128 _v = (__m128)__lsx_vldrepl_d(p0a, 0); + if (psa) + _v = __lsx_vfmul_s(_v, __lsx_vreplfr2vr_s(*psa++)); _absmax = __lsx_vfmax_s(_absmax, (__m128)__lsx_vand_v((__m128i)_v, _abs_mask)); + p0a += A_hstep; } float absmax[2]; __lsx_vstelm_d((__m128i)_absmax, absmax, 0, 0); pd[0] = absmax[0] / 127.f; pd[1] = absmax[1] / 127.f; pd += 2; - volatile double scale0_fp64 = absmax[0] == 0.f ? 0.0 : 127.0 / (double)absmax[0]; - volatile double scale1_fp64 = absmax[1] == 0.f ? 0.0 : 127.0 / (double)absmax[1]; - const float scale0 = (float)scale0_fp64; - const float scale1 = (float)scale1_fp64; + const float scale0 = absmax[0] == 0.f ? 0.f : 127.f / absmax[0]; + const float scale1 = absmax[1] == 0.f ? 0.f : 127.f / absmax[1]; + const float* p0q = p0g; + const float* psq = sg; int kk = 0; for (; kk + 3 < max_kk; kk += 4) { - const int k = k0 + kk; - const float* p0 = ptrA + (size_t)k * A_hstep; + const float* p0 = p0q; const float* p1 = p0 + A_hstep; const float* p2 = p1 + A_hstep; const float* p3 = p2 + A_hstep; - const float s0 = input_scale_ptr ? input_scale_ptr[k] : 1.f; - const float s1 = input_scale_ptr ? input_scale_ptr[k + 1] : 1.f; - const float s2 = input_scale_ptr ? input_scale_ptr[k + 2] : 1.f; - const float s3 = input_scale_ptr ? input_scale_ptr[k + 3] : 1.f; + const float s0 = psq ? psq[0] : 1.f; + const float s1 = psq ? psq[1] : 1.f; + const float s2 = psq ? psq[2] : 1.f; + const float s3 = psq ? psq[3] : 1.f; pp[0] = float2int8(p0[0] * s0 * scale0); pp[1] = float2int8(p1[0] * s1 * scale0); pp[2] = float2int8(p2[0] * s2 * scale0); @@ -896,92 +1295,178 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A pp[6] = float2int8(p2[1] * s2 * scale1); pp[7] = float2int8(p3[1] * s3 * scale1); pp += 8; + p0q += (size_t)4 * A_hstep; + if (psq) + psq += 4; } if (kk + 1 < max_kk) { - const int k = k0 + kk; - const float* p0 = ptrA + (size_t)k * A_hstep; + const float* p0 = p0q; const float* p1 = p0 + A_hstep; - const float s0 = input_scale_ptr ? input_scale_ptr[k] : 1.f; - const float s1 = input_scale_ptr ? input_scale_ptr[k + 1] : 1.f; + const float s0 = psq ? psq[0] : 1.f; + const float s1 = psq ? psq[1] : 1.f; pp[0] = float2int8(p0[0] * s0 * scale0); pp[1] = float2int8(p1[0] * s1 * scale0); pp[2] = float2int8(p0[1] * s0 * scale1); pp[3] = float2int8(p1[1] * s1 * scale1); pp += 4; + p0q += (size_t)2 * A_hstep; + if (psq) + psq += 2; kk += 2; } if (kk < max_kk) { - const int k = k0 + kk; - const float* p = ptrA + (size_t)k * A_hstep; - const float s = input_scale_ptr ? input_scale_ptr[k] : 1.f; - pp[0] = float2int8(p[0] * s * scale0); - pp[1] = float2int8(p[1] * s * scale1); + const float s = psq ? *psq : 1.f; + pp[0] = float2int8(p0q[0] * s * scale0); + pp[1] = float2int8(p0q[1] * s * scale1); pp += 2; } } } #endif // __loongarch_sx + for (; ii + 1 < max_ii; ii += 2) + { + const float* ptrA = A_data + i + ii; + signed char* pp = outptr + ii * out_hstep; + float* pd = descales + ii * descales_hstep; + + for (int g = 0; g < block_count; g++) + { + const int k0 = g * block_size; + const int max_kk = std::min(K - k0, block_size); + const float* p0g = ptrA + (size_t)k0 * A_hstep; + const float* sg = input_scale_ptr ? input_scale_ptr + k0 : 0; + float absmax0 = 0.f; + float absmax1 = 0.f; + const float* p0a = p0g; + const float* psa = sg; + for (int kk = 0; kk < max_kk; kk++) + { + const float s = psa ? *psa++ : 1.f; + absmax0 = std::max(absmax0, fabsf(p0a[0] * s)); + absmax1 = std::max(absmax1, fabsf(p0a[1] * s)); + p0a += A_hstep; + } + pd[0] = absmax0 / 127.f; + pd[1] = absmax1 / 127.f; + pd += 2; + const float scale0 = absmax0 == 0.f ? 0.f : 127.f / absmax0; + const float scale1 = absmax1 == 0.f ? 0.f : 127.f / absmax1; + const float* p0q = p0g; + const float* psq = sg; + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + const float* p0 = p0q; + const float* p1 = p0 + A_hstep; + const float* p2 = p1 + A_hstep; + const float* p3 = p2 + A_hstep; + const float s0 = psq ? psq[0] : 1.f; + const float s1 = psq ? psq[1] : 1.f; + const float s2 = psq ? psq[2] : 1.f; + const float s3 = psq ? psq[3] : 1.f; + pp[0] = float2int8(p0[0] * s0 * scale0); + pp[1] = float2int8(p1[0] * s1 * scale0); + pp[2] = float2int8(p2[0] * s2 * scale0); + pp[3] = float2int8(p3[0] * s3 * scale0); + pp[4] = float2int8(p0[1] * s0 * scale1); + pp[5] = float2int8(p1[1] * s1 * scale1); + pp[6] = float2int8(p2[1] * s2 * scale1); + pp[7] = float2int8(p3[1] * s3 * scale1); + pp += 8; + p0q += (size_t)4 * A_hstep; + if (psq) + psq += 4; + } + if (kk + 1 < max_kk) + { + const float* p0 = p0q; + const float* p1 = p0 + A_hstep; + const float s0 = psq ? psq[0] : 1.f; + const float s1 = psq ? psq[1] : 1.f; + pp[0] = float2int8(p0[0] * s0 * scale0); + pp[1] = float2int8(p1[0] * s1 * scale0); + pp[2] = float2int8(p0[1] * s0 * scale1); + pp[3] = float2int8(p1[1] * s1 * scale1); + pp += 4; + p0q += (size_t)2 * A_hstep; + if (psq) + psq += 2; + kk += 2; + } + if (kk < max_kk) + { + const float s = psq ? *psq : 1.f; + pp[0] = float2int8(p0q[0] * s * scale0); + pp[1] = float2int8(p0q[1] * s * scale1); + pp += 2; + } + } + } for (; ii < max_ii; ii++) { - const float* ptrA = (const float*)A + i + ii; - signed char* outptr0 = outptr + ii * out_hstep; - float* descale_ptr = descales + ii * descales_hstep; + const float* p0 = A_data + i + ii; + signed char* pp = outptr + ii * out_hstep; + float* pd = descales + ii * descales_hstep; for (int g = 0; g < block_count; g++) { const int k0 = g * block_size; const int max_kk = std::min(K - k0, block_size); + const float* p0g = p0 + (size_t)k0 * A_hstep; + const float* sg = input_scale_ptr ? input_scale_ptr + k0 : 0; float absmax = 0.f; + const float* p0a = p0g; + const float* psa = sg; for (int kk = 0; kk < max_kk; kk++) { - const int k = k0 + kk; - float v = ptrA[(size_t)k * A_hstep]; - if (input_scale_ptr) - v *= input_scale_ptr[k]; + float v = *p0a; + if (psa) + v *= *psa++; absmax = std::max(absmax, fabsf(v)); + p0a += A_hstep; } if (absmax == 0.f) { - descale_ptr[g] = 0.f; - for (int k = 0; k < max_kk; k++) - outptr0[k0 + k] = 0; + *pd++ = 0.f; + for (int kk = 0; kk < max_kk; kk++) + *pp++ = 0; continue; } - volatile double scale_fp64 = 127.0 / (double)absmax; - const float scale = (float)scale_fp64; - descale_ptr[g] = absmax / 127.f; + const float scale = 127.f / absmax; + *pd++ = absmax / 127.f; + const float* p0q = p0g; + const float* psq = sg; for (int kk = 0; kk < max_kk; kk++) { - const int k = k0 + kk; - float v = ptrA[(size_t)k * A_hstep]; - if (input_scale_ptr) - { - v *= input_scale_ptr[k]; - // preserve multiplication order for consistent rounding - asm volatile("" - : "+f"(v)); - } - outptr0[k] = float2int8(v * scale); + float v = *p0q; + if (psq) + v *= *psq++; + *pp++ = float2int8(v * scale); + p0q += A_hstep; } } } } -static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int block_size) +static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int full_K, int k0, int max_kk0, int block_size) { const signed char* pAT = AT_tile; - const int A_hstep = AT_tile.w; + const int A_hstep = max_kk0; const float* pAT_descales = AT_descales_tile; - const int A_descales_hstep = AT_descales_tile.w; + const int A_descales_hstep = (max_kk0 + block_size - 1) / block_size; const signed char* pBT = BT_tile; const float* pBT_descales = BT_descales_tile; float* outptr = topT_tile; + const int K = max_kk0; + const int block_count = (full_K + block_size - 1) / block_size; + const int block_start = k0 / block_size; + const int tile_blocks = (max_kk0 + block_size - 1) / block_size; int ii = 0; #if __loongarch_sx @@ -994,6 +1479,16 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #if __loongarch_asx for (; jj + 7 < max_jj; jj += 8) { + pB += (size_t)8 * k0; + pB_descales += (size_t)8 * block_start; + __m256 _out0 = k0 == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr, 0); + __m256 _out1 = k0 == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr + 8, 0); + __m256 _out2 = k0 == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr + 16, 0); + __m256 _out3 = k0 == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr + 24, 0); + __m256 _out4 = k0 == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr + 32, 0); + __m256 _out5 = k0 == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr + 40, 0); + __m256 _out6 = k0 == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr + 48, 0); + __m256 _out7 = k0 == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr + 56, 0); const signed char* pA = pAT; const float* pA_descales = pAT_descales; for (int k = 0; k < K; k += block_size) @@ -1019,28 +1514,28 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de __m256i _s0 = __lasx_xvmulwev_h_b(_pA, _pB0); __m256i _s1 = __lasx_xvmulwev_h_b(_pA1, _pB0); - __m256i _s2 = __lasx_xvmulwev_h_b(_pA, _pB1); - __m256i _s3 = __lasx_xvmulwev_h_b(_pA1, _pB1); - __m256i _s4 = __lasx_xvmulwev_h_b(_pA2, _pB0); - __m256i _s5 = __lasx_xvmulwev_h_b(_pA3, _pB0); - __m256i _s6 = __lasx_xvmulwev_h_b(_pA2, _pB1); - __m256i _s7 = __lasx_xvmulwev_h_b(_pA3, _pB1); _s0 = __lasx_xvmaddwod_h_b(_s0, _pA, _pB0); _s1 = __lasx_xvmaddwod_h_b(_s1, _pA1, _pB0); - _s2 = __lasx_xvmaddwod_h_b(_s2, _pA, _pB1); - _s3 = __lasx_xvmaddwod_h_b(_s3, _pA1, _pB1); - _s4 = __lasx_xvmaddwod_h_b(_s4, _pA2, _pB0); - _s5 = __lasx_xvmaddwod_h_b(_s5, _pA3, _pB0); - _s6 = __lasx_xvmaddwod_h_b(_s6, _pA2, _pB1); - _s7 = __lasx_xvmaddwod_h_b(_s7, _pA3, _pB1); _sum0 = __lasx_xvadd_w(_sum0, __lasx_xvhaddw_w_h(_s0, _s0)); _sum1 = __lasx_xvadd_w(_sum1, __lasx_xvhaddw_w_h(_s1, _s1)); - _sum2 = __lasx_xvadd_w(_sum2, __lasx_xvhaddw_w_h(_s2, _s2)); - _sum3 = __lasx_xvadd_w(_sum3, __lasx_xvhaddw_w_h(_s3, _s3)); - _sum4 = __lasx_xvadd_w(_sum4, __lasx_xvhaddw_w_h(_s4, _s4)); - _sum5 = __lasx_xvadd_w(_sum5, __lasx_xvhaddw_w_h(_s5, _s5)); - _sum6 = __lasx_xvadd_w(_sum6, __lasx_xvhaddw_w_h(_s6, _s6)); - _sum7 = __lasx_xvadd_w(_sum7, __lasx_xvhaddw_w_h(_s7, _s7)); + _s0 = __lasx_xvmulwev_h_b(_pA, _pB1); + _s1 = __lasx_xvmulwev_h_b(_pA1, _pB1); + _s0 = __lasx_xvmaddwod_h_b(_s0, _pA, _pB1); + _s1 = __lasx_xvmaddwod_h_b(_s1, _pA1, _pB1); + _sum2 = __lasx_xvadd_w(_sum2, __lasx_xvhaddw_w_h(_s0, _s0)); + _sum3 = __lasx_xvadd_w(_sum3, __lasx_xvhaddw_w_h(_s1, _s1)); + _s0 = __lasx_xvmulwev_h_b(_pA2, _pB0); + _s1 = __lasx_xvmulwev_h_b(_pA3, _pB0); + _s0 = __lasx_xvmaddwod_h_b(_s0, _pA2, _pB0); + _s1 = __lasx_xvmaddwod_h_b(_s1, _pA3, _pB0); + _sum4 = __lasx_xvadd_w(_sum4, __lasx_xvhaddw_w_h(_s0, _s0)); + _sum5 = __lasx_xvadd_w(_sum5, __lasx_xvhaddw_w_h(_s1, _s1)); + _s0 = __lasx_xvmulwev_h_b(_pA2, _pB1); + _s1 = __lasx_xvmulwev_h_b(_pA3, _pB1); + _s0 = __lasx_xvmaddwod_h_b(_s0, _pA2, _pB1); + _s1 = __lasx_xvmaddwod_h_b(_s1, _pA3, _pB1); + _sum6 = __lasx_xvadd_w(_sum6, __lasx_xvhaddw_w_h(_s0, _s0)); + _sum7 = __lasx_xvadd_w(_sum7, __lasx_xvhaddw_w_h(_s1, _s1)); pB += 32; pA += 32; } @@ -1057,20 +1552,20 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de __m256i _pB1 = __lasx_xvshuf4i_h(_pB0, _LSX_SHUFFLE(1, 0, 3, 2)); __m256i _s0 = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA, _pB0), _pA, _pB0); __m256i _s1 = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA1, _pB0), _pA1, _pB0); - __m256i _s2 = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA, _pB1), _pA, _pB1); - __m256i _s3 = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA1, _pB1), _pA1, _pB1); - __m256i _s4 = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA2, _pB0), _pA2, _pB0); - __m256i _s5 = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA3, _pB0), _pA3, _pB0); - __m256i _s6 = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA2, _pB1), _pA2, _pB1); - __m256i _s7 = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA3, _pB1), _pA3, _pB1); _sum0 = __lasx_xvadd_w(_sum0, __lasx_vext2xv_w_h(_s0)); _sum1 = __lasx_xvadd_w(_sum1, __lasx_vext2xv_w_h(_s1)); - _sum2 = __lasx_xvadd_w(_sum2, __lasx_vext2xv_w_h(_s2)); - _sum3 = __lasx_xvadd_w(_sum3, __lasx_vext2xv_w_h(_s3)); - _sum4 = __lasx_xvadd_w(_sum4, __lasx_vext2xv_w_h(_s4)); - _sum5 = __lasx_xvadd_w(_sum5, __lasx_vext2xv_w_h(_s5)); - _sum6 = __lasx_xvadd_w(_sum6, __lasx_vext2xv_w_h(_s6)); - _sum7 = __lasx_xvadd_w(_sum7, __lasx_vext2xv_w_h(_s7)); + _s0 = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA, _pB1), _pA, _pB1); + _s1 = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA1, _pB1), _pA1, _pB1); + _sum2 = __lasx_xvadd_w(_sum2, __lasx_vext2xv_w_h(_s0)); + _sum3 = __lasx_xvadd_w(_sum3, __lasx_vext2xv_w_h(_s1)); + _s0 = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA2, _pB0), _pA2, _pB0); + _s1 = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA3, _pB0), _pA3, _pB0); + _sum4 = __lasx_xvadd_w(_sum4, __lasx_vext2xv_w_h(_s0)); + _sum5 = __lasx_xvadd_w(_sum5, __lasx_vext2xv_w_h(_s1)); + _s0 = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA2, _pB1), _pA2, _pB1); + _s1 = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA3, _pB1), _pA3, _pB1); + _sum6 = __lasx_xvadd_w(_sum6, __lasx_vext2xv_w_h(_s0)); + _sum7 = __lasx_xvadd_w(_sum7, __lasx_vext2xv_w_h(_s1)); pB += 16; pA += 16; kk += 2; @@ -1087,20 +1582,20 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de __m256i _pB1 = __lasx_xvshuf4i_h(_pB0, _LSX_SHUFFLE(1, 0, 3, 2)); __m256i _s0 = __lasx_xvmul_h(_pA, _pB0); __m256i _s1 = __lasx_xvmul_h(_pA1, _pB0); - __m256i _s2 = __lasx_xvmul_h(_pA, _pB1); - __m256i _s3 = __lasx_xvmul_h(_pA1, _pB1); - __m256i _s4 = __lasx_xvmul_h(_pA2, _pB0); - __m256i _s5 = __lasx_xvmul_h(_pA3, _pB0); - __m256i _s6 = __lasx_xvmul_h(_pA2, _pB1); - __m256i _s7 = __lasx_xvmul_h(_pA3, _pB1); _sum0 = __lasx_xvadd_w(_sum0, __lasx_vext2xv_w_h(_s0)); _sum1 = __lasx_xvadd_w(_sum1, __lasx_vext2xv_w_h(_s1)); - _sum2 = __lasx_xvadd_w(_sum2, __lasx_vext2xv_w_h(_s2)); - _sum3 = __lasx_xvadd_w(_sum3, __lasx_vext2xv_w_h(_s3)); - _sum4 = __lasx_xvadd_w(_sum4, __lasx_vext2xv_w_h(_s4)); - _sum5 = __lasx_xvadd_w(_sum5, __lasx_vext2xv_w_h(_s5)); - _sum6 = __lasx_xvadd_w(_sum6, __lasx_vext2xv_w_h(_s6)); - _sum7 = __lasx_xvadd_w(_sum7, __lasx_vext2xv_w_h(_s7)); + _s0 = __lasx_xvmul_h(_pA, _pB1); + _s1 = __lasx_xvmul_h(_pA1, _pB1); + _sum2 = __lasx_xvadd_w(_sum2, __lasx_vext2xv_w_h(_s0)); + _sum3 = __lasx_xvadd_w(_sum3, __lasx_vext2xv_w_h(_s1)); + _s0 = __lasx_xvmul_h(_pA2, _pB0); + _s1 = __lasx_xvmul_h(_pA3, _pB0); + _sum4 = __lasx_xvadd_w(_sum4, __lasx_vext2xv_w_h(_s0)); + _sum5 = __lasx_xvadd_w(_sum5, __lasx_vext2xv_w_h(_s1)); + _s0 = __lasx_xvmul_h(_pA2, _pB1); + _s1 = __lasx_xvmul_h(_pA3, _pB1); + _sum6 = __lasx_xvadd_w(_sum6, __lasx_vext2xv_w_h(_s0)); + _sum7 = __lasx_xvadd_w(_sum7, __lasx_vext2xv_w_h(_s1)); pB += 8; pA += 8; } @@ -1111,14 +1606,6 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de __m256 _ascale1 = (__m256)__lasx_xvshuf4i_w((__m256i)_ascale, _LSX_SHUFFLE(0, 3, 2, 1)); __m256 _ascale2 = (__m256)__lasx_xvpermi_q((__m256i)_ascale, (__m256i)_ascale, _LSX_SHUFFLE(0, 0, 0, 1)); __m256 _ascale3 = (__m256)__lasx_xvshuf4i_w((__m256i)_ascale2, _LSX_SHUFFLE(0, 3, 2, 1)); - __m256 _out0 = k == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr, 0); - __m256 _out1 = k == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr + 8, 0); - __m256 _out2 = k == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr + 16, 0); - __m256 _out3 = k == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr + 24, 0); - __m256 _out4 = k == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr + 32, 0); - __m256 _out5 = k == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr + 40, 0); - __m256 _out6 = k == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr + 48, 0); - __m256 _out7 = k == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr + 56, 0); _out0 = __lasx_xvfmadd_s((__m256)__lasx_xvffint_s_w(_sum0), __lasx_xvfmul_s(_ascale, _bscale), _out0); _out1 = __lasx_xvfmadd_s((__m256)__lasx_xvffint_s_w(_sum1), __lasx_xvfmul_s(_ascale1, _bscale), _out1); _out2 = __lasx_xvfmadd_s((__m256)__lasx_xvffint_s_w(_sum2), __lasx_xvfmul_s(_ascale, _bscale1), _out2); @@ -1127,22 +1614,34 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de _out5 = __lasx_xvfmadd_s((__m256)__lasx_xvffint_s_w(_sum5), __lasx_xvfmul_s(_ascale3, _bscale), _out5); _out6 = __lasx_xvfmadd_s((__m256)__lasx_xvffint_s_w(_sum6), __lasx_xvfmul_s(_ascale2, _bscale1), _out6); _out7 = __lasx_xvfmadd_s((__m256)__lasx_xvffint_s_w(_sum7), __lasx_xvfmul_s(_ascale3, _bscale1), _out7); - __lasx_xvst(_out0, outptr, 0); - __lasx_xvst(_out1, outptr + 8, 0); - __lasx_xvst(_out2, outptr + 16, 0); - __lasx_xvst(_out3, outptr + 24, 0); - __lasx_xvst(_out4, outptr + 32, 0); - __lasx_xvst(_out5, outptr + 40, 0); - __lasx_xvst(_out6, outptr + 48, 0); - __lasx_xvst(_out7, outptr + 56, 0); pA_descales += 8; pB_descales += 8; } + __lasx_xvst(_out0, outptr, 0); + __lasx_xvst(_out1, outptr + 8, 0); + __lasx_xvst(_out2, outptr + 16, 0); + __lasx_xvst(_out3, outptr + 24, 0); + __lasx_xvst(_out4, outptr + 32, 0); + __lasx_xvst(_out5, outptr + 40, 0); + __lasx_xvst(_out6, outptr + 48, 0); + __lasx_xvst(_out7, outptr + 56, 0); outptr += 64; + pB += (size_t)8 * (full_K - k0 - max_kk0); + pB_descales += (size_t)8 * (block_count - block_start - tile_blocks); } #endif // __loongarch_asx for (; jj + 3 < max_jj; jj += 4) { + pB += (size_t)4 * k0; + pB_descales += (size_t)4 * block_start; + __m128 _out00 = k0 == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr, 0); + __m128 _out01 = k0 == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 4, 0); + __m128 _out10 = k0 == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 8, 0); + __m128 _out11 = k0 == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 12, 0); + __m128 _out20 = k0 == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 16, 0); + __m128 _out21 = k0 == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 20, 0); + __m128 _out30 = k0 == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 24, 0); + __m128 _out31 = k0 == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 28, 0); const signed char* pA = pAT; const float* pA_descales = pAT_descales; for (int k = 0; k < K; k += block_size) @@ -1264,14 +1763,6 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de __m128 _ascale1 = (__m128)__lsx_vld(pA_descales + 4, 0); __m128 _ascale0r = (__m128)__lsx_vshuf4i_w((__m128i)_ascale0, _LSX_SHUFFLE(1, 0, 3, 2)); __m128 _ascale1r = (__m128)__lsx_vshuf4i_w((__m128i)_ascale1, _LSX_SHUFFLE(1, 0, 3, 2)); - __m128 _out00 = k == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr, 0); - __m128 _out01 = k == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 4, 0); - __m128 _out10 = k == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 8, 0); - __m128 _out11 = k == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 12, 0); - __m128 _out20 = k == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 16, 0); - __m128 _out21 = k == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 20, 0); - __m128 _out30 = k == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 24, 0); - __m128 _out31 = k == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 28, 0); _out00 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum00), __lsx_vfmul_s(_ascale0, _bscale), _out00); _out01 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum01), __lsx_vfmul_s(_ascale1, _bscale), _out01); _out10 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum10), __lsx_vfmul_s(_ascale0, _bscaler), _out10); @@ -1280,26 +1771,30 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de _out21 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum21), __lsx_vfmul_s(_ascale1r, _bscale), _out21); _out30 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum30), __lsx_vfmul_s(_ascale0r, _bscaler), _out30); _out31 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum31), __lsx_vfmul_s(_ascale1r, _bscaler), _out31); - __lsx_vst((__m128i)_out00, outptr, 0); - __lsx_vst((__m128i)_out01, outptr + 4, 0); - __lsx_vst((__m128i)_out10, outptr + 8, 0); - __lsx_vst((__m128i)_out11, outptr + 12, 0); - __lsx_vst((__m128i)_out20, outptr + 16, 0); - __lsx_vst((__m128i)_out21, outptr + 20, 0); - __lsx_vst((__m128i)_out30, outptr + 24, 0); - __lsx_vst((__m128i)_out31, outptr + 28, 0); pA_descales += 8; pB_descales += 4; } + __lsx_vst((__m128i)_out00, outptr, 0); + __lsx_vst((__m128i)_out01, outptr + 4, 0); + __lsx_vst((__m128i)_out10, outptr + 8, 0); + __lsx_vst((__m128i)_out11, outptr + 12, 0); + __lsx_vst((__m128i)_out20, outptr + 16, 0); + __lsx_vst((__m128i)_out21, outptr + 20, 0); + __lsx_vst((__m128i)_out30, outptr + 24, 0); + __lsx_vst((__m128i)_out31, outptr + 28, 0); outptr += 32; + pB += (size_t)4 * (full_K - k0 - max_kk0); + pB_descales += (size_t)4 * (block_count - block_start - tile_blocks); } for (; jj + 1 < max_jj; jj += 2) { - __m128 _out00 = (__m128)__lsx_vldi(0); - __m128 _out01 = (__m128)__lsx_vldi(0); - __m128 _out10 = (__m128)__lsx_vldi(0); - __m128 _out11 = (__m128)__lsx_vldi(0); + pB += (size_t)2 * k0; + pB_descales += (size_t)2 * block_start; + __m128 _out00 = k0 == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr, 0); + __m128 _out01 = k0 == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 4, 0); + __m128 _out10 = k0 == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 8, 0); + __m128 _out11 = k0 == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 12, 0); const signed char* pA = pAT; const float* pA_descales = pAT_descales; for (int k = 0; k < K; k += block_size) @@ -1335,7 +1830,7 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de __m128i _pA = __lsx_vilvl_b(_pA8, _pAs); __m128i _pA0 = __lsx_vreplvei_d(_pA, 0); __m128i _pA1 = __lsx_vreplvei_d(_pA, 1); - __m128i _pB = __lsx_vreplgr2vr_w(*(const int*)pB); + __m128i _pB = __lsx_vldrepl_w(pB, 0); __m128i _pB0 = __lsx_vreplvei_h(_pB, 0); __m128i _pB1 = __lsx_vreplvei_h(_pB, 1); __m128i _s0 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB0), _pA0, _pB0); @@ -1356,8 +1851,10 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de _pA = __lsx_vilvl_b(__lsx_vslti_b(_pA, 0), _pA); __m128i _pA0 = __lsx_vreplvei_d(_pA, 0); __m128i _pA1 = __lsx_vreplvei_d(_pA, 1); - __m128i _pB0 = __lsx_vreplgr2vr_h((signed char)pB[0]); - __m128i _pB1 = __lsx_vreplgr2vr_h((signed char)pB[1]); + __m128i _pB = __lsx_vldrepl_h(pB, 0); + _pB = __lsx_vilvl_b(__lsx_vslti_b(_pB, 0), _pB); + __m128i _pB0 = __lsx_vreplvei_h(_pB, 0); + __m128i _pB1 = __lsx_vreplvei_h(_pB, 1); __m128i _s0 = __lsx_vmul_h(_pA0, _pB0); __m128i _s1 = __lsx_vmul_h(_pA1, _pB0); _sum00 = __lsx_vadd_w(_sum00, __lsx_vilvl_h(__lsx_vslti_h(_s0, 0), _s0)); @@ -1385,11 +1882,15 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de __lsx_vst((__m128i)_out10, outptr + 8, 0); __lsx_vst((__m128i)_out11, outptr + 12, 0); outptr += 16; + pB += (size_t)2 * (full_K - k0 - max_kk0); + pB_descales += (size_t)2 * (block_count - block_start - tile_blocks); } for (; jj < max_jj; jj++) { - __m128 _out0 = (__m128)__lsx_vldi(0); - __m128 _out1 = (__m128)__lsx_vldi(0); + pB += k0; + pB_descales += block_start; + __m128 _out0 = k0 == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr, 0); + __m128 _out1 = k0 == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 4, 0); const signed char* pA = pAT; const float* pA_descales = pAT_descales; for (int k = 0; k < K; k += block_size) @@ -1402,7 +1903,7 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de { __m128i _pA0 = __lsx_vld(pA, 0); __m128i _pA1 = __lsx_vld(pA + 16, 0); - __m128i _pB = __lsx_vreplgr2vr_w(*(const int*)pB); + __m128i _pB = __lsx_vldrepl_w(pB, 0); __m128i _s0 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB), _pA0, _pB); __m128i _s1 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA1, _pB), _pA1, _pB); _sum0 = __lsx_vadd_w(_sum0, __lsx_vhaddw_w_h(_s0, _s0)); @@ -1417,7 +1918,7 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de __m128i _pA = __lsx_vilvl_b(_pA8, _pAs); __m128i _pA0 = __lsx_vreplvei_d(_pA, 0); __m128i _pA1 = __lsx_vreplvei_d(_pA, 1); - __m128i _pB = __lsx_vreplgr2vr_h((unsigned char)pB[0] | ((unsigned char)pB[1] << 8)); + __m128i _pB = __lsx_vldrepl_h(pB, 0); __m128i _s0 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB), _pA0, _pB); __m128i _s1 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA1, _pB), _pA1, _pB); _sum0 = __lsx_vadd_w(_sum0, __lsx_vilvl_h(__lsx_vslti_h(_s0, 0), _s0)); @@ -1448,13 +1949,13 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de __lsx_vst((__m128i)_out0, outptr, 0); __lsx_vst((__m128i)_out1, outptr + 4, 0); outptr += 8; + pB += full_K - k0 - max_kk0; + pB_descales += block_count - block_start - tile_blocks; } pAT += K * 8; pAT_descales += (K + block_size - 1) / block_size * 8; } -#endif // __loongarch_sx -#if __loongarch_sx for (; ii + 3 < max_ii; ii += 4) { const signed char* pB = pBT; @@ -1464,18 +1965,18 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #if __loongarch_asx for (; jj + 15 < max_jj; jj += 16) { - const signed char* pB0 = pB; - const signed char* pB1 = pB + 8 * K; - const float* pB_descales0 = pB_descales; - const float* pB_descales1 = pB_descales + 8 * ((K + block_size - 1) / block_size); - __m256 _out00 = (__m256)__lasx_xvldi(0); - __m256 _out01 = (__m256)__lasx_xvldi(0); - __m256 _out10 = (__m256)__lasx_xvldi(0); - __m256 _out11 = (__m256)__lasx_xvldi(0); - __m256 _out20 = (__m256)__lasx_xvldi(0); - __m256 _out21 = (__m256)__lasx_xvldi(0); - __m256 _out30 = (__m256)__lasx_xvldi(0); - __m256 _out31 = (__m256)__lasx_xvldi(0); + const signed char* pB0 = pB + 8 * k0; + const signed char* pB1 = pB + 8 * full_K + 8 * k0; + const float* pB_descales0 = pB_descales + 8 * block_start; + const float* pB_descales1 = pB_descales + 8 * block_count + 8 * block_start; + __m256 _out00 = k0 == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr, 0); + __m256 _out01 = k0 == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr + 8, 0); + __m256 _out10 = k0 == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr + 16, 0); + __m256 _out11 = k0 == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr + 24, 0); + __m256 _out20 = k0 == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr + 32, 0); + __m256 _out21 = k0 == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr + 40, 0); + __m256 _out30 = k0 == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr + 48, 0); + __m256 _out31 = k0 == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr + 56, 0); const signed char* pA = pAT; const float* pA_descales = pAT_descales; @@ -1493,37 +1994,64 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de int kk = 0; for (; kk + 3 < max_kk; kk += 4) { + __m128i _pA4 = __lsx_vld(pA, 0); + __m256i _pA = __lasx_concat_128(_pA4, _pA4); + __m256i _pA1 = __lasx_xvshuf4i_w(_pA, _LSX_SHUFFLE(1, 0, 3, 2)); __m256i _pB0 = __lasx_xvld(pB0, 0); __m256i _pB1 = __lasx_xvld(pB1, 0); - __m128i _pA = __lsx_vld(pA, 0); - __m128i _pA0_128 = __lsx_vreplvei_w(_pA, 0); - __m128i _pA1_128 = __lsx_vreplvei_w(_pA, 1); - __m128i _pA2_128 = __lsx_vreplvei_w(_pA, 2); - __m128i _pA3_128 = __lsx_vreplvei_w(_pA, 3); - __m256i _pA0 = __lasx_concat_128(_pA0_128, _pA0_128); - __m256i _pA1 = __lasx_concat_128(_pA1_128, _pA1_128); - __m256i _pA2 = __lasx_concat_128(_pA2_128, _pA2_128); - __m256i _pA3 = __lasx_concat_128(_pA3_128, _pA3_128); - __m256i _s = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA0, _pB0), _pA0, _pB0); + __m256i _pB0r = __lasx_xvshuf4i_w(_pB0, _LSX_SHUFFLE(0, 3, 2, 1)); + __m256i _pB1r = __lasx_xvshuf4i_w(_pB1, _LSX_SHUFFLE(0, 3, 2, 1)); + __m256i _s = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA, _pB0), _pA, _pB0); _sum00 = __lasx_xvadd_w(_sum00, __lasx_xvhaddw_w_h(_s, _s)); - _s = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA0, _pB1), _pA0, _pB1); - _sum01 = __lasx_xvadd_w(_sum01, __lasx_xvhaddw_w_h(_s, _s)); - _s = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA1, _pB0), _pA1, _pB0); + _s = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA, _pB0r), _pA, _pB0r); _sum10 = __lasx_xvadd_w(_sum10, __lasx_xvhaddw_w_h(_s, _s)); - _s = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA1, _pB1), _pA1, _pB1); - _sum11 = __lasx_xvadd_w(_sum11, __lasx_xvhaddw_w_h(_s, _s)); - _s = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA2, _pB0), _pA2, _pB0); + _s = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA1, _pB0), _pA1, _pB0); _sum20 = __lasx_xvadd_w(_sum20, __lasx_xvhaddw_w_h(_s, _s)); - _s = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA2, _pB1), _pA2, _pB1); - _sum21 = __lasx_xvadd_w(_sum21, __lasx_xvhaddw_w_h(_s, _s)); - _s = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA3, _pB0), _pA3, _pB0); + _s = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA1, _pB0r), _pA1, _pB0r); _sum30 = __lasx_xvadd_w(_sum30, __lasx_xvhaddw_w_h(_s, _s)); - _s = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA3, _pB1), _pA3, _pB1); + _s = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA, _pB1), _pA, _pB1); + _sum01 = __lasx_xvadd_w(_sum01, __lasx_xvhaddw_w_h(_s, _s)); + _s = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA, _pB1r), _pA, _pB1r); + _sum11 = __lasx_xvadd_w(_sum11, __lasx_xvhaddw_w_h(_s, _s)); + _s = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA1, _pB1), _pA1, _pB1); + _sum21 = __lasx_xvadd_w(_sum21, __lasx_xvhaddw_w_h(_s, _s)); + _s = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA1, _pB1r), _pA1, _pB1r); _sum31 = __lasx_xvadd_w(_sum31, __lasx_xvhaddw_w_h(_s, _s)); pB0 += 32; pB1 += 32; pA += 16; } + _sum20 = __lasx_xvshuf4i_w(_sum20, _LSX_SHUFFLE(1, 0, 3, 2)); + _sum30 = __lasx_xvshuf4i_w(_sum30, _LSX_SHUFFLE(1, 0, 3, 2)); + { + __m256i _tmp0 = __lasx_xvilvl_w(_sum10, _sum00); + __m256i _tmp1 = __lasx_xvilvh_w(_sum10, _sum00); + __m256i _tmp2 = __lasx_xvilvl_w(_sum30, _sum20); + __m256i _tmp3 = __lasx_xvilvh_w(_sum30, _sum20); + _sum00 = __lasx_xvilvl_d(_tmp2, _tmp0); + _sum10 = __lasx_xvilvh_d(_tmp2, _tmp0); + _sum20 = __lasx_xvilvl_d(_tmp3, _tmp1); + _sum30 = __lasx_xvilvh_d(_tmp3, _tmp1); + } + _sum10 = __lasx_xvshuf4i_w(_sum10, _LSX_SHUFFLE(2, 1, 0, 3)); + _sum20 = __lasx_xvshuf4i_w(_sum20, _LSX_SHUFFLE(1, 0, 3, 2)); + _sum30 = __lasx_xvshuf4i_w(_sum30, _LSX_SHUFFLE(0, 3, 2, 1)); + + _sum21 = __lasx_xvshuf4i_w(_sum21, _LSX_SHUFFLE(1, 0, 3, 2)); + _sum31 = __lasx_xvshuf4i_w(_sum31, _LSX_SHUFFLE(1, 0, 3, 2)); + { + __m256i _tmp0 = __lasx_xvilvl_w(_sum11, _sum01); + __m256i _tmp1 = __lasx_xvilvh_w(_sum11, _sum01); + __m256i _tmp2 = __lasx_xvilvl_w(_sum31, _sum21); + __m256i _tmp3 = __lasx_xvilvh_w(_sum31, _sum21); + _sum01 = __lasx_xvilvl_d(_tmp2, _tmp0); + _sum11 = __lasx_xvilvh_d(_tmp2, _tmp0); + _sum21 = __lasx_xvilvl_d(_tmp3, _tmp1); + _sum31 = __lasx_xvilvh_d(_tmp3, _tmp1); + } + _sum11 = __lasx_xvshuf4i_w(_sum11, _LSX_SHUFFLE(2, 1, 0, 3)); + _sum21 = __lasx_xvshuf4i_w(_sum21, _LSX_SHUFFLE(1, 0, 3, 2)); + _sum31 = __lasx_xvshuf4i_w(_sum31, _LSX_SHUFFLE(0, 3, 2, 1)); if (kk + 1 < max_kk) { __m128i _pB0 = __lsx_vld(pB0, 0); @@ -1560,22 +2088,27 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de __m128i _pB1 = __lsx_vldrepl_d(pB1, 0); _pB0 = __lsx_vilvl_b(__lsx_vslti_b(_pB0, 0), _pB0); _pB1 = __lsx_vilvl_b(__lsx_vslti_b(_pB1, 0), _pB1); - const int a0123 = __lsx_vpickve2gr_w(__lsx_vldrepl_w(pA, 0), 0); - __m128i _s = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)a0123), _pB0); + __m128i _pA = __lsx_vldrepl_w(pA, 0); + _pA = __lsx_vilvl_b(__lsx_vslti_b(_pA, 0), _pA); + __m128i _pA0 = __lsx_vreplvei_h(_pA, 0); + __m128i _pA1 = __lsx_vreplvei_h(_pA, 1); + __m128i _pA2 = __lsx_vreplvei_h(_pA, 2); + __m128i _pA3 = __lsx_vreplvei_h(_pA, 3); + __m128i _s = __lsx_vmul_h(_pA0, _pB0); _sum00 = __lasx_xvadd_w(_sum00, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); - _s = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)a0123), _pB1); + _s = __lsx_vmul_h(_pA0, _pB1); _sum01 = __lasx_xvadd_w(_sum01, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); - _s = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)(a0123 >> 8)), _pB0); + _s = __lsx_vmul_h(_pA1, _pB0); _sum10 = __lasx_xvadd_w(_sum10, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); - _s = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)(a0123 >> 8)), _pB1); + _s = __lsx_vmul_h(_pA1, _pB1); _sum11 = __lasx_xvadd_w(_sum11, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); - _s = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)(a0123 >> 16)), _pB0); + _s = __lsx_vmul_h(_pA2, _pB0); _sum20 = __lasx_xvadd_w(_sum20, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); - _s = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)(a0123 >> 16)), _pB1); + _s = __lsx_vmul_h(_pA2, _pB1); _sum21 = __lasx_xvadd_w(_sum21, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); - _s = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)(a0123 >> 24)), _pB0); + _s = __lsx_vmul_h(_pA3, _pB0); _sum30 = __lasx_xvadd_w(_sum30, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); - _s = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)(a0123 >> 24)), _pB1); + _s = __lsx_vmul_h(_pA3, _pB1); _sum31 = __lasx_xvadd_w(_sum31, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); pB0 += 8; pB1 += 8; @@ -1601,24 +2134,27 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB_descales1 += 8; } - pB = pB1; - pB_descales = pB_descales1; + pB = pB1 + 8 * (full_K - k0 - max_kk0); + pB_descales = pB_descales1 + 8 * (block_count - block_start - tile_blocks); - __lasx_xvst(_out00, outptr + (ii + 0) * max_jj + jj, 0); - __lasx_xvst(_out01, outptr + (ii + 0) * max_jj + jj + 8, 0); - __lasx_xvst(_out10, outptr + (ii + 1) * max_jj + jj, 0); - __lasx_xvst(_out11, outptr + (ii + 1) * max_jj + jj + 8, 0); - __lasx_xvst(_out20, outptr + (ii + 2) * max_jj + jj, 0); - __lasx_xvst(_out21, outptr + (ii + 2) * max_jj + jj + 8, 0); - __lasx_xvst(_out30, outptr + (ii + 3) * max_jj + jj, 0); - __lasx_xvst(_out31, outptr + (ii + 3) * max_jj + jj + 8, 0); + __lasx_xvst(_out00, outptr, 0); + __lasx_xvst(_out01, outptr + 8, 0); + __lasx_xvst(_out10, outptr + 16, 0); + __lasx_xvst(_out11, outptr + 24, 0); + __lasx_xvst(_out20, outptr + 32, 0); + __lasx_xvst(_out21, outptr + 40, 0); + __lasx_xvst(_out30, outptr + 48, 0); + __lasx_xvst(_out31, outptr + 56, 0); + outptr += 64; } for (; jj + 7 < max_jj; jj += 8) { - __m256 _out0 = (__m256)__lasx_xvldi(0); - __m256 _out1 = (__m256)__lasx_xvldi(0); - __m256 _out2 = (__m256)__lasx_xvldi(0); - __m256 _out3 = (__m256)__lasx_xvldi(0); + pB += (size_t)8 * k0; + pB_descales += (size_t)8 * block_start; + __m256 _out0 = k0 == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr, 0); + __m256 _out1 = k0 == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr + 8, 0); + __m256 _out2 = k0 == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr + 16, 0); + __m256 _out3 = k0 == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr + 24, 0); const signed char* pA = pAT; const float* pA_descales = pAT_descales; @@ -1632,31 +2168,37 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de int kk = 0; for (; kk + 3 < max_kk; kk += 4) { + __m128i _pA4 = __lsx_vld(pA, 0); + __m256i _pA = __lasx_concat_128(_pA4, _pA4); + __m256i _pA1 = __lasx_xvshuf4i_w(_pA, _LSX_SHUFFLE(1, 0, 3, 2)); __m256i _pB = __lasx_xvld(pB, 0); - __m128i _pA = __lsx_vld(pA, 0); - __m128i _pA0_128 = __lsx_vreplvei_w(_pA, 0); - __m128i _pA1_128 = __lsx_vreplvei_w(_pA, 1); - __m128i _pA2_128 = __lsx_vreplvei_w(_pA, 2); - __m128i _pA3_128 = __lsx_vreplvei_w(_pA, 3); - __m256i _pA0 = __lasx_concat_128(_pA0_128, _pA0_128); - __m256i _pA1 = __lasx_concat_128(_pA1_128, _pA1_128); - __m256i _pA2 = __lasx_concat_128(_pA2_128, _pA2_128); - __m256i _pA3 = __lasx_concat_128(_pA3_128, _pA3_128); - __m256i _s0 = __lasx_xvmulwev_h_b(_pA0, _pB); - __m256i _s1 = __lasx_xvmulwev_h_b(_pA1, _pB); - __m256i _s2 = __lasx_xvmulwev_h_b(_pA2, _pB); - __m256i _s3 = __lasx_xvmulwev_h_b(_pA3, _pB); - _s0 = __lasx_xvmaddwod_h_b(_s0, _pA0, _pB); - _s1 = __lasx_xvmaddwod_h_b(_s1, _pA1, _pB); - _s2 = __lasx_xvmaddwod_h_b(_s2, _pA2, _pB); - _s3 = __lasx_xvmaddwod_h_b(_s3, _pA3, _pB); - _sum0 = __lasx_xvadd_w(_sum0, __lasx_xvhaddw_w_h(_s0, _s0)); - _sum1 = __lasx_xvadd_w(_sum1, __lasx_xvhaddw_w_h(_s1, _s1)); - _sum2 = __lasx_xvadd_w(_sum2, __lasx_xvhaddw_w_h(_s2, _s2)); - _sum3 = __lasx_xvadd_w(_sum3, __lasx_xvhaddw_w_h(_s3, _s3)); + __m256i _pB1 = __lasx_xvshuf4i_w(_pB, _LSX_SHUFFLE(0, 3, 2, 1)); + __m256i _s = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA, _pB), _pA, _pB); + _sum0 = __lasx_xvadd_w(_sum0, __lasx_xvhaddw_w_h(_s, _s)); + _s = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA, _pB1), _pA, _pB1); + _sum1 = __lasx_xvadd_w(_sum1, __lasx_xvhaddw_w_h(_s, _s)); + _s = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA1, _pB), _pA1, _pB); + _sum2 = __lasx_xvadd_w(_sum2, __lasx_xvhaddw_w_h(_s, _s)); + _s = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA1, _pB1), _pA1, _pB1); + _sum3 = __lasx_xvadd_w(_sum3, __lasx_xvhaddw_w_h(_s, _s)); pB += 32; pA += 16; } + _sum2 = __lasx_xvshuf4i_w(_sum2, _LSX_SHUFFLE(1, 0, 3, 2)); + _sum3 = __lasx_xvshuf4i_w(_sum3, _LSX_SHUFFLE(1, 0, 3, 2)); + { + __m256i _tmp0 = __lasx_xvilvl_w(_sum1, _sum0); + __m256i _tmp1 = __lasx_xvilvh_w(_sum1, _sum0); + __m256i _tmp2 = __lasx_xvilvl_w(_sum3, _sum2); + __m256i _tmp3 = __lasx_xvilvh_w(_sum3, _sum2); + _sum0 = __lasx_xvilvl_d(_tmp2, _tmp0); + _sum1 = __lasx_xvilvh_d(_tmp2, _tmp0); + _sum2 = __lasx_xvilvl_d(_tmp3, _tmp1); + _sum3 = __lasx_xvilvh_d(_tmp3, _tmp1); + } + _sum1 = __lasx_xvshuf4i_w(_sum1, _LSX_SHUFFLE(2, 1, 0, 3)); + _sum2 = __lasx_xvshuf4i_w(_sum2, _LSX_SHUFFLE(1, 0, 3, 2)); + _sum3 = __lasx_xvshuf4i_w(_sum3, _LSX_SHUFFLE(0, 3, 2, 1)); if (kk + 1 < max_kk) { __m128i _pB = __lsx_vld(pB, 0); @@ -1685,11 +2227,12 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de { __m128i _pB = __lsx_vldrepl_d(pB, 0); _pB = __lsx_vilvl_b(__lsx_vslti_b(_pB, 0), _pB); - const int a0123 = __lsx_vpickve2gr_w(__lsx_vldrepl_w(pA, 0), 0); - __m128i _s0 = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)a0123), _pB); - __m128i _s1 = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)(a0123 >> 8)), _pB); - __m128i _s2 = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)(a0123 >> 16)), _pB); - __m128i _s3 = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)(a0123 >> 24)), _pB); + __m128i _pA = __lsx_vldrepl_w(pA, 0); + _pA = __lsx_vilvl_b(__lsx_vslti_b(_pA, 0), _pA); + __m128i _s0 = __lsx_vmul_h(__lsx_vreplvei_h(_pA, 0), _pB); + __m128i _s1 = __lsx_vmul_h(__lsx_vreplvei_h(_pA, 1), _pB); + __m128i _s2 = __lsx_vmul_h(__lsx_vreplvei_h(_pA, 2), _pB); + __m128i _s3 = __lsx_vmul_h(__lsx_vreplvei_h(_pA, 3), _pB); _sum0 = __lasx_xvadd_w(_sum0, __lasx_vext2xv_w_h(__lasx_cast_128(_s0))); _sum1 = __lasx_xvadd_w(_sum1, __lasx_vext2xv_w_h(__lasx_cast_128(_s1))); _sum2 = __lasx_xvadd_w(_sum2, __lasx_vext2xv_w_h(__lasx_cast_128(_s2))); @@ -1711,27 +2254,29 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB_descales += 8; } - __lasx_xvst(_out0, outptr + (ii + 0) * max_jj + jj, 0); - __lasx_xvst(_out1, outptr + (ii + 1) * max_jj + jj, 0); - __lasx_xvst(_out2, outptr + (ii + 2) * max_jj + jj, 0); - __lasx_xvst(_out3, outptr + (ii + 3) * max_jj + jj, 0); + __lasx_xvst(_out0, outptr, 0); + __lasx_xvst(_out1, outptr + 8, 0); + __lasx_xvst(_out2, outptr + 16, 0); + __lasx_xvst(_out3, outptr + 24, 0); + outptr += 32; + pB += (size_t)8 * (full_K - k0 - max_kk0); + pB_descales += (size_t)8 * (block_count - block_start - tile_blocks); } -#endif -#if __loongarch_sx +#endif // __loongarch_asx for (; jj + 7 < max_jj; jj += 8) { - const signed char* pB0 = pB; - const signed char* pB1 = pB + 4 * K; - const float* pB_descales0 = pB_descales; - const float* pB_descales1 = pB_descales + 4 * ((K + block_size - 1) / block_size); - __m128 _out00 = (__m128)__lsx_vldi(0); - __m128 _out01 = (__m128)__lsx_vldi(0); - __m128 _out10 = (__m128)__lsx_vldi(0); - __m128 _out11 = (__m128)__lsx_vldi(0); - __m128 _out20 = (__m128)__lsx_vldi(0); - __m128 _out21 = (__m128)__lsx_vldi(0); - __m128 _out30 = (__m128)__lsx_vldi(0); - __m128 _out31 = (__m128)__lsx_vldi(0); + const signed char* pB0 = pB + 4 * k0; + const signed char* pB1 = pB + 4 * full_K + 4 * k0; + const float* pB_descales0 = pB_descales + 4 * block_start; + const float* pB_descales1 = pB_descales + 4 * block_count + 4 * block_start; + __m128 _out00 = k0 == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr, 0); + __m128 _out01 = k0 == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 4, 0); + __m128 _out10 = k0 == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 8, 0); + __m128 _out11 = k0 == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 12, 0); + __m128 _out20 = k0 == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 16, 0); + __m128 _out21 = k0 == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 20, 0); + __m128 _out30 = k0 == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 24, 0); + __m128 _out31 = k0 == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 28, 0); const signed char* pA = pAT; const float* pA_descales = pAT_descales; for (int k = 0; k < K; k += block_size) @@ -1748,33 +2293,45 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de int kk = 0; for (; kk + 3 < max_kk; kk += 4) { + __m128i _pA = __lsx_vld(pA, 0); + __m128i _pA1 = __lsx_vshuf4i_w(_pA, _LSX_SHUFFLE(1, 0, 3, 2)); __m128i _pB0 = __lsx_vld(pB0, 0); __m128i _pB1 = __lsx_vld(pB1, 0); - __m128i _pA = __lsx_vld(pA, 0); - __m128i _pA0 = __lsx_vreplvei_w(_pA, 0); - __m128i _pA1 = __lsx_vreplvei_w(_pA, 1); - __m128i _pA2 = __lsx_vreplvei_w(_pA, 2); - __m128i _pA3 = __lsx_vreplvei_w(_pA, 3); - __m128i _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB0), _pA0, _pB0); + __m128i _pB0r = __lsx_vshuf4i_w(_pB0, _LSX_SHUFFLE(0, 3, 2, 1)); + __m128i _pB1r = __lsx_vshuf4i_w(_pB1, _LSX_SHUFFLE(0, 3, 2, 1)); + __m128i _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA, _pB0), _pA, _pB0); _sum00 = __lsx_vadd_w(_sum00, __lsx_vhaddw_w_h(_s, _s)); - _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB1), _pA0, _pB1); - _sum01 = __lsx_vadd_w(_sum01, __lsx_vhaddw_w_h(_s, _s)); - _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA1, _pB0), _pA1, _pB0); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA, _pB0r), _pA, _pB0r); _sum10 = __lsx_vadd_w(_sum10, __lsx_vhaddw_w_h(_s, _s)); - _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA1, _pB1), _pA1, _pB1); - _sum11 = __lsx_vadd_w(_sum11, __lsx_vhaddw_w_h(_s, _s)); - _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA2, _pB0), _pA2, _pB0); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA1, _pB0), _pA1, _pB0); _sum20 = __lsx_vadd_w(_sum20, __lsx_vhaddw_w_h(_s, _s)); - _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA2, _pB1), _pA2, _pB1); - _sum21 = __lsx_vadd_w(_sum21, __lsx_vhaddw_w_h(_s, _s)); - _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA3, _pB0), _pA3, _pB0); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA1, _pB0r), _pA1, _pB0r); _sum30 = __lsx_vadd_w(_sum30, __lsx_vhaddw_w_h(_s, _s)); - _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA3, _pB1), _pA3, _pB1); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA, _pB1), _pA, _pB1); + _sum01 = __lsx_vadd_w(_sum01, __lsx_vhaddw_w_h(_s, _s)); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA, _pB1r), _pA, _pB1r); + _sum11 = __lsx_vadd_w(_sum11, __lsx_vhaddw_w_h(_s, _s)); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA1, _pB1), _pA1, _pB1); + _sum21 = __lsx_vadd_w(_sum21, __lsx_vhaddw_w_h(_s, _s)); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA1, _pB1r), _pA1, _pB1r); _sum31 = __lsx_vadd_w(_sum31, __lsx_vhaddw_w_h(_s, _s)); pB0 += 16; pB1 += 16; pA += 16; } + _sum20 = __lsx_vshuf4i_w(_sum20, _LSX_SHUFFLE(1, 0, 3, 2)); + _sum30 = __lsx_vshuf4i_w(_sum30, _LSX_SHUFFLE(1, 0, 3, 2)); + transpose4x4_epi32(_sum00, _sum10, _sum20, _sum30); + _sum10 = __lsx_vshuf4i_w(_sum10, _LSX_SHUFFLE(2, 1, 0, 3)); + _sum20 = __lsx_vshuf4i_w(_sum20, _LSX_SHUFFLE(1, 0, 3, 2)); + _sum30 = __lsx_vshuf4i_w(_sum30, _LSX_SHUFFLE(0, 3, 2, 1)); + + _sum21 = __lsx_vshuf4i_w(_sum21, _LSX_SHUFFLE(1, 0, 3, 2)); + _sum31 = __lsx_vshuf4i_w(_sum31, _LSX_SHUFFLE(1, 0, 3, 2)); + transpose4x4_epi32(_sum01, _sum11, _sum21, _sum31); + _sum11 = __lsx_vshuf4i_w(_sum11, _LSX_SHUFFLE(2, 1, 0, 3)); + _sum21 = __lsx_vshuf4i_w(_sum21, _LSX_SHUFFLE(1, 0, 3, 2)); + _sum31 = __lsx_vshuf4i_w(_sum31, _LSX_SHUFFLE(0, 3, 2, 1)); if (kk + 1 < max_kk) { __m128i _pB0 = __lsx_vldrepl_d(pB0, 0); @@ -1811,22 +2368,27 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de __m128i _pB1 = __lsx_vldrepl_w(pB1, 0); _pB0 = __lsx_vilvl_b(__lsx_vslti_b(_pB0, 0), _pB0); _pB1 = __lsx_vilvl_b(__lsx_vslti_b(_pB1, 0), _pB1); - const int a0123 = __lsx_vpickve2gr_w(__lsx_vldrepl_w(pA, 0), 0); - __m128i _s = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)a0123), _pB0); + __m128i _pA = __lsx_vldrepl_w(pA, 0); + _pA = __lsx_vilvl_b(__lsx_vslti_b(_pA, 0), _pA); + __m128i _pA0 = __lsx_vreplvei_h(_pA, 0); + __m128i _pA1 = __lsx_vreplvei_h(_pA, 1); + __m128i _pA2 = __lsx_vreplvei_h(_pA, 2); + __m128i _pA3 = __lsx_vreplvei_h(_pA, 3); + __m128i _s = __lsx_vmul_h(_pA0, _pB0); _sum00 = __lsx_vadd_w(_sum00, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); - _s = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)a0123), _pB1); + _s = __lsx_vmul_h(_pA0, _pB1); _sum01 = __lsx_vadd_w(_sum01, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); - _s = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)(a0123 >> 8)), _pB0); + _s = __lsx_vmul_h(_pA1, _pB0); _sum10 = __lsx_vadd_w(_sum10, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); - _s = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)(a0123 >> 8)), _pB1); + _s = __lsx_vmul_h(_pA1, _pB1); _sum11 = __lsx_vadd_w(_sum11, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); - _s = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)(a0123 >> 16)), _pB0); + _s = __lsx_vmul_h(_pA2, _pB0); _sum20 = __lsx_vadd_w(_sum20, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); - _s = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)(a0123 >> 16)), _pB1); + _s = __lsx_vmul_h(_pA2, _pB1); _sum21 = __lsx_vadd_w(_sum21, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); - _s = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)(a0123 >> 24)), _pB0); + _s = __lsx_vmul_h(_pA3, _pB0); _sum30 = __lsx_vadd_w(_sum30, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); - _s = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)(a0123 >> 24)), _pB1); + _s = __lsx_vmul_h(_pA3, _pB1); _sum31 = __lsx_vadd_w(_sum31, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); pB0 += 4; pB1 += 4; @@ -1850,23 +2412,26 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB_descales0 += 4; pB_descales1 += 4; } - pB = pB1; - pB_descales = pB_descales1; - __lsx_vst((__m128i)_out00, outptr + (ii + 0) * max_jj + jj, 0); - __lsx_vst((__m128i)_out01, outptr + (ii + 0) * max_jj + jj + 4, 0); - __lsx_vst((__m128i)_out10, outptr + (ii + 1) * max_jj + jj, 0); - __lsx_vst((__m128i)_out11, outptr + (ii + 1) * max_jj + jj + 4, 0); - __lsx_vst((__m128i)_out20, outptr + (ii + 2) * max_jj + jj, 0); - __lsx_vst((__m128i)_out21, outptr + (ii + 2) * max_jj + jj + 4, 0); - __lsx_vst((__m128i)_out30, outptr + (ii + 3) * max_jj + jj, 0); - __lsx_vst((__m128i)_out31, outptr + (ii + 3) * max_jj + jj + 4, 0); + pB = pB1 + 4 * (full_K - k0 - max_kk0); + pB_descales = pB_descales1 + 4 * (block_count - block_start - tile_blocks); + __lsx_vst((__m128i)_out00, outptr, 0); + __lsx_vst((__m128i)_out01, outptr + 4, 0); + __lsx_vst((__m128i)_out10, outptr + 8, 0); + __lsx_vst((__m128i)_out11, outptr + 12, 0); + __lsx_vst((__m128i)_out20, outptr + 16, 0); + __lsx_vst((__m128i)_out21, outptr + 20, 0); + __lsx_vst((__m128i)_out30, outptr + 24, 0); + __lsx_vst((__m128i)_out31, outptr + 28, 0); + outptr += 32; } for (; jj + 3 < max_jj; jj += 4) { - __m128 _out0 = (__m128)__lsx_vldi(0); - __m128 _out1 = (__m128)__lsx_vldi(0); - __m128 _out2 = (__m128)__lsx_vldi(0); - __m128 _out3 = (__m128)__lsx_vldi(0); + pB += (size_t)4 * k0; + pB_descales += (size_t)4 * block_start; + __m128 _out0 = k0 == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr, 0); + __m128 _out1 = k0 == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 4, 0); + __m128 _out2 = k0 == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 8, 0); + __m128 _out3 = k0 == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 12, 0); const signed char* pA = pAT; const float* pA_descales = pAT_descales; for (int k = 0; k < K; k += block_size) @@ -1879,27 +2444,27 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de int kk = 0; for (; kk + 3 < max_kk; kk += 4) { - __m128i _pB = __lsx_vld(pB, 0); __m128i _pA = __lsx_vld(pA, 0); - __m128i _pA0 = __lsx_vreplvei_w(_pA, 0); - __m128i _pA1 = __lsx_vreplvei_w(_pA, 1); - __m128i _pA2 = __lsx_vreplvei_w(_pA, 2); - __m128i _pA3 = __lsx_vreplvei_w(_pA, 3); - __m128i _s0 = __lsx_vmulwev_h_b(_pA0, _pB); - __m128i _s1 = __lsx_vmulwev_h_b(_pA1, _pB); - __m128i _s2 = __lsx_vmulwev_h_b(_pA2, _pB); - __m128i _s3 = __lsx_vmulwev_h_b(_pA3, _pB); - _s0 = __lsx_vmaddwod_h_b(_s0, _pA0, _pB); - _s1 = __lsx_vmaddwod_h_b(_s1, _pA1, _pB); - _s2 = __lsx_vmaddwod_h_b(_s2, _pA2, _pB); - _s3 = __lsx_vmaddwod_h_b(_s3, _pA3, _pB); - _sum0 = __lsx_vadd_w(_sum0, __lsx_vhaddw_w_h(_s0, _s0)); - _sum1 = __lsx_vadd_w(_sum1, __lsx_vhaddw_w_h(_s1, _s1)); - _sum2 = __lsx_vadd_w(_sum2, __lsx_vhaddw_w_h(_s2, _s2)); - _sum3 = __lsx_vadd_w(_sum3, __lsx_vhaddw_w_h(_s3, _s3)); + __m128i _pA1 = __lsx_vshuf4i_w(_pA, _LSX_SHUFFLE(1, 0, 3, 2)); + __m128i _pB = __lsx_vld(pB, 0); + __m128i _pB1 = __lsx_vshuf4i_w(_pB, _LSX_SHUFFLE(0, 3, 2, 1)); + __m128i _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA, _pB), _pA, _pB); + _sum0 = __lsx_vadd_w(_sum0, __lsx_vhaddw_w_h(_s, _s)); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA, _pB1), _pA, _pB1); + _sum1 = __lsx_vadd_w(_sum1, __lsx_vhaddw_w_h(_s, _s)); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA1, _pB), _pA1, _pB); + _sum2 = __lsx_vadd_w(_sum2, __lsx_vhaddw_w_h(_s, _s)); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA1, _pB1), _pA1, _pB1); + _sum3 = __lsx_vadd_w(_sum3, __lsx_vhaddw_w_h(_s, _s)); pB += 16; pA += 16; } + _sum2 = __lsx_vshuf4i_w(_sum2, _LSX_SHUFFLE(1, 0, 3, 2)); + _sum3 = __lsx_vshuf4i_w(_sum3, _LSX_SHUFFLE(1, 0, 3, 2)); + transpose4x4_epi32(_sum0, _sum1, _sum2, _sum3); + _sum1 = __lsx_vshuf4i_w(_sum1, _LSX_SHUFFLE(2, 1, 0, 3)); + _sum2 = __lsx_vshuf4i_w(_sum2, _LSX_SHUFFLE(1, 0, 3, 2)); + _sum3 = __lsx_vshuf4i_w(_sum3, _LSX_SHUFFLE(0, 3, 2, 1)); if (kk + 1 < max_kk) { __m128i _pB = __lsx_vldrepl_d(pB, 0); @@ -1928,11 +2493,12 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de { __m128i _pB = __lsx_vldrepl_w(pB, 0); _pB = __lsx_vilvl_b(__lsx_vslti_b(_pB, 0), _pB); - const int a0123 = __lsx_vpickve2gr_w(__lsx_vldrepl_w(pA, 0), 0); - __m128i _s0 = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)a0123), _pB); - __m128i _s1 = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)(a0123 >> 8)), _pB); - __m128i _s2 = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)(a0123 >> 16)), _pB); - __m128i _s3 = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)(a0123 >> 24)), _pB); + __m128i _pA = __lsx_vldrepl_w(pA, 0); + _pA = __lsx_vilvl_b(__lsx_vslti_b(_pA, 0), _pA); + __m128i _s0 = __lsx_vmul_h(__lsx_vreplvei_h(_pA, 0), _pB); + __m128i _s1 = __lsx_vmul_h(__lsx_vreplvei_h(_pA, 1), _pB); + __m128i _s2 = __lsx_vmul_h(__lsx_vreplvei_h(_pA, 2), _pB); + __m128i _s3 = __lsx_vmul_h(__lsx_vreplvei_h(_pA, 3), _pB); _sum0 = __lsx_vadd_w(_sum0, __lsx_vilvl_h(__lsx_vslti_h(_s0, 0), _s0)); _sum1 = __lsx_vadd_w(_sum1, __lsx_vilvl_h(__lsx_vslti_h(_s1, 0), _s1)); _sum2 = __lsx_vadd_w(_sum2, __lsx_vilvl_h(__lsx_vslti_h(_s2, 0), _s2)); @@ -1948,17 +2514,27 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pA_descales += 4; pB_descales += 4; } - __lsx_vst((__m128i)_out0, outptr + (ii + 0) * max_jj + jj, 0); - __lsx_vst((__m128i)_out1, outptr + (ii + 1) * max_jj + jj, 0); - __lsx_vst((__m128i)_out2, outptr + (ii + 2) * max_jj + jj, 0); - __lsx_vst((__m128i)_out3, outptr + (ii + 3) * max_jj + jj, 0); + __lsx_vst((__m128i)_out0, outptr, 0); + __lsx_vst((__m128i)_out1, outptr + 4, 0); + __lsx_vst((__m128i)_out2, outptr + 8, 0); + __lsx_vst((__m128i)_out3, outptr + 12, 0); + outptr += 16; + pB += (size_t)4 * (full_K - k0 - max_kk0); + pB_descales += (size_t)4 * (block_count - block_start - tile_blocks); } -#endif -#if __loongarch_sx for (; jj + 1 < max_jj; jj += 2) { + pB += (size_t)2 * k0; + pB_descales += (size_t)2 * block_start; __m128 _out0 = (__m128)__lsx_vldi(0); __m128 _out1 = (__m128)__lsx_vldi(0); + if (k0 != 0) + { + __m128i _out01 = __lsx_vld(outptr, 0); + __m128i _out23 = __lsx_vld(outptr + 4, 0); + _out0 = (__m128)__lsx_vpickev_w(_out23, _out01); + _out1 = (__m128)__lsx_vpickod_w(_out23, _out01); + } const signed char* pA = pAT; const float* pA_descales = pAT_descales; for (int k = 0; k < K; k += block_size) @@ -1986,8 +2562,8 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de __m128i _pA1 = __lsx_vreplvei_w(__lsx_vpickod_b(_pA, _pA), 0); _pA0 = __lsx_vilvl_b(__lsx_vslti_b(_pA0, 0), _pA0); _pA1 = __lsx_vilvl_b(__lsx_vslti_b(_pA1, 0), _pA1); - int b01 = (unsigned char)pB[0] | ((unsigned char)pB[2] << 8); - __m128i _pB0 = __lsx_vreplgr2vr_w(b01); + __m128i _pBs = __lsx_vldrepl_w(pB, 0); + __m128i _pB0 = __lsx_vreplvei_w(__lsx_vpickev_b(_pBs, _pBs), 0); _pB0 = __lsx_vilvl_b(__lsx_vslti_b(_pB0, 0), _pB0); _pB0 = __lsx_vshuf4i_h(_pB0, _LSX_SHUFFLE(1, 0, 1, 0)); __m128i _pB1 = __lsx_vshuf4i_h(_pB0, _LSX_SHUFFLE(0, 3, 2, 1)); @@ -1995,8 +2571,7 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de __m128i _s1 = __lsx_vmul_h(_pA0, _pB1); _sum0 = __lsx_vadd_w(_sum0, __lsx_vilvl_h(__lsx_vslti_h(_s0, 0), _s0)); _sum1 = __lsx_vadd_w(_sum1, __lsx_vilvl_h(__lsx_vslti_h(_s1, 0), _s1)); - b01 = (unsigned char)pB[1] | ((unsigned char)pB[3] << 8); - _pB0 = __lsx_vreplgr2vr_w(b01); + _pB0 = __lsx_vreplvei_w(__lsx_vpickod_b(_pBs, _pBs), 0); _pB0 = __lsx_vilvl_b(__lsx_vslti_b(_pB0, 0), _pB0); _pB0 = __lsx_vshuf4i_h(_pB0, _LSX_SHUFFLE(1, 0, 1, 0)); _pB1 = __lsx_vshuf4i_h(_pB0, _LSX_SHUFFLE(0, 3, 2, 1)); @@ -2012,10 +2587,8 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de { __m128i _pA = __lsx_vldrepl_w(pA, 0); _pA = __lsx_vilvl_b(__lsx_vslti_b(_pA, 0), _pA); - int b01 = (unsigned char)pB[0] | ((unsigned char)pB[1] << 8); - __m128i _pB0 = __lsx_vreplgr2vr_w(b01); + __m128i _pB0 = __lsx_vldrepl_h(pB, 0); _pB0 = __lsx_vilvl_b(__lsx_vslti_b(_pB0, 0), _pB0); - _pB0 = __lsx_vshuf4i_h(_pB0, _LSX_SHUFFLE(1, 0, 1, 0)); __m128i _pB1 = __lsx_vshuf4i_h(_pB0, _LSX_SHUFFLE(0, 3, 2, 1)); __m128i _s0 = __lsx_vmul_h(_pA, _pB0); __m128i _s1 = __lsx_vmul_h(_pA, _pB1); @@ -2030,24 +2603,29 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de __m128i _sum1o = __lsx_vshuf4i_w(_sum1, _LSX_SHUFFLE(2, 0, 3, 1)); __m128i _sumc0 = __lsx_vilvl_w(_sum1o, _sum0e); __m128i _sumc1 = __lsx_vilvl_w(_sum0o, _sum1e); - const __m128 _ascale = (__m128)__lsx_vld(pA_descales, 0); + __m128 _ascale = (__m128)__lsx_vld(pA_descales, 0); _out0 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sumc0), __lsx_vfmul_s(_ascale, __lsx_vreplfr2vr_s(pB_descales[0])), _out0); _out1 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sumc1), __lsx_vfmul_s(_ascale, __lsx_vreplfr2vr_s(pB_descales[1])), _out1); pA_descales += 4; pB_descales += 2; } - __lsx_vstelm_w((__m128i)_out0, outptr + (ii + 0) * max_jj + jj, 0, 0); - __lsx_vstelm_w((__m128i)_out1, outptr + (ii + 0) * max_jj + jj + 1, 0, 0); - __lsx_vstelm_w((__m128i)_out0, outptr + (ii + 1) * max_jj + jj, 0, 1); - __lsx_vstelm_w((__m128i)_out1, outptr + (ii + 1) * max_jj + jj + 1, 0, 1); - __lsx_vstelm_w((__m128i)_out0, outptr + (ii + 2) * max_jj + jj, 0, 2); - __lsx_vstelm_w((__m128i)_out1, outptr + (ii + 2) * max_jj + jj + 1, 0, 2); - __lsx_vstelm_w((__m128i)_out0, outptr + (ii + 3) * max_jj + jj, 0, 3); - __lsx_vstelm_w((__m128i)_out1, outptr + (ii + 3) * max_jj + jj + 1, 0, 3); + __lsx_vstelm_w((__m128i)_out0, outptr, 0, 0); + __lsx_vstelm_w((__m128i)_out1, outptr + 1, 0, 0); + __lsx_vstelm_w((__m128i)_out0, outptr + 2, 0, 1); + __lsx_vstelm_w((__m128i)_out1, outptr + 3, 0, 1); + __lsx_vstelm_w((__m128i)_out0, outptr + 4, 0, 2); + __lsx_vstelm_w((__m128i)_out1, outptr + 5, 0, 2); + __lsx_vstelm_w((__m128i)_out0, outptr + 6, 0, 3); + __lsx_vstelm_w((__m128i)_out1, outptr + 7, 0, 3); + outptr += 8; + pB += (size_t)2 * (full_K - k0 - max_kk0); + pB_descales += (size_t)2 * (block_count - block_start - tile_blocks); } for (; jj < max_jj; jj++) { - __m128 _out0 = (__m128)__lsx_vldi(0); + pB += k0; + pB_descales += block_start; + __m128 _out0 = k0 == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr, 0); const signed char* pA = pAT; const float* pA_descales = pAT_descales; for (int k = 0; k < K; k += block_size) @@ -2067,7 +2645,7 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de if (kk + 1 < max_kk) { __m128i _pA = __lsx_vldrepl_d(pA, 0); - __m128i _pB = __lsx_vreplgr2vr_h((unsigned char)pB[0] | ((unsigned char)pB[1] << 8)); + __m128i _pB = __lsx_vldrepl_h(pB, 0); __m128i _s0 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA, _pB), _pA, _pB); _sum0 = __lsx_vadd_w(_sum0, __lsx_vilvl_h(__lsx_vslti_h(_s0, 0), _s0)); pA += 8; @@ -2083,16 +2661,15 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pA += 4; pB++; } - const __m128 _scale = __lsx_vfmul_s((__m128)__lsx_vld(pA_descales, 0), __lsx_vreplfr2vr_s(*pB_descales++)); + __m128 _scale = __lsx_vfmul_s((__m128)__lsx_vld(pA_descales, 0), __lsx_vreplfr2vr_s(*pB_descales++)); _out0 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum0), _scale, _out0); pA_descales += 4; } - __lsx_vstelm_w((__m128i)_out0, outptr + (ii + 0) * max_jj + jj, 0, 0); - __lsx_vstelm_w((__m128i)_out0, outptr + (ii + 1) * max_jj + jj, 0, 1); - __lsx_vstelm_w((__m128i)_out0, outptr + (ii + 2) * max_jj + jj, 0, 2); - __lsx_vstelm_w((__m128i)_out0, outptr + (ii + 3) * max_jj + jj, 0, 3); + __lsx_vst((__m128i)_out0, outptr, 0); + outptr += 4; + pB += full_K - k0 - max_kk0; + pB_descales += block_count - block_start - tile_blocks; } -#endif pAT += A_hstep * 4; pAT_descales += A_descales_hstep * 4; } @@ -2105,14 +2682,14 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #if __loongarch_asx for (; jj + 15 < max_jj; jj += 16) { - const signed char* pB0 = pB; - const signed char* pB1 = pB + 8 * K; - const float* pB_descales0 = pB_descales; - const float* pB_descales1 = pB_descales + 8 * ((K + block_size - 1) / block_size); - __m256 _out00 = (__m256)__lasx_xvldi(0); - __m256 _out01 = (__m256)__lasx_xvldi(0); - __m256 _out10 = (__m256)__lasx_xvldi(0); - __m256 _out11 = (__m256)__lasx_xvldi(0); + const signed char* pB0 = pB + 8 * k0; + const signed char* pB1 = pB + 8 * full_K + 8 * k0; + const float* pB_descales0 = pB_descales + 8 * block_start; + const float* pB_descales1 = pB_descales + 8 * block_count + 8 * block_start; + __m256 _out00 = k0 == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr, 0); + __m256 _out01 = k0 == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr + 8, 0); + __m256 _out10 = k0 == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr + 16, 0); + __m256 _out11 = k0 == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr + 24, 0); const signed char* pA = pAT; const float* pA_descales = pAT_descales; for (int k = 0; k < K; k += block_size) @@ -2127,9 +2704,8 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de { __m256i _pB0 = __lasx_xvld(pB0, 0); __m256i _pB1 = __lasx_xvld(pB1, 0); - __m128i _pAs = __lsx_vldrepl_d(pA, 0); - __m256i _pA0 = __lasx_xvreplgr2vr_w(__lsx_vpickve2gr_w(_pAs, 0)); - __m256i _pA1 = __lasx_xvreplgr2vr_w(__lsx_vpickve2gr_w(_pAs, 1)); + __m256i _pA0 = __lasx_xvldrepl_w(pA, 0); + __m256i _pA1 = __lasx_xvldrepl_w(pA + 4, 0); __m256i _s = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA0, _pB0), _pA0, _pB0); _sum00 = __lasx_xvadd_w(_sum00, __lasx_xvhaddw_w_h(_s, _s)); _s = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA0, _pB1), _pA0, _pB1); @@ -2168,16 +2744,17 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de __m128i _pB1 = __lsx_vldrepl_d(pB1, 0); _pB0 = __lsx_vilvl_b(__lsx_vslti_b(_pB0, 0), _pB0); _pB1 = __lsx_vilvl_b(__lsx_vslti_b(_pB1, 0), _pB1); - __m128i _pAs = __lsx_vldrepl_h(pA, 0); - const int a0 = (signed char)__lsx_vpickve2gr_b(_pAs, 0); - const int a1 = (signed char)__lsx_vpickve2gr_b(_pAs, 1); - __m128i _s = __lsx_vmul_h(__lsx_vreplgr2vr_h(a0), _pB0); + __m128i _pA = __lsx_vldrepl_h(pA, 0); + _pA = __lsx_vilvl_b(__lsx_vslti_b(_pA, 0), _pA); + __m128i _pA0 = __lsx_vreplvei_h(_pA, 0); + __m128i _pA1 = __lsx_vreplvei_h(_pA, 1); + __m128i _s = __lsx_vmul_h(_pA0, _pB0); _sum00 = __lasx_xvadd_w(_sum00, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); - _s = __lsx_vmul_h(__lsx_vreplgr2vr_h(a0), _pB1); + _s = __lsx_vmul_h(_pA0, _pB1); _sum01 = __lasx_xvadd_w(_sum01, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); - _s = __lsx_vmul_h(__lsx_vreplgr2vr_h(a1), _pB0); + _s = __lsx_vmul_h(_pA1, _pB0); _sum10 = __lasx_xvadd_w(_sum10, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); - _s = __lsx_vmul_h(__lsx_vreplgr2vr_h(a1), _pB1); + _s = __lsx_vmul_h(_pA1, _pB1); _sum11 = __lasx_xvadd_w(_sum11, __lasx_vext2xv_w_h(__lasx_cast_128(_s))); pB0 += 8; pB1 += 8; @@ -2195,17 +2772,20 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB_descales0 += 8; pB_descales1 += 8; } - pB = pB1; - pB_descales = pB_descales1; - __lasx_xvst(_out00, outptr + (ii + 0) * max_jj + jj, 0); - __lasx_xvst(_out01, outptr + (ii + 0) * max_jj + jj + 8, 0); - __lasx_xvst(_out10, outptr + (ii + 1) * max_jj + jj, 0); - __lasx_xvst(_out11, outptr + (ii + 1) * max_jj + jj + 8, 0); + pB = pB1 + 8 * (full_K - k0 - max_kk0); + pB_descales = pB_descales1 + 8 * (block_count - block_start - tile_blocks); + __lasx_xvst(_out00, outptr, 0); + __lasx_xvst(_out01, outptr + 8, 0); + __lasx_xvst(_out10, outptr + 16, 0); + __lasx_xvst(_out11, outptr + 24, 0); + outptr += 32; } for (; jj + 7 < max_jj; jj += 8) { - __m256 _out0 = (__m256)__lasx_xvldi(0); - __m256 _out1 = (__m256)__lasx_xvldi(0); + pB += (size_t)8 * k0; + pB_descales += (size_t)8 * block_start; + __m256 _out0 = k0 == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr, 0); + __m256 _out1 = k0 == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr + 8, 0); const signed char* pA = pAT; const float* pA_descales = pAT_descales; for (int k = 0; k < K; k += block_size) @@ -2217,15 +2797,12 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de for (; kk + 3 < max_kk; kk += 4) { __m256i _pB = __lasx_xvld(pB, 0); - __m128i _pAs = __lsx_vldrepl_d(pA, 0); - __m256i _pA0 = __lasx_xvreplgr2vr_w(__lsx_vpickve2gr_w(_pAs, 0)); - __m256i _pA1 = __lasx_xvreplgr2vr_w(__lsx_vpickve2gr_w(_pAs, 1)); - __m256i _s0 = __lasx_xvmulwev_h_b(_pA0, _pB); - __m256i _s1 = __lasx_xvmulwev_h_b(_pA1, _pB); - _s0 = __lasx_xvmaddwod_h_b(_s0, _pA0, _pB); - _s1 = __lasx_xvmaddwod_h_b(_s1, _pA1, _pB); - _sum0 = __lasx_xvadd_w(_sum0, __lasx_xvhaddw_w_h(_s0, _s0)); - _sum1 = __lasx_xvadd_w(_sum1, __lasx_xvhaddw_w_h(_s1, _s1)); + __m256i _pA0 = __lasx_xvldrepl_w(pA, 0); + __m256i _pA1 = __lasx_xvldrepl_w(pA + 4, 0); + __m256i _s = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA0, _pB), _pA0, _pB); + _sum0 = __lasx_xvadd_w(_sum0, __lasx_xvhaddw_w_h(_s, _s)); + _s = __lasx_xvmaddwod_h_b(__lasx_xvmulwev_h_b(_pA1, _pB), _pA1, _pB); + _sum1 = __lasx_xvadd_w(_sum1, __lasx_xvhaddw_w_h(_s, _s)); pB += 32; pA += 8; } @@ -2247,9 +2824,10 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de { __m128i _pB = __lsx_vldrepl_d(pB, 0); _pB = __lsx_vilvl_b(__lsx_vslti_b(_pB, 0), _pB); - __m128i _pAs = __lsx_vldrepl_h(pA, 0); - __m128i _s0 = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)__lsx_vpickve2gr_b(_pAs, 0)), _pB); - __m128i _s1 = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)__lsx_vpickve2gr_b(_pAs, 1)), _pB); + __m128i _pA = __lsx_vldrepl_h(pA, 0); + _pA = __lsx_vilvl_b(__lsx_vslti_b(_pA, 0), _pA); + __m128i _s0 = __lsx_vmul_h(__lsx_vreplvei_h(_pA, 0), _pB); + __m128i _s1 = __lsx_vmul_h(__lsx_vreplvei_h(_pA, 1), _pB); _sum0 = __lasx_xvadd_w(_sum0, __lasx_vext2xv_w_h(__lasx_cast_128(_s0))); _sum1 = __lasx_xvadd_w(_sum1, __lasx_vext2xv_w_h(__lasx_cast_128(_s1))); pB += 8; @@ -2261,21 +2839,23 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pA_descales += 2; pB_descales += 8; } - __lasx_xvst(_out0, outptr + (ii + 0) * max_jj + jj, 0); - __lasx_xvst(_out1, outptr + (ii + 1) * max_jj + jj, 0); + __lasx_xvst(_out0, outptr, 0); + __lasx_xvst(_out1, outptr + 8, 0); + outptr += 16; + pB += (size_t)8 * (full_K - k0 - max_kk0); + pB_descales += (size_t)8 * (block_count - block_start - tile_blocks); } -#endif -#if __loongarch_sx +#endif // __loongarch_asx for (; jj + 7 < max_jj; jj += 8) { - const signed char* pB0 = pB; - const signed char* pB1 = pB + 4 * K; - const float* pB_descales0 = pB_descales; - const float* pB_descales1 = pB_descales + 4 * ((K + block_size - 1) / block_size); - __m128 _out00 = (__m128)__lsx_vldi(0); - __m128 _out01 = (__m128)__lsx_vldi(0); - __m128 _out10 = (__m128)__lsx_vldi(0); - __m128 _out11 = (__m128)__lsx_vldi(0); + const signed char* pB0 = pB + 4 * k0; + const signed char* pB1 = pB + 4 * full_K + 4 * k0; + const float* pB_descales0 = pB_descales + 4 * block_start; + const float* pB_descales1 = pB_descales + 4 * block_count + 4 * block_start; + __m128 _out00 = k0 == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr, 0); + __m128 _out01 = k0 == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 4, 0); + __m128 _out10 = k0 == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 8, 0); + __m128 _out11 = k0 == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 12, 0); const signed char* pA = pAT; const float* pA_descales = pAT_descales; for (int k = 0; k < K; k += block_size) @@ -2290,9 +2870,8 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de { __m128i _pB0 = __lsx_vld(pB0, 0); __m128i _pB1 = __lsx_vld(pB1, 0); - __m128i _pAs = __lsx_vldrepl_d(pA, 0); - __m128i _pA0 = __lsx_vreplvei_w(_pAs, 0); - __m128i _pA1 = __lsx_vreplvei_w(_pAs, 1); + __m128i _pA0 = __lsx_vldrepl_w(pA, 0); + __m128i _pA1 = __lsx_vldrepl_w(pA + 4, 0); __m128i _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB0), _pA0, _pB0); _sum00 = __lsx_vadd_w(_sum00, __lsx_vhaddw_w_h(_s, _s)); _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB1), _pA0, _pB1); @@ -2331,16 +2910,17 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de __m128i _pB1 = __lsx_vldrepl_w(pB1, 0); _pB0 = __lsx_vilvl_b(__lsx_vslti_b(_pB0, 0), _pB0); _pB1 = __lsx_vilvl_b(__lsx_vslti_b(_pB1, 0), _pB1); - __m128i _pAs = __lsx_vldrepl_h(pA, 0); - const int a0 = (signed char)__lsx_vpickve2gr_b(_pAs, 0); - const int a1 = (signed char)__lsx_vpickve2gr_b(_pAs, 1); - __m128i _s = __lsx_vmul_h(__lsx_vreplgr2vr_h(a0), _pB0); + __m128i _pA = __lsx_vldrepl_h(pA, 0); + _pA = __lsx_vilvl_b(__lsx_vslti_b(_pA, 0), _pA); + __m128i _pA0 = __lsx_vreplvei_h(_pA, 0); + __m128i _pA1 = __lsx_vreplvei_h(_pA, 1); + __m128i _s = __lsx_vmul_h(_pA0, _pB0); _sum00 = __lsx_vadd_w(_sum00, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); - _s = __lsx_vmul_h(__lsx_vreplgr2vr_h(a0), _pB1); + _s = __lsx_vmul_h(_pA0, _pB1); _sum01 = __lsx_vadd_w(_sum01, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); - _s = __lsx_vmul_h(__lsx_vreplgr2vr_h(a1), _pB0); + _s = __lsx_vmul_h(_pA1, _pB0); _sum10 = __lsx_vadd_w(_sum10, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); - _s = __lsx_vmul_h(__lsx_vreplgr2vr_h(a1), _pB1); + _s = __lsx_vmul_h(_pA1, _pB1); _sum11 = __lsx_vadd_w(_sum11, __lsx_vilvl_h(__lsx_vslti_h(_s, 0), _s)); pB0 += 4; pB1 += 4; @@ -2358,17 +2938,20 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB_descales0 += 4; pB_descales1 += 4; } - pB = pB1; - pB_descales = pB_descales1; - __lsx_vst((__m128i)_out00, outptr + (ii + 0) * max_jj + jj, 0); - __lsx_vst((__m128i)_out01, outptr + (ii + 0) * max_jj + jj + 4, 0); - __lsx_vst((__m128i)_out10, outptr + (ii + 1) * max_jj + jj, 0); - __lsx_vst((__m128i)_out11, outptr + (ii + 1) * max_jj + jj + 4, 0); + pB = pB1 + 4 * (full_K - k0 - max_kk0); + pB_descales = pB_descales1 + 4 * (block_count - block_start - tile_blocks); + __lsx_vst((__m128i)_out00, outptr, 0); + __lsx_vst((__m128i)_out01, outptr + 4, 0); + __lsx_vst((__m128i)_out10, outptr + 8, 0); + __lsx_vst((__m128i)_out11, outptr + 12, 0); + outptr += 16; } for (; jj + 3 < max_jj; jj += 4) { - __m128 _out0 = (__m128)__lsx_vldi(0); - __m128 _out1 = (__m128)__lsx_vldi(0); + pB += (size_t)4 * k0; + pB_descales += (size_t)4 * block_start; + __m128 _out0 = k0 == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr, 0); + __m128 _out1 = k0 == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 4, 0); const signed char* pA = pAT; const float* pA_descales = pAT_descales; for (int k = 0; k < K; k += block_size) @@ -2380,13 +2963,12 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de for (; kk + 3 < max_kk; kk += 4) { __m128i _pB = __lsx_vld(pB, 0); - __m128i _pAs = __lsx_vldrepl_d(pA, 0); - __m128i _pA0 = __lsx_vreplvei_w(_pAs, 0); - __m128i _pA1 = __lsx_vreplvei_w(_pAs, 1); - __m128i _s0 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB), _pA0, _pB); - __m128i _s1 = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA1, _pB), _pA1, _pB); - _sum0 = __lsx_vadd_w(_sum0, __lsx_vhaddw_w_h(_s0, _s0)); - _sum1 = __lsx_vadd_w(_sum1, __lsx_vhaddw_w_h(_s1, _s1)); + __m128i _pA0 = __lsx_vldrepl_w(pA, 0); + __m128i _pA1 = __lsx_vldrepl_w(pA + 4, 0); + __m128i _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA0, _pB), _pA0, _pB); + _sum0 = __lsx_vadd_w(_sum0, __lsx_vhaddw_w_h(_s, _s)); + _s = __lsx_vmaddwod_h_b(__lsx_vmulwev_h_b(_pA1, _pB), _pA1, _pB); + _sum1 = __lsx_vadd_w(_sum1, __lsx_vhaddw_w_h(_s, _s)); pB += 16; pA += 8; } @@ -2408,9 +2990,10 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de { __m128i _pB = __lsx_vldrepl_w(pB, 0); _pB = __lsx_vilvl_b(__lsx_vslti_b(_pB, 0), _pB); - __m128i _pAs = __lsx_vldrepl_h(pA, 0); - __m128i _s0 = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)__lsx_vpickve2gr_b(_pAs, 0)), _pB); - __m128i _s1 = __lsx_vmul_h(__lsx_vreplgr2vr_h((signed char)__lsx_vpickve2gr_b(_pAs, 1)), _pB); + __m128i _pA = __lsx_vldrepl_h(pA, 0); + _pA = __lsx_vilvl_b(__lsx_vslti_b(_pA, 0), _pA); + __m128i _s0 = __lsx_vmul_h(__lsx_vreplvei_h(_pA, 0), _pB); + __m128i _s1 = __lsx_vmul_h(__lsx_vreplvei_h(_pA, 1), _pB); _sum0 = __lsx_vadd_w(_sum0, __lsx_vilvl_h(__lsx_vslti_h(_s0, 0), _s0)); _sum1 = __lsx_vadd_w(_sum1, __lsx_vilvl_h(__lsx_vslti_h(_s1, 0), _s1)); pB += 4; @@ -2422,16 +3005,20 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pA_descales += 2; pB_descales += 4; } - __lsx_vst((__m128i)_out0, outptr + (ii + 0) * max_jj + jj, 0); - __lsx_vst((__m128i)_out1, outptr + (ii + 1) * max_jj + jj, 0); + __lsx_vst((__m128i)_out0, outptr, 0); + __lsx_vst((__m128i)_out1, outptr + 4, 0); + outptr += 8; + pB += (size_t)4 * (full_K - k0 - max_kk0); + pB_descales += (size_t)4 * (block_count - block_start - tile_blocks); } -#endif for (; jj + 1 < max_jj; jj += 2) { - float _out00 = 0.f; - float _out01 = 0.f; - float _out10 = 0.f; - float _out11 = 0.f; + pB += (size_t)2 * k0; + pB_descales += (size_t)2 * block_start; + float _out00 = k0 == 0 ? 0.f : outptr[0]; + float _out01 = k0 == 0 ? 0.f : outptr[1]; + float _out10 = k0 == 0 ? 0.f : outptr[2]; + float _out11 = k0 == 0 ? 0.f : outptr[3]; const signed char* pA = pAT; const float* pA_descales = pAT_descales; for (int k = 0; k < K; k += block_size) @@ -2486,15 +3073,20 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pA_descales += 2; pB_descales += 2; } - outptr[(ii + 0) * max_jj + jj] = _out00; - outptr[(ii + 0) * max_jj + jj + 1] = _out01; - outptr[(ii + 1) * max_jj + jj] = _out10; - outptr[(ii + 1) * max_jj + jj + 1] = _out11; + outptr[0] = _out00; + outptr[1] = _out01; + outptr[2] = _out10; + outptr[3] = _out11; + outptr += 4; + pB += (size_t)2 * (full_K - k0 - max_kk0); + pB_descales += (size_t)2 * (block_count - block_start - tile_blocks); } for (; jj < max_jj; jj++) { - float _out0 = 0.f; - float _out1 = 0.f; + pB += k0; + pB_descales += block_start; + float _out0 = k0 == 0 ? 0.f : outptr[0]; + float _out1 = k0 == 0 ? 0.f : outptr[1]; const signed char* pA = pAT; const float* pA_descales = pAT_descales; for (int k = 0; k < K; k += block_size) @@ -2535,28 +3127,165 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de _out1 += _sum1 * pA_descales[1] * bscale; pA_descales += 2; } - outptr[(ii + 0) * max_jj + jj] = _out0; - outptr[(ii + 1) * max_jj + jj] = _out1; + outptr[0] = _out0; + outptr[1] = _out1; + outptr += 2; + pB += full_K - k0 - max_kk0; + pB_descales += block_count - block_start - tile_blocks; } pAT += A_hstep * 2; pAT_descales += A_descales_hstep * 2; } #endif // __loongarch_sx +#if !__loongarch_sx + for (; ii + 1 < max_ii; ii += 2) + { + const signed char* pB = pBT; + const float* pB_descales = pBT_descales; + + int jj = 0; + for (; jj + 1 < max_jj; jj += 2) + { + pB += (size_t)2 * k0; + pB_descales += (size_t)2 * block_start; + float _out00 = k0 == 0 ? 0.f : outptr[0]; + float _out01 = k0 == 0 ? 0.f : outptr[1]; + float _out10 = k0 == 0 ? 0.f : outptr[2]; + float _out11 = k0 == 0 ? 0.f : outptr[3]; + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + int _sum00 = 0; + int _sum01 = 0; + int _sum10 = 0; + int _sum11 = 0; + const int max_kk = std::min(K - k, block_size); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + asm volatile("" + : + : + : "memory"); + + _sum00 += pA[0] * pB[0] + pA[1] * pB[1] + pA[2] * pB[2] + pA[3] * pB[3]; + _sum01 += pA[0] * pB[4] + pA[1] * pB[5] + pA[2] * pB[6] + pA[3] * pB[7]; + _sum10 += pA[4] * pB[0] + pA[5] * pB[1] + pA[6] * pB[2] + pA[7] * pB[3]; + _sum11 += pA[4] * pB[4] + pA[5] * pB[5] + pA[6] * pB[6] + pA[7] * pB[7]; + pB += 8; + pA += 8; + } + if (kk + 1 < max_kk) + { + _sum00 += pA[0] * pB[0] + pA[1] * pB[1]; + _sum01 += pA[0] * pB[2] + pA[1] * pB[3]; + _sum10 += pA[2] * pB[0] + pA[3] * pB[1]; + _sum11 += pA[2] * pB[2] + pA[3] * pB[3]; + pB += 4; + pA += 4; + kk += 2; + } + if (kk < max_kk) + { + _sum00 += pA[0] * pB[0]; + _sum01 += pA[0] * pB[1]; + _sum10 += pA[1] * pB[0]; + _sum11 += pA[1] * pB[1]; + pB += 2; + pA += 2; + } + const float bscale0 = pB_descales[0]; + const float bscale1 = pB_descales[1]; + const float ascale0 = pA_descales[0]; + const float ascale1 = pA_descales[1]; + _out00 += _sum00 * ascale0 * bscale0; + _out01 += _sum01 * ascale0 * bscale1; + _out10 += _sum10 * ascale1 * bscale0; + _out11 += _sum11 * ascale1 * bscale1; + pA_descales += 2; + pB_descales += 2; + } + outptr[0] = _out00; + outptr[1] = _out01; + outptr[2] = _out10; + outptr[3] = _out11; + outptr += 4; + pB += (size_t)2 * (full_K - k0 - max_kk0); + pB_descales += (size_t)2 * (block_count - block_start - tile_blocks); + } + for (; jj < max_jj; jj++) + { + pB += k0; + pB_descales += block_start; + float _out0 = k0 == 0 ? 0.f : outptr[0]; + float _out1 = k0 == 0 ? 0.f : outptr[1]; + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + int _sum0 = 0; + int _sum1 = 0; + const int max_kk = std::min(K - k, block_size); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + asm volatile("" + : + : + : "memory"); + + _sum0 += pA[0] * pB[0] + pA[1] * pB[1] + pA[2] * pB[2] + pA[3] * pB[3]; + _sum1 += pA[4] * pB[0] + pA[5] * pB[1] + pA[6] * pB[2] + pA[7] * pB[3]; + pB += 4; + pA += 8; + } + if (kk + 1 < max_kk) + { + _sum0 += pA[0] * pB[0] + pA[1] * pB[1]; + _sum1 += pA[2] * pB[0] + pA[3] * pB[1]; + pB += 2; + pA += 4; + kk += 2; + } + if (kk < max_kk) + { + _sum0 += pA[0] * pB[0]; + _sum1 += pA[1] * pB[0]; + pB++; + pA += 2; + } + const float bscale = *pB_descales++; + _out0 += _sum0 * pA_descales[0] * bscale; + _out1 += _sum1 * pA_descales[1] * bscale; + pA_descales += 2; + } + outptr[0] = _out0; + outptr[1] = _out1; + outptr += 2; + pB += full_K - k0 - max_kk0; + pB_descales += block_count - block_start - tile_blocks; + } + pAT += A_hstep * 2; + pAT_descales += A_descales_hstep * 2; + } +#endif // !__loongarch_sx for (; ii < max_ii; ii++) { const signed char* pB = pBT; const float* pB_descales = pBT_descales; int jj = 0; +#if __loongarch_sx #if __loongarch_asx for (; jj + 15 < max_jj; jj += 16) { - const signed char* pB0 = pB; - const signed char* pB1 = pB + 8 * K; - const float* pB_descales0 = pB_descales; - const float* pB_descales1 = pB_descales + 8 * ((K + block_size - 1) / block_size); - __m256 _out00 = (__m256)__lasx_xvldi(0); - __m256 _out01 = (__m256)__lasx_xvldi(0); + const signed char* pB0 = pB + 8 * k0; + const signed char* pB1 = pB + 8 * full_K + 8 * k0; + const float* pB_descales0 = pB_descales + 8 * block_start; + const float* pB_descales1 = pB_descales + 8 * block_count + 8 * block_start; + __m256 _out00 = k0 == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr, 0); + __m256 _out01 = k0 == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr + 8, 0); const signed char* pA = pAT; const float* pA_descales = pAT_descales; for (int k = 0; k < K; k += block_size) @@ -2615,14 +3344,17 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB_descales0 += 8; pB_descales1 += 8; } - pB = pB1; - pB_descales = pB_descales1; - __lasx_xvst(_out00, outptr + ii * max_jj + jj, 0); - __lasx_xvst(_out01, outptr + ii * max_jj + jj + 8, 0); + pB = pB1 + 8 * (full_K - k0 - max_kk0); + pB_descales = pB_descales1 + 8 * (block_count - block_start - tile_blocks); + __lasx_xvst(_out00, outptr, 0); + __lasx_xvst(_out01, outptr + 8, 0); + outptr += 16; } for (; jj + 7 < max_jj; jj += 8) { - __m256 _out0 = (__m256)__lasx_xvldi(0); + pB += (size_t)8 * k0; + pB_descales += (size_t)8 * block_start; + __m256 _out0 = k0 == 0 ? (__m256)__lasx_xvldi(0) : (__m256)__lasx_xvld(outptr, 0); const signed char* pA = pAT; const float* pA_descales = pAT_descales; for (int k = 0; k < K; k += block_size) @@ -2662,18 +3394,20 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de _out0 = __lasx_xvfmadd_s((__m256)__lasx_xvffint_s_w(_sum0), __lasx_xvfmul_s(_bscale, (__m256)__lasx_xvreplfr2vr_s(*pA_descales++)), _out0); pB_descales += 8; } - __lasx_xvst(_out0, outptr + ii * max_jj + jj, 0); + __lasx_xvst(_out0, outptr, 0); + outptr += 8; + pB += (size_t)8 * (full_K - k0 - max_kk0); + pB_descales += (size_t)8 * (block_count - block_start - tile_blocks); } #endif -#if __loongarch_sx for (; jj + 7 < max_jj; jj += 8) { - const signed char* pB0 = pB; - const signed char* pB1 = pB + 4 * K; - const float* pB_descales0 = pB_descales; - const float* pB_descales1 = pB_descales + 4 * ((K + block_size - 1) / block_size); - __m128 _out00 = (__m128)__lsx_vldi(0); - __m128 _out01 = (__m128)__lsx_vldi(0); + const signed char* pB0 = pB + 4 * k0; + const signed char* pB1 = pB + 4 * full_K + 4 * k0; + const float* pB_descales0 = pB_descales + 4 * block_start; + const float* pB_descales1 = pB_descales + 4 * block_count + 4 * block_start; + __m128 _out00 = k0 == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr, 0); + __m128 _out01 = k0 == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr + 4, 0); const signed char* pA = pAT; const float* pA_descales = pAT_descales; for (int k = 0; k < K; k += block_size) @@ -2732,14 +3466,17 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB_descales0 += 4; pB_descales1 += 4; } - pB = pB1; - pB_descales = pB_descales1; - __lsx_vst((__m128i)_out00, outptr + ii * max_jj + jj, 0); - __lsx_vst((__m128i)_out01, outptr + ii * max_jj + jj + 4, 0); + pB = pB1 + 4 * (full_K - k0 - max_kk0); + pB_descales = pB_descales1 + 4 * (block_count - block_start - tile_blocks); + __lsx_vst((__m128i)_out00, outptr, 0); + __lsx_vst((__m128i)_out01, outptr + 4, 0); + outptr += 8; } for (; jj + 3 < max_jj; jj += 4) { - __m128 _out0 = (__m128)__lsx_vldi(0); + pB += (size_t)4 * k0; + pB_descales += (size_t)4 * block_start; + __m128 _out0 = k0 == 0 ? (__m128)__lsx_vldi(0) : (__m128)__lsx_vld(outptr, 0); const signed char* pA = pAT; const float* pA_descales = pAT_descales; for (int k = 0; k < K; k += block_size) @@ -2779,13 +3516,18 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de _out0 = __lsx_vfmadd_s((__m128)__lsx_vffint_s_w(_sum0), __lsx_vfmul_s(_bscale, __lsx_vreplfr2vr_s(*pA_descales++)), _out0); pB_descales += 4; } - __lsx_vst((__m128i)_out0, outptr + ii * max_jj + jj, 0); + __lsx_vst((__m128i)_out0, outptr, 0); + outptr += 4; + pB += (size_t)4 * (full_K - k0 - max_kk0); + pB_descales += (size_t)4 * (block_count - block_start - tile_blocks); } #endif for (; jj + 1 < max_jj; jj += 2) { - float _out0 = 0.f; - float _out1 = 0.f; + pB += (size_t)2 * k0; + pB_descales += (size_t)2 * block_start; + float _out0 = k0 == 0 ? 0.f : outptr[0]; + float _out1 = k0 == 0 ? 0.f : outptr[1]; const signed char* pA = pAT; const float* pA_descales = pAT_descales; for (int k = 0; k < K; k += block_size) @@ -2826,12 +3568,17 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de _out1 += _sum1 * ascale * pB_descales[1]; pB_descales += 2; } - outptr[ii * max_jj + jj] = _out0; - outptr[ii * max_jj + jj + 1] = _out1; + outptr[0] = _out0; + outptr[1] = _out1; + outptr += 2; + pB += (size_t)2 * (full_K - k0 - max_kk0); + pB_descales += (size_t)2 * (block_count - block_start - tile_blocks); } for (; jj < max_jj; jj++) { - float _out0 = 0.f; + pB += k0; + pB_descales += block_start; + float _out0 = k0 == 0 ? 0.f : outptr[0]; const signed char* pA = pAT; const float* pA_descales = pAT_descales; for (int k = 0; k < K; k += block_size) @@ -2860,7 +3607,9 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } _out0 += _sum0 * *pA_descales++ * *pB_descales++; } - outptr[ii * max_jj + jj] = _out0; + *outptr++ = _out0; + pB += full_K - k0 - max_kk0; + pB_descales += block_count - block_start - tile_blocks; } pAT += A_hstep; pAT_descales += A_descales_hstep; @@ -2929,8 +3678,8 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b int jj = 0; #if __loongarch_asx - const __m256 _alpha256 = (__m256)__lasx_xvreplfr2vr_s(alpha); - const __m256 _beta256 = (__m256)__lasx_xvreplfr2vr_s(beta); + __m256 _alpha256 = (__m256)__lasx_xvreplfr2vr_s(alpha); + __m256 _beta256 = (__m256)__lasx_xvreplfr2vr_s(beta); for (; jj + 7 < max_jj; jj += 8) { __m256i _sum0 = __lasx_xvld(pp, 0); @@ -2982,7 +3731,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b { if (broadcast_type_C == 0) { - const __m256 _c = (__m256)__lasx_xvreplfr2vr_s(c0); + __m256 _c = (__m256)__lasx_xvreplfr2vr_s(c0); _f0 = __lasx_xvfadd_s(_f0, _c); _f1 = __lasx_xvfadd_s(_f1, _c); _f2 = __lasx_xvfadd_s(_f2, _c); @@ -3005,14 +3754,14 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b } if (broadcast_type_C == 3) { - const __m256 _c0 = (__m256)__lasx_xvld(pC, 0); - const __m256 _c1 = (__m256)__lasx_xvld(pC + c_hstep, 0); - const __m256 _c2 = (__m256)__lasx_xvld(pC + c_hstep * 2, 0); - const __m256 _c3 = (__m256)__lasx_xvld(pC + c_hstep * 3, 0); - const __m256 _c4 = (__m256)__lasx_xvld(pC + c_hstep * 4, 0); - const __m256 _c5 = (__m256)__lasx_xvld(pC + c_hstep * 5, 0); - const __m256 _c6 = (__m256)__lasx_xvld(pC + c_hstep * 6, 0); - const __m256 _c7 = (__m256)__lasx_xvld(pC + c_hstep * 7, 0); + __m256 _c0 = (__m256)__lasx_xvld(pC, 0); + __m256 _c1 = (__m256)__lasx_xvld(pC + c_hstep, 0); + __m256 _c2 = (__m256)__lasx_xvld(pC + c_hstep * 2, 0); + __m256 _c3 = (__m256)__lasx_xvld(pC + c_hstep * 3, 0); + __m256 _c4 = (__m256)__lasx_xvld(pC + c_hstep * 4, 0); + __m256 _c5 = (__m256)__lasx_xvld(pC + c_hstep * 5, 0); + __m256 _c6 = (__m256)__lasx_xvld(pC + c_hstep * 6, 0); + __m256 _c7 = (__m256)__lasx_xvld(pC + c_hstep * 7, 0); if (beta == 1.f) { _f0 = __lasx_xvfadd_s(_f0, _c0); @@ -3035,6 +3784,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _f6 = __lasx_xvfmadd_s(_c6, _beta256, _f6); _f7 = __lasx_xvfmadd_s(_c7, _beta256, _f7); } + pC += 8; } if (broadcast_type_C == 4) { @@ -3049,6 +3799,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _f5 = __lasx_xvfadd_s(_f5, _c); _f6 = __lasx_xvfadd_s(_f6, _c); _f7 = __lasx_xvfadd_s(_f7, _c); + pC += 8; } } if (alpha != 1.f) @@ -3078,12 +3829,10 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b p5 += 8; p6 += 8; p7 += 8; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 8; } -#endif - const __m128 _alpha128 = __lsx_vreplfr2vr_s(alpha); - const __m128 _beta128 = __lsx_vreplfr2vr_s(beta); +#endif // __loongarch_asx + __m128 _alpha128 = __lsx_vreplfr2vr_s(alpha); + __m128 _beta128 = __lsx_vreplfr2vr_s(beta); for (; jj + 3 < max_jj; jj += 4) { __m128i _sum0 = __lsx_vld(pp, 0); @@ -3119,7 +3868,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b { if (broadcast_type_C == 0) { - const __m128 _c = __lsx_vreplfr2vr_s(c0); + __m128 _c = __lsx_vreplfr2vr_s(c0); _f0 = __lsx_vfadd_s(_f0, _c); _f1 = __lsx_vfadd_s(_f1, _c); _f2 = __lsx_vfadd_s(_f2, _c); @@ -3142,14 +3891,14 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b } if (broadcast_type_C == 3) { - const __m128 _c0 = (__m128)__lsx_vld(pC, 0); - const __m128 _c1 = (__m128)__lsx_vld(pC + c_hstep, 0); - const __m128 _c2 = (__m128)__lsx_vld(pC + c_hstep * 2, 0); - const __m128 _c3 = (__m128)__lsx_vld(pC + c_hstep * 3, 0); - const __m128 _c4 = (__m128)__lsx_vld(pC + c_hstep * 4, 0); - const __m128 _c5 = (__m128)__lsx_vld(pC + c_hstep * 5, 0); - const __m128 _c6 = (__m128)__lsx_vld(pC + c_hstep * 6, 0); - const __m128 _c7 = (__m128)__lsx_vld(pC + c_hstep * 7, 0); + __m128 _c0 = (__m128)__lsx_vld(pC, 0); + __m128 _c1 = (__m128)__lsx_vld(pC + c_hstep, 0); + __m128 _c2 = (__m128)__lsx_vld(pC + c_hstep * 2, 0); + __m128 _c3 = (__m128)__lsx_vld(pC + c_hstep * 3, 0); + __m128 _c4 = (__m128)__lsx_vld(pC + c_hstep * 4, 0); + __m128 _c5 = (__m128)__lsx_vld(pC + c_hstep * 5, 0); + __m128 _c6 = (__m128)__lsx_vld(pC + c_hstep * 6, 0); + __m128 _c7 = (__m128)__lsx_vld(pC + c_hstep * 7, 0); if (beta == 1.f) { _f0 = __lsx_vfadd_s(_f0, _c0); @@ -3172,6 +3921,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _f6 = __lsx_vfmadd_s(_c6, _beta128, _f6); _f7 = __lsx_vfmadd_s(_c7, _beta128, _f7); } + pC += 4; } if (broadcast_type_C == 4) { @@ -3186,6 +3936,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _f5 = __lsx_vfadd_s(_f5, _c); _f6 = __lsx_vfadd_s(_f6, _c); _f7 = __lsx_vfadd_s(_f7, _c); + pC += 4; } } if (alpha != 1.f) @@ -3215,8 +3966,6 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b p5 += 4; p6 += 4; p7 += 4; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 4; } for (; jj + 1 < max_jj; jj += 2) { @@ -3249,7 +3998,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b { if (broadcast_type_C == 0) { - const __m128 _c = __lsx_vreplfr2vr_s(c0); + __m128 _c = __lsx_vreplfr2vr_s(c0); _f0 = __lsx_vfadd_s(_f0, _c); _f1 = __lsx_vfadd_s(_f1, _c); _f2 = __lsx_vfadd_s(_f2, _c); @@ -3272,14 +4021,14 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b } if (broadcast_type_C == 3) { - const __m128 _c0 = (__m128)__lsx_vldrepl_d(pC, 0); - const __m128 _c1 = (__m128)__lsx_vldrepl_d(pC + c_hstep, 0); - const __m128 _c2 = (__m128)__lsx_vldrepl_d(pC + c_hstep * 2, 0); - const __m128 _c3 = (__m128)__lsx_vldrepl_d(pC + c_hstep * 3, 0); - const __m128 _c4 = (__m128)__lsx_vldrepl_d(pC + c_hstep * 4, 0); - const __m128 _c5 = (__m128)__lsx_vldrepl_d(pC + c_hstep * 5, 0); - const __m128 _c6 = (__m128)__lsx_vldrepl_d(pC + c_hstep * 6, 0); - const __m128 _c7 = (__m128)__lsx_vldrepl_d(pC + c_hstep * 7, 0); + __m128 _c0 = (__m128)__lsx_vldrepl_d(pC, 0); + __m128 _c1 = (__m128)__lsx_vldrepl_d(pC + c_hstep, 0); + __m128 _c2 = (__m128)__lsx_vldrepl_d(pC + c_hstep * 2, 0); + __m128 _c3 = (__m128)__lsx_vldrepl_d(pC + c_hstep * 3, 0); + __m128 _c4 = (__m128)__lsx_vldrepl_d(pC + c_hstep * 4, 0); + __m128 _c5 = (__m128)__lsx_vldrepl_d(pC + c_hstep * 5, 0); + __m128 _c6 = (__m128)__lsx_vldrepl_d(pC + c_hstep * 6, 0); + __m128 _c7 = (__m128)__lsx_vldrepl_d(pC + c_hstep * 7, 0); if (beta == 1.f) { _f0 = __lsx_vfadd_s(_f0, _c0); @@ -3302,6 +4051,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _f6 = __lsx_vfmadd_s(_c6, _beta128, _f6); _f7 = __lsx_vfmadd_s(_c7, _beta128, _f7); } + pC += 2; } if (broadcast_type_C == 4) { @@ -3316,6 +4066,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _f5 = __lsx_vfadd_s(_f5, _c); _f6 = __lsx_vfadd_s(_f6, _c); _f7 = __lsx_vfadd_s(_f7, _c); + pC += 2; } } if (alpha != 1.f) @@ -3345,8 +4096,6 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b p5 += 2; p6 += 2; p7 += 2; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 2; } for (; jj < max_jj; jj++) { @@ -3357,7 +4106,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b { if (broadcast_type_C == 0) { - const __m128 _c = __lsx_vreplfr2vr_s(c0); + __m128 _c = __lsx_vreplfr2vr_s(c0); _f0 = __lsx_vfadd_s(_f0, _c); _f4 = __lsx_vfadd_s(_f4, _c); } @@ -3386,12 +4135,14 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _f0 = __lsx_vfmadd_s((__m128)_c0, _beta128, _f0); _f4 = __lsx_vfmadd_s((__m128)_c4, _beta128, _f4); } + pC++; } if (broadcast_type_C == 4) { - const __m128 _c = __lsx_vreplfr2vr_s(beta == 1.f ? pC[0] : pC[0] * beta); + __m128 _c = __lsx_vreplfr2vr_s(beta == 1.f ? pC[0] : pC[0] * beta); _f0 = __lsx_vfadd_s(_f0, _c); _f4 = __lsx_vfadd_s(_f4, _c); + pC++; } } if (alpha != 1.f) @@ -3415,11 +4166,8 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b p5++; p6++; p7++; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC++; } } -#endif // __loongarch_sx for (; ii + 3 < max_ii; ii += 4) { float* p0 = outptr + (size_t)(i + ii) * N + j; @@ -3460,23 +4208,24 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b int jj = 0; #if __loongarch_asx - const __m256 _alpha256 = (__m256)__lasx_xvreplfr2vr_s(alpha); - const __m256 _beta256 = (__m256)__lasx_xvreplfr2vr_s(beta); + __m256 _alpha256 = (__m256)__lasx_xvreplfr2vr_s(alpha); + __m256 _beta256 = (__m256)__lasx_xvreplfr2vr_s(beta); for (; jj + 15 < max_jj; jj += 16) { - __m256 _f00 = (__m256)__lasx_xvld(pp + (ii + 0) * max_jj + jj, 0); - __m256 _f01 = (__m256)__lasx_xvld(pp + (ii + 0) * max_jj + jj + 8, 0); - __m256 _f10 = (__m256)__lasx_xvld(pp + (ii + 1) * max_jj + jj, 0); - __m256 _f11 = (__m256)__lasx_xvld(pp + (ii + 1) * max_jj + jj + 8, 0); - __m256 _f20 = (__m256)__lasx_xvld(pp + (ii + 2) * max_jj + jj, 0); - __m256 _f21 = (__m256)__lasx_xvld(pp + (ii + 2) * max_jj + jj + 8, 0); - __m256 _f30 = (__m256)__lasx_xvld(pp + (ii + 3) * max_jj + jj, 0); - __m256 _f31 = (__m256)__lasx_xvld(pp + (ii + 3) * max_jj + jj + 8, 0); + __m256 _f00 = (__m256)__lasx_xvld(pp, 0); + __m256 _f01 = (__m256)__lasx_xvld(pp + 8, 0); + __m256 _f10 = (__m256)__lasx_xvld(pp + 16, 0); + __m256 _f11 = (__m256)__lasx_xvld(pp + 24, 0); + __m256 _f20 = (__m256)__lasx_xvld(pp + 32, 0); + __m256 _f21 = (__m256)__lasx_xvld(pp + 40, 0); + __m256 _f30 = (__m256)__lasx_xvld(pp + 48, 0); + __m256 _f31 = (__m256)__lasx_xvld(pp + 56, 0); + pp += 64; if (pC) { if (broadcast_type_C == 0) { - const __m256 _c = (__m256)__lasx_xvreplfr2vr_s(c0); + __m256 _c = (__m256)__lasx_xvreplfr2vr_s(c0); _f00 = __lasx_xvfadd_s(_f00, _c); _f01 = __lasx_xvfadd_s(_f01, _c); _f10 = __lasx_xvfadd_s(_f10, _c); @@ -3488,10 +4237,10 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b } if (broadcast_type_C == 1 || broadcast_type_C == 2) { - const __m256 _c0 = (__m256)__lasx_xvreplfr2vr_s(c0); - const __m256 _c1 = (__m256)__lasx_xvreplfr2vr_s(c1); - const __m256 _c2 = (__m256)__lasx_xvreplfr2vr_s(c2); - const __m256 _c3 = (__m256)__lasx_xvreplfr2vr_s(c3); + __m256 _c0 = (__m256)__lasx_xvreplfr2vr_s(c0); + __m256 _c1 = (__m256)__lasx_xvreplfr2vr_s(c1); + __m256 _c2 = (__m256)__lasx_xvreplfr2vr_s(c2); + __m256 _c3 = (__m256)__lasx_xvreplfr2vr_s(c3); _f00 = __lasx_xvfadd_s(_f00, _c0); _f01 = __lasx_xvfadd_s(_f01, _c0); _f10 = __lasx_xvfadd_s(_f10, _c1); @@ -3503,14 +4252,14 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b } if (broadcast_type_C == 3) { - const __m256 _c00 = (__m256)__lasx_xvld(pC, 0); - const __m256 _c01 = (__m256)__lasx_xvld(pC + 8, 0); - const __m256 _c10 = (__m256)__lasx_xvld(pC + c_hstep, 0); - const __m256 _c11 = (__m256)__lasx_xvld(pC + c_hstep + 8, 0); - const __m256 _c20 = (__m256)__lasx_xvld(pC + c_hstep * 2, 0); - const __m256 _c21 = (__m256)__lasx_xvld(pC + c_hstep * 2 + 8, 0); - const __m256 _c30 = (__m256)__lasx_xvld(pC + c_hstep * 3, 0); - const __m256 _c31 = (__m256)__lasx_xvld(pC + c_hstep * 3 + 8, 0); + __m256 _c00 = (__m256)__lasx_xvld(pC, 0); + __m256 _c01 = (__m256)__lasx_xvld(pC + 8, 0); + __m256 _c10 = (__m256)__lasx_xvld(pC + c_hstep, 0); + __m256 _c11 = (__m256)__lasx_xvld(pC + c_hstep + 8, 0); + __m256 _c20 = (__m256)__lasx_xvld(pC + c_hstep * 2, 0); + __m256 _c21 = (__m256)__lasx_xvld(pC + c_hstep * 2 + 8, 0); + __m256 _c30 = (__m256)__lasx_xvld(pC + c_hstep * 3, 0); + __m256 _c31 = (__m256)__lasx_xvld(pC + c_hstep * 3 + 8, 0); if (beta == 1.f) { _f00 = __lasx_xvfadd_s(_f00, _c00); @@ -3533,6 +4282,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _f30 = __lasx_xvfmadd_s(_c30, _beta256, _f30); _f31 = __lasx_xvfmadd_s(_c31, _beta256, _f31); } + pC += 16; } if (broadcast_type_C == 4) { @@ -3551,6 +4301,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _f21 = __lasx_xvfadd_s(_f21, _c1); _f30 = __lasx_xvfadd_s(_f30, _c0); _f31 = __lasx_xvfadd_s(_f31, _c1); + pC += 16; } } if (alpha != 1.f) @@ -3576,15 +4327,14 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b p1 += 16; p2 += 16; p3 += 16; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 16; } for (; jj + 7 < max_jj; jj += 8) { - __m256 _f0 = (__m256)__lasx_xvld(pp + (ii + 0) * max_jj + jj, 0); - __m256 _f1 = (__m256)__lasx_xvld(pp + (ii + 1) * max_jj + jj, 0); - __m256 _f2 = (__m256)__lasx_xvld(pp + (ii + 2) * max_jj + jj, 0); - __m256 _f3 = (__m256)__lasx_xvld(pp + (ii + 3) * max_jj + jj, 0); + __m256 _f0 = (__m256)__lasx_xvld(pp, 0); + __m256 _f1 = (__m256)__lasx_xvld(pp + 8, 0); + __m256 _f2 = (__m256)__lasx_xvld(pp + 16, 0); + __m256 _f3 = (__m256)__lasx_xvld(pp + 24, 0); + pp += 32; if (pC) { if (broadcast_type_C == 0) @@ -3621,6 +4371,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _f2 = __lasx_xvfmadd_s(_c2, _beta256, _f2); _f3 = __lasx_xvfmadd_s(_c3, _beta256, _f3); } + pC += 8; } if (broadcast_type_C == 4) { @@ -3631,6 +4382,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _f1 = __lasx_xvfadd_s(_f1, _c); _f2 = __lasx_xvfadd_s(_f2, _c); _f3 = __lasx_xvfadd_s(_f3, _c); + pC += 8; } } if (alpha != 1.f) @@ -3648,28 +4400,26 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b p1 += 8; p2 += 8; p3 += 8; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 8; } -#endif -#if __loongarch_sx - const __m128 _alpha128 = __lsx_vreplfr2vr_s(alpha); - const __m128 _beta128 = __lsx_vreplfr2vr_s(beta); +#endif // __loongarch_asx + __m128 _alpha128 = __lsx_vreplfr2vr_s(alpha); + __m128 _beta128 = __lsx_vreplfr2vr_s(beta); for (; jj + 7 < max_jj; jj += 8) { - __m128 _f00 = (__m128)__lsx_vld(pp + (ii + 0) * max_jj + jj, 0); - __m128 _f01 = (__m128)__lsx_vld(pp + (ii + 0) * max_jj + jj + 4, 0); - __m128 _f10 = (__m128)__lsx_vld(pp + (ii + 1) * max_jj + jj, 0); - __m128 _f11 = (__m128)__lsx_vld(pp + (ii + 1) * max_jj + jj + 4, 0); - __m128 _f20 = (__m128)__lsx_vld(pp + (ii + 2) * max_jj + jj, 0); - __m128 _f21 = (__m128)__lsx_vld(pp + (ii + 2) * max_jj + jj + 4, 0); - __m128 _f30 = (__m128)__lsx_vld(pp + (ii + 3) * max_jj + jj, 0); - __m128 _f31 = (__m128)__lsx_vld(pp + (ii + 3) * max_jj + jj + 4, 0); + __m128 _f00 = (__m128)__lsx_vld(pp, 0); + __m128 _f01 = (__m128)__lsx_vld(pp + 4, 0); + __m128 _f10 = (__m128)__lsx_vld(pp + 8, 0); + __m128 _f11 = (__m128)__lsx_vld(pp + 12, 0); + __m128 _f20 = (__m128)__lsx_vld(pp + 16, 0); + __m128 _f21 = (__m128)__lsx_vld(pp + 20, 0); + __m128 _f30 = (__m128)__lsx_vld(pp + 24, 0); + __m128 _f31 = (__m128)__lsx_vld(pp + 28, 0); + pp += 32; if (pC) { if (broadcast_type_C == 0) { - const __m128 _c = __lsx_vreplfr2vr_s(c0); + __m128 _c = __lsx_vreplfr2vr_s(c0); _f00 = __lsx_vfadd_s(_f00, _c); _f01 = __lsx_vfadd_s(_f01, _c); _f10 = __lsx_vfadd_s(_f10, _c); @@ -3681,10 +4431,10 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b } if (broadcast_type_C == 1 || broadcast_type_C == 2) { - const __m128 _c0 = __lsx_vreplfr2vr_s(c0); - const __m128 _c1 = __lsx_vreplfr2vr_s(c1); - const __m128 _c2 = __lsx_vreplfr2vr_s(c2); - const __m128 _c3 = __lsx_vreplfr2vr_s(c3); + __m128 _c0 = __lsx_vreplfr2vr_s(c0); + __m128 _c1 = __lsx_vreplfr2vr_s(c1); + __m128 _c2 = __lsx_vreplfr2vr_s(c2); + __m128 _c3 = __lsx_vreplfr2vr_s(c3); _f00 = __lsx_vfadd_s(_f00, _c0); _f01 = __lsx_vfadd_s(_f01, _c0); _f10 = __lsx_vfadd_s(_f10, _c1); @@ -3696,14 +4446,14 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b } if (broadcast_type_C == 3) { - const __m128 _c00 = (__m128)__lsx_vld(pC, 0); - const __m128 _c01 = (__m128)__lsx_vld(pC + 4, 0); - const __m128 _c10 = (__m128)__lsx_vld(pC + c_hstep, 0); - const __m128 _c11 = (__m128)__lsx_vld(pC + c_hstep + 4, 0); - const __m128 _c20 = (__m128)__lsx_vld(pC + c_hstep * 2, 0); - const __m128 _c21 = (__m128)__lsx_vld(pC + c_hstep * 2 + 4, 0); - const __m128 _c30 = (__m128)__lsx_vld(pC + c_hstep * 3, 0); - const __m128 _c31 = (__m128)__lsx_vld(pC + c_hstep * 3 + 4, 0); + __m128 _c00 = (__m128)__lsx_vld(pC, 0); + __m128 _c01 = (__m128)__lsx_vld(pC + 4, 0); + __m128 _c10 = (__m128)__lsx_vld(pC + c_hstep, 0); + __m128 _c11 = (__m128)__lsx_vld(pC + c_hstep + 4, 0); + __m128 _c20 = (__m128)__lsx_vld(pC + c_hstep * 2, 0); + __m128 _c21 = (__m128)__lsx_vld(pC + c_hstep * 2 + 4, 0); + __m128 _c30 = (__m128)__lsx_vld(pC + c_hstep * 3, 0); + __m128 _c31 = (__m128)__lsx_vld(pC + c_hstep * 3 + 4, 0); if (beta == 1.f) { _f00 = __lsx_vfadd_s(_f00, _c00); @@ -3726,6 +4476,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _f30 = __lsx_vfmadd_s(_c30, _beta128, _f30); _f31 = __lsx_vfmadd_s(_c31, _beta128, _f31); } + pC += 8; } if (broadcast_type_C == 4) { @@ -3744,6 +4495,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _f21 = __lsx_vfadd_s(_f21, _c1); _f30 = __lsx_vfadd_s(_f30, _c0); _f31 = __lsx_vfadd_s(_f31, _c1); + pC += 8; } } if (alpha != 1.f) @@ -3769,15 +4521,14 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b p1 += 8; p2 += 8; p3 += 8; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 8; } for (; jj + 3 < max_jj; jj += 4) { - __m128 _f0 = (__m128)__lsx_vld(pp + (ii + 0) * max_jj + jj, 0); - __m128 _f1 = (__m128)__lsx_vld(pp + (ii + 1) * max_jj + jj, 0); - __m128 _f2 = (__m128)__lsx_vld(pp + (ii + 2) * max_jj + jj, 0); - __m128 _f3 = (__m128)__lsx_vld(pp + (ii + 3) * max_jj + jj, 0); + __m128 _f0 = (__m128)__lsx_vld(pp, 0); + __m128 _f1 = (__m128)__lsx_vld(pp + 4, 0); + __m128 _f2 = (__m128)__lsx_vld(pp + 8, 0); + __m128 _f3 = (__m128)__lsx_vld(pp + 12, 0); + pp += 16; if (pC) { if (broadcast_type_C == 0) @@ -3814,6 +4565,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _f2 = __lsx_vfmadd_s(_c2, _beta128, _f2); _f3 = __lsx_vfmadd_s(_c3, _beta128, _f3); } + pC += 4; } if (broadcast_type_C == 4) { @@ -3824,6 +4576,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _f1 = __lsx_vfadd_s(_f1, _c); _f2 = __lsx_vfadd_s(_f2, _c); _f3 = __lsx_vfadd_s(_f3, _c); + pC += 4; } } if (alpha != 1.f) @@ -3841,15 +4594,14 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b p1 += 4; p2 += 4; p3 += 4; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 4; } for (; jj + 1 < max_jj; jj += 2) { - __m128 _f0 = (__m128)__lsx_vldrepl_d(pp + (ii + 0) * max_jj + jj, 0); - __m128 _f1 = (__m128)__lsx_vldrepl_d(pp + (ii + 1) * max_jj + jj, 0); - __m128 _f2 = (__m128)__lsx_vldrepl_d(pp + (ii + 2) * max_jj + jj, 0); - __m128 _f3 = (__m128)__lsx_vldrepl_d(pp + (ii + 3) * max_jj + jj, 0); + __m128 _f0 = (__m128)__lsx_vldrepl_d(pp, 0); + __m128 _f1 = (__m128)__lsx_vldrepl_d(pp + 2, 0); + __m128 _f2 = (__m128)__lsx_vldrepl_d(pp + 4, 0); + __m128 _f3 = (__m128)__lsx_vldrepl_d(pp + 6, 0); + pp += 8; if (pC) { if (broadcast_type_C == 0) @@ -3887,6 +4639,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _f2 = __lsx_vfmadd_s(_c2, _beta128, _f2); _f3 = __lsx_vfmadd_s(_c3, _beta128, _f3); } + pC += 2; } if (broadcast_type_C == 4) { @@ -3897,6 +4650,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _f1 = __lsx_vfadd_s(_f1, _c); _f2 = __lsx_vfadd_s(_f2, _c); _f3 = __lsx_vfadd_s(_f3, _c); + pC += 2; } } if (alpha != 1.f) @@ -3914,17 +4668,11 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b p1 += 2; p2 += 2; p3 += 2; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 2; } -#endif -#if __loongarch_sx for (; jj < max_jj; jj++) { - __m128i _fi = __lsx_vreplgr2vr_w(((const int*)(pp + (ii + 0) * max_jj + jj))[0]); - _fi = __lsx_vinsgr2vr_w(_fi, ((const int*)(pp + (ii + 1) * max_jj + jj))[0], 1); - _fi = __lsx_vinsgr2vr_w(_fi, ((const int*)(pp + (ii + 2) * max_jj + jj))[0], 2); - _fi = __lsx_vinsgr2vr_w(_fi, ((const int*)(pp + (ii + 3) * max_jj + jj))[0], 3); + __m128i _fi = __lsx_vld(pp, 0); + pp += 4; __m128 _f0 = (__m128)_fi; if (pC) { @@ -3942,9 +4690,13 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _f0 = __lsx_vfadd_s(_f0, (__m128)_c0); else _f0 = __lsx_vfmadd_s((__m128)_c0, _beta128, _f0); + pC++; } if (broadcast_type_C == 4) + { _f0 = __lsx_vfadd_s(_f0, __lsx_vreplfr2vr_s(beta == 1.f ? pC[0] : pC[0] * beta)); + pC++; + } } if (alpha != 1.f) _f0 = __lsx_vfmul_s(_f0, _alpha128); @@ -3956,78 +4708,9 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b p1++; p2++; p3++; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC++; - } -#else - for (; jj < max_jj; jj++) - { - float f0 = pp[(ii + 0) * max_jj + jj]; - float f1 = pp[(ii + 1) * max_jj + jj]; - float f2 = pp[(ii + 2) * max_jj + jj]; - float f3 = pp[(ii + 3) * max_jj + jj]; - if (pC) - { - if (broadcast_type_C == 0) - { - f0 += c0; - f1 += c0; - f2 += c0; - f3 += c0; - } - if (broadcast_type_C == 1 || broadcast_type_C == 2) - { - f0 += c0; - f1 += c1; - f2 += c2; - f3 += c3; - } - if (broadcast_type_C == 3) - { - if (beta == 1.f) - { - f0 += pC[0]; - f1 += pC[c_hstep]; - f2 += pC[c_hstep * 2]; - f3 += pC[c_hstep * 3]; - } - else - { - f0 += pC[0] * beta; - f1 += pC[c_hstep] * beta; - f2 += pC[c_hstep * 2] * beta; - f3 += pC[c_hstep * 3] * beta; - } - } - if (broadcast_type_C == 4) - { - float c = beta == 1.f ? pC[0] : pC[0] * beta; - f0 += c; - f1 += c; - f2 += c; - f3 += c; - } - } - if (alpha != 1.f) - { - f0 *= alpha; - f1 *= alpha; - f2 *= alpha; - f3 *= alpha; - } - p0[0] = f0; - p1[0] = f1; - p2[0] = f2; - p3[0] = f3; - p0++; - p1++; - p2++; - p3++; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC++; } -#endif } +#endif // __loongarch_sx for (; ii + 1 < max_ii; ii += 2) { float* p0 = outptr + (size_t)(i + ii) * N + j; @@ -4059,20 +4742,21 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b } int jj = 0; +#if __loongarch_sx #if __loongarch_asx - const __m256 _alpha256 = (__m256)__lasx_xvreplfr2vr_s(alpha); - const __m256 _beta256 = (__m256)__lasx_xvreplfr2vr_s(beta); + __m256 _alpha256 = (__m256)__lasx_xvreplfr2vr_s(alpha); + __m256 _beta256 = (__m256)__lasx_xvreplfr2vr_s(beta); for (; jj + 15 < max_jj; jj += 16) { - __m256 _f00 = (__m256)__lasx_xvld(pp + (ii + 0) * max_jj + jj, 0); - __m256 _f01 = (__m256)__lasx_xvld(pp + (ii + 0) * max_jj + jj + 8, 0); - __m256 _f10 = (__m256)__lasx_xvld(pp + (ii + 1) * max_jj + jj, 0); - __m256 _f11 = (__m256)__lasx_xvld(pp + (ii + 1) * max_jj + jj + 8, 0); + __m256 _f00 = (__m256)__lasx_xvld(pp, 0); + __m256 _f01 = (__m256)__lasx_xvld(pp + 8, 0); + __m256 _f10 = (__m256)__lasx_xvld(pp + 16, 0); + __m256 _f11 = (__m256)__lasx_xvld(pp + 24, 0); if (pC) { if (broadcast_type_C == 0) { - const __m256 _c = (__m256)__lasx_xvreplfr2vr_s(c0); + __m256 _c = (__m256)__lasx_xvreplfr2vr_s(c0); _f00 = __lasx_xvfadd_s(_f00, _c); _f01 = __lasx_xvfadd_s(_f01, _c); _f10 = __lasx_xvfadd_s(_f10, _c); @@ -4080,8 +4764,8 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b } if (broadcast_type_C == 1 || broadcast_type_C == 2) { - const __m256 _c0 = (__m256)__lasx_xvreplfr2vr_s(c0); - const __m256 _c1 = (__m256)__lasx_xvreplfr2vr_s(c1); + __m256 _c0 = (__m256)__lasx_xvreplfr2vr_s(c0); + __m256 _c1 = (__m256)__lasx_xvreplfr2vr_s(c1); _f00 = __lasx_xvfadd_s(_f00, _c0); _f01 = __lasx_xvfadd_s(_f01, _c0); _f10 = __lasx_xvfadd_s(_f10, _c1); @@ -4089,10 +4773,10 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b } if (broadcast_type_C == 3) { - const __m256 _c00 = (__m256)__lasx_xvld(pC, 0); - const __m256 _c01 = (__m256)__lasx_xvld(pC + 8, 0); - const __m256 _c10 = (__m256)__lasx_xvld(pC + c_hstep, 0); - const __m256 _c11 = (__m256)__lasx_xvld(pC + c_hstep + 8, 0); + __m256 _c00 = (__m256)__lasx_xvld(pC, 0); + __m256 _c01 = (__m256)__lasx_xvld(pC + 8, 0); + __m256 _c10 = (__m256)__lasx_xvld(pC + c_hstep, 0); + __m256 _c11 = (__m256)__lasx_xvld(pC + c_hstep + 8, 0); if (beta == 1.f) { _f00 = __lasx_xvfadd_s(_f00, _c00); @@ -4107,6 +4791,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _f10 = __lasx_xvfmadd_s(_c10, _beta256, _f10); _f11 = __lasx_xvfmadd_s(_c11, _beta256, _f11); } + pC += 16; } if (broadcast_type_C == 4) { @@ -4121,6 +4806,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _f01 = __lasx_xvfadd_s(_f01, _c1); _f10 = __lasx_xvfadd_s(_f10, _c0); _f11 = __lasx_xvfadd_s(_f11, _c1); + pC += 16; } } if (alpha != 1.f) @@ -4136,13 +4822,12 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b __lasx_xvst(_f11, p1 + 8, 0); p0 += 16; p1 += 16; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 16; + pp += 32; } for (; jj + 7 < max_jj; jj += 8) { - __m256 _f0 = (__m256)__lasx_xvld(pp + (ii + 0) * max_jj + jj, 0); - __m256 _f1 = (__m256)__lasx_xvld(pp + (ii + 1) * max_jj + jj, 0); + __m256 _f0 = (__m256)__lasx_xvld(pp, 0); + __m256 _f1 = (__m256)__lasx_xvld(pp + 8, 0); if (pC) { if (broadcast_type_C == 0) @@ -4169,6 +4854,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _f0 = __lasx_xvfmadd_s(_c0, _beta256, _f0); _f1 = __lasx_xvfmadd_s(_c1, _beta256, _f1); } + pC += 8; } if (broadcast_type_C == 4) { @@ -4177,6 +4863,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _c = __lasx_xvfmul_s(_c, _beta256); _f0 = __lasx_xvfadd_s(_f0, _c); _f1 = __lasx_xvfadd_s(_f1, _c); + pC += 8; } } if (alpha != 1.f) @@ -4188,24 +4875,22 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b __lasx_xvst(_f1, p1, 0); p0 += 8; p1 += 8; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 8; + pp += 16; } -#endif -#if __loongarch_sx - const __m128 _alpha128 = __lsx_vreplfr2vr_s(alpha); - const __m128 _beta128 = __lsx_vreplfr2vr_s(beta); +#endif // __loongarch_asx + __m128 _alpha128 = __lsx_vreplfr2vr_s(alpha); + __m128 _beta128 = __lsx_vreplfr2vr_s(beta); for (; jj + 7 < max_jj; jj += 8) { - __m128 _f00 = (__m128)__lsx_vld(pp + (ii + 0) * max_jj + jj, 0); - __m128 _f01 = (__m128)__lsx_vld(pp + (ii + 0) * max_jj + jj + 4, 0); - __m128 _f10 = (__m128)__lsx_vld(pp + (ii + 1) * max_jj + jj, 0); - __m128 _f11 = (__m128)__lsx_vld(pp + (ii + 1) * max_jj + jj + 4, 0); + __m128 _f00 = (__m128)__lsx_vld(pp, 0); + __m128 _f01 = (__m128)__lsx_vld(pp + 4, 0); + __m128 _f10 = (__m128)__lsx_vld(pp + 8, 0); + __m128 _f11 = (__m128)__lsx_vld(pp + 12, 0); if (pC) { if (broadcast_type_C == 0) { - const __m128 _c = __lsx_vreplfr2vr_s(c0); + __m128 _c = __lsx_vreplfr2vr_s(c0); _f00 = __lsx_vfadd_s(_f00, _c); _f01 = __lsx_vfadd_s(_f01, _c); _f10 = __lsx_vfadd_s(_f10, _c); @@ -4213,8 +4898,8 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b } if (broadcast_type_C == 1 || broadcast_type_C == 2) { - const __m128 _c0 = __lsx_vreplfr2vr_s(c0); - const __m128 _c1 = __lsx_vreplfr2vr_s(c1); + __m128 _c0 = __lsx_vreplfr2vr_s(c0); + __m128 _c1 = __lsx_vreplfr2vr_s(c1); _f00 = __lsx_vfadd_s(_f00, _c0); _f01 = __lsx_vfadd_s(_f01, _c0); _f10 = __lsx_vfadd_s(_f10, _c1); @@ -4222,10 +4907,10 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b } if (broadcast_type_C == 3) { - const __m128 _c00 = (__m128)__lsx_vld(pC, 0); - const __m128 _c01 = (__m128)__lsx_vld(pC + 4, 0); - const __m128 _c10 = (__m128)__lsx_vld(pC + c_hstep, 0); - const __m128 _c11 = (__m128)__lsx_vld(pC + c_hstep + 4, 0); + __m128 _c00 = (__m128)__lsx_vld(pC, 0); + __m128 _c01 = (__m128)__lsx_vld(pC + 4, 0); + __m128 _c10 = (__m128)__lsx_vld(pC + c_hstep, 0); + __m128 _c11 = (__m128)__lsx_vld(pC + c_hstep + 4, 0); if (beta == 1.f) { _f00 = __lsx_vfadd_s(_f00, _c00); @@ -4240,6 +4925,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _f10 = __lsx_vfmadd_s(_c10, _beta128, _f10); _f11 = __lsx_vfmadd_s(_c11, _beta128, _f11); } + pC += 8; } if (broadcast_type_C == 4) { @@ -4254,6 +4940,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _f01 = __lsx_vfadd_s(_f01, _c1); _f10 = __lsx_vfadd_s(_f10, _c0); _f11 = __lsx_vfadd_s(_f11, _c1); + pC += 8; } } if (alpha != 1.f) @@ -4269,13 +4956,12 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b __lsx_vst((__m128i)_f11, p1 + 4, 0); p0 += 8; p1 += 8; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 8; + pp += 16; } for (; jj + 3 < max_jj; jj += 4) { - __m128 _f0 = (__m128)__lsx_vld(pp + (ii + 0) * max_jj + jj, 0); - __m128 _f1 = (__m128)__lsx_vld(pp + (ii + 1) * max_jj + jj, 0); + __m128 _f0 = (__m128)__lsx_vld(pp, 0); + __m128 _f1 = (__m128)__lsx_vld(pp + 4, 0); if (pC) { if (broadcast_type_C == 0) @@ -4302,6 +4988,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _f0 = __lsx_vfmadd_s(_c0, _beta128, _f0); _f1 = __lsx_vfmadd_s(_c1, _beta128, _f1); } + pC += 4; } if (broadcast_type_C == 4) { @@ -4310,6 +4997,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _c = __lsx_vfmul_s(_c, _beta128); _f0 = __lsx_vfadd_s(_f0, _c); _f1 = __lsx_vfadd_s(_f1, _c); + pC += 4; } } if (alpha != 1.f) @@ -4321,13 +5009,12 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b __lsx_vst((__m128i)_f1, p1, 0); p0 += 4; p1 += 4; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 4; + pp += 8; } for (; jj + 1 < max_jj; jj += 2) { - __m128 _f0 = (__m128)__lsx_vldrepl_d(pp + (ii + 0) * max_jj + jj, 0); - __m128 _f1 = (__m128)__lsx_vldrepl_d(pp + (ii + 1) * max_jj + jj, 0); + __m128 _f0 = (__m128)__lsx_vldrepl_d(pp, 0); + __m128 _f1 = (__m128)__lsx_vldrepl_d(pp + 2, 0); if (pC) { if (broadcast_type_C == 0) @@ -4355,6 +5042,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _f0 = __lsx_vfmadd_s(_c0, _beta128, _f0); _f1 = __lsx_vfmadd_s(_c1, _beta128, _f1); } + pC += 2; } if (broadcast_type_C == 4) { @@ -4363,6 +5051,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _c = __lsx_vfmul_s(_c, _beta128); _f0 = __lsx_vfadd_s(_f0, _c); _f1 = __lsx_vfadd_s(_f1, _c); + pC += 2; } } if (alpha != 1.f) @@ -4374,14 +5063,79 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b __lsx_vstelm_d((__m128i)_f1, p1, 0, 0); p0 += 2; p1 += 2; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 2; + pp += 4; + } +#endif // __loongarch_sx + for (; jj + 1 < max_jj; jj += 2) + { + float f00 = pp[0]; + float f01 = pp[1]; + float f10 = pp[2]; + float f11 = pp[3]; + if (pC) + { + if (broadcast_type_C == 0) + { + f00 += c0; + f01 += c0; + f10 += c0; + f11 += c0; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + f00 += c0; + f01 += c0; + f10 += c1; + f11 += c1; + } + if (broadcast_type_C == 3) + { + if (beta == 1.f) + { + f00 += pC[0]; + f01 += pC[1]; + f10 += pC[c_hstep]; + f11 += pC[c_hstep + 1]; + } + else + { + f00 += pC[0] * beta; + f01 += pC[1] * beta; + f10 += pC[c_hstep] * beta; + f11 += pC[c_hstep + 1] * beta; + } + pC += 2; + } + if (broadcast_type_C == 4) + { + const float cc0 = beta == 1.f ? pC[0] : pC[0] * beta; + const float cc1 = beta == 1.f ? pC[1] : pC[1] * beta; + f00 += cc0; + f01 += cc1; + f10 += cc0; + f11 += cc1; + pC += 2; + } + } + if (alpha != 1.f) + { + f00 *= alpha; + f01 *= alpha; + f10 *= alpha; + f11 *= alpha; + } + p0[0] = f00; + p0[1] = f01; + p1[0] = f10; + p1[1] = f11; + p0 += 2; + p1 += 2; + pp += 4; } -#endif for (; jj < max_jj; jj++) { - float f0 = pp[(ii + 0) * max_jj + jj]; - float f1 = pp[(ii + 1) * max_jj + jj]; + float f0 = pp[0]; + float f1 = pp[1]; if (pC) { if (broadcast_type_C == 0) @@ -4406,12 +5160,14 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b f0 += pC[0] * beta; f1 += pC[c_hstep] * beta; } + pC++; } if (broadcast_type_C == 4) { float c = beta == 1.f ? pC[0] : pC[0] * beta; f0 += c; f1 += c; + pC++; } } if (alpha != 1.f) @@ -4423,8 +5179,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b p1[0] = f1; p0++; p1++; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC++; + pp += 2; } } for (; ii < max_ii; ii++) @@ -4446,25 +5201,26 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b } int jj = 0; +#if __loongarch_sx #if __loongarch_asx - const __m256 _alpha256 = (__m256)__lasx_xvreplfr2vr_s(alpha); - const __m256 _beta256 = (__m256)__lasx_xvreplfr2vr_s(beta); + __m256 _alpha256 = (__m256)__lasx_xvreplfr2vr_s(alpha); + __m256 _beta256 = (__m256)__lasx_xvreplfr2vr_s(beta); for (; jj + 15 < max_jj; jj += 16) { - __m256 _f0 = (__m256)__lasx_xvld(pp + ii * max_jj + jj, 0); - __m256 _f1 = (__m256)__lasx_xvld(pp + ii * max_jj + jj + 8, 0); + __m256 _f0 = (__m256)__lasx_xvld(pp, 0); + __m256 _f1 = (__m256)__lasx_xvld(pp + 8, 0); if (pC) { if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) { - const __m256 _c = (__m256)__lasx_xvreplfr2vr_s(c0); + __m256 _c = (__m256)__lasx_xvreplfr2vr_s(c0); _f0 = __lasx_xvfadd_s(_f0, _c); _f1 = __lasx_xvfadd_s(_f1, _c); } if (broadcast_type_C == 3 || broadcast_type_C == 4) { - const __m256 _c0 = (__m256)__lasx_xvld(pC, 0); - const __m256 _c1 = (__m256)__lasx_xvld(pC + 8, 0); + __m256 _c0 = (__m256)__lasx_xvld(pC, 0); + __m256 _c1 = (__m256)__lasx_xvld(pC + 8, 0); if (beta == 1.f) { _f0 = __lasx_xvfadd_s(_f0, _c0); @@ -4475,6 +5231,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _f0 = __lasx_xvfmadd_s(_c0, _beta256, _f0); _f1 = __lasx_xvfmadd_s(_c1, _beta256, _f1); } + pC += 16; } } if (alpha != 1.f) @@ -4485,12 +5242,11 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b __lasx_xvst(_f0, p0, 0); __lasx_xvst(_f1, p0 + 8, 0); p0 += 16; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 16; + pp += 16; } for (; jj + 7 < max_jj; jj += 8) { - __m256 _f0 = (__m256)__lasx_xvld(pp + ii * max_jj + jj, 0); + __m256 _f0 = (__m256)__lasx_xvld(pp, 0); if (pC) { if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) @@ -4502,34 +5258,33 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _f0 = __lasx_xvfadd_s(_f0, _c0); else _f0 = __lasx_xvfmadd_s(_c0, _beta256, _f0); + pC += 8; } } if (alpha != 1.f) _f0 = __lasx_xvfmul_s(_f0, _alpha256); __lasx_xvst(_f0, p0, 0); p0 += 8; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 8; + pp += 8; } -#endif -#if __loongarch_sx - const __m128 _alpha128 = __lsx_vreplfr2vr_s(alpha); - const __m128 _beta128 = __lsx_vreplfr2vr_s(beta); +#endif // __loongarch_asx + __m128 _alpha128 = __lsx_vreplfr2vr_s(alpha); + __m128 _beta128 = __lsx_vreplfr2vr_s(beta); for (; jj + 7 < max_jj; jj += 8) { - __m128 _f0 = (__m128)__lsx_vld(pp + ii * max_jj + jj, 0); - __m128 _f1 = (__m128)__lsx_vld(pp + ii * max_jj + jj + 4, 0); + __m128 _f0 = (__m128)__lsx_vld(pp, 0); + __m128 _f1 = (__m128)__lsx_vld(pp + 4, 0); if (pC) { if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) { - const __m128 _c = __lsx_vreplfr2vr_s(c0); + __m128 _c = __lsx_vreplfr2vr_s(c0); _f0 = __lsx_vfadd_s(_f0, _c); _f1 = __lsx_vfadd_s(_f1, _c); } if (broadcast_type_C == 3 || broadcast_type_C == 4) { - const __m128 _c0 = (__m128)__lsx_vld(pC, 0); - const __m128 _c1 = (__m128)__lsx_vld(pC + 4, 0); + __m128 _c0 = (__m128)__lsx_vld(pC, 0); + __m128 _c1 = (__m128)__lsx_vld(pC + 4, 0); if (beta == 1.f) { _f0 = __lsx_vfadd_s(_f0, _c0); @@ -4540,6 +5295,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _f0 = __lsx_vfmadd_s(_c0, _beta128, _f0); _f1 = __lsx_vfmadd_s(_c1, _beta128, _f1); } + pC += 8; } } if (alpha != 1.f) @@ -4550,12 +5306,11 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b __lsx_vst((__m128i)_f0, p0, 0); __lsx_vst((__m128i)_f1, p0 + 4, 0); p0 += 8; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 8; + pp += 8; } for (; jj + 3 < max_jj; jj += 4) { - __m128 _f0 = (__m128)__lsx_vld(pp + ii * max_jj + jj, 0); + __m128 _f0 = (__m128)__lsx_vld(pp, 0); if (pC) { if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) @@ -4567,17 +5322,17 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _f0 = __lsx_vfadd_s(_f0, _c0); else _f0 = __lsx_vfmadd_s(_c0, _beta128, _f0); + pC += 4; } } if (alpha != 1.f) _f0 = __lsx_vfmul_s(_f0, _alpha128); __lsx_vst((__m128i)_f0, p0, 0); p0 += 4; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 4; + pp += 4; } for (; jj + 1 < max_jj; jj += 2) { - __m128 _f0 = (__m128)__lsx_vldrepl_d(pp + ii * max_jj + jj, 0); + __m128 _f0 = (__m128)__lsx_vldrepl_d(pp, 0); if (pC) { if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) @@ -4589,31 +5344,32 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _f0 = __lsx_vfadd_s(_f0, _c0); else _f0 = __lsx_vfmadd_s(_c0, _beta128, _f0); + pC += 2; } } if (alpha != 1.f) _f0 = __lsx_vfmul_s(_f0, _alpha128); __lsx_vstelm_d((__m128i)_f0, p0, 0, 0); p0 += 2; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 2; + pp += 2; } -#endif +#endif // __loongarch_sx for (; jj < max_jj; jj++) { - float f0 = pp[ii * max_jj + jj]; + float f0 = *pp++; if (pC) { if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) f0 += c0; if (broadcast_type_C == 3 || broadcast_type_C == 4) + { f0 += beta == 1.f ? pC[0] : pC[0] * beta; + pC++; + } } if (alpha != 1.f) f0 *= alpha; p0[0] = f0; p0++; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC++; } } } @@ -4644,7 +5400,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _c1 = (__m128)__lsx_vld(pC + i + ii + 4, 0); if (beta != 1.f) { - const __m128 _beta = __lsx_vreplfr2vr_s(beta); + __m128 _beta = __lsx_vreplfr2vr_s(beta); _c0 = __lsx_vfmul_s(_c0, _beta); _c1 = __lsx_vfmul_s(_c1, _beta); } @@ -4655,13 +5411,13 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, if (pC && broadcast_type_C == 4) pC += j; - const __m128 _alpha = __lsx_vreplfr2vr_s(alpha); - const __m128 _beta = __lsx_vreplfr2vr_s(beta); + __m128 _alpha = __lsx_vreplfr2vr_s(alpha); + __m128 _beta = __lsx_vreplfr2vr_s(beta); int jj = 0; #if __loongarch_asx - const __m256 _alpha256 = (__m256)__lasx_xvreplfr2vr_s(alpha); - const __m256 _beta256 = (__m256)__lasx_xvreplfr2vr_s(beta); - const __m256 _c256 = __lasx_concat_128_s(_c0, _c1); + __m256 _alpha256 = (__m256)__lasx_xvreplfr2vr_s(alpha); + __m256 _beta256 = (__m256)__lasx_xvreplfr2vr_s(beta); + __m256 _c256 = __lasx_concat_128_s(_c0, _c1); for (; jj + 7 < max_jj; jj += 8) { __m256i _sum0 = __lasx_xvld(pp, 0); @@ -4756,6 +5512,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _f6 = __lasx_xvfmadd_s(_cc6, _beta256, _f6); _f7 = __lasx_xvfmadd_s(_cc7, _beta256, _f7); } + pC += 8; } if (broadcast_type_C == 4) { @@ -4767,6 +5524,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _f5 = __lasx_xvfadd_s(_f5, (__m256)__lasx_xvreplfr2vr_s(pC[5] * beta)); _f6 = __lasx_xvfadd_s(_f6, (__m256)__lasx_xvreplfr2vr_s(pC[6] * beta)); _f7 = __lasx_xvfadd_s(_f7, (__m256)__lasx_xvreplfr2vr_s(pC[7] * beta)); + pC += 8; } } if (alpha != 1.f) @@ -4789,10 +5547,8 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, __lasx_xvst(_f6, p0 + M * 6, 0); __lasx_xvst(_f7, p0 + M * 7, 0); p0 += M * 8; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 8; } -#endif +#endif // __loongarch_asx for (; jj + 3 < max_jj; jj += 4) { __m128i _sum0 = __lsx_vld(pp, 0); @@ -4873,6 +5629,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _f6 = __lsx_vfmadd_s(_cc6, _beta, _f6); _f7 = __lsx_vfmadd_s(_cc7, _beta, _f7); } + pC += 4; } if (broadcast_type_C == 4) { @@ -4887,6 +5644,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _f5 = __lsx_vfadd_s(_f5, (__m128)__lsx_vreplvei_w((__m128i)_cc, 1)); _f6 = __lsx_vfadd_s(_f6, (__m128)__lsx_vreplvei_w((__m128i)_cc, 2)); _f7 = __lsx_vfadd_s(_f7, (__m128)__lsx_vreplvei_w((__m128i)_cc, 3)); + pC += 4; } } if (alpha != 1.f) @@ -4909,8 +5667,6 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, __lsx_vst((__m128i)_f3, p0 + M * 3, 0); __lsx_vst((__m128i)_f7, p0 + M * 3 + 4, 0); p0 += M * 4; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 4; } for (; jj + 1 < max_jj; jj += 2) { @@ -4960,15 +5716,17 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _f2 = __lsx_vfmadd_s((__m128)_ci2, _beta, _f2); _f3 = __lsx_vfmadd_s((__m128)_ci3, _beta, _f3); } + pC += 2; } if (broadcast_type_C == 4) { - const __m128 _cc0 = __lsx_vreplfr2vr_s(pC[0] * beta); - const __m128 _cc1 = __lsx_vreplfr2vr_s(pC[1] * beta); + __m128 _cc0 = __lsx_vreplfr2vr_s(pC[0] * beta); + __m128 _cc1 = __lsx_vreplfr2vr_s(pC[1] * beta); _f0 = __lsx_vfadd_s(_f0, _cc0); _f1 = __lsx_vfadd_s(_f1, _cc0); _f2 = __lsx_vfadd_s(_f2, _cc1); _f3 = __lsx_vfadd_s(_f3, _cc1); + pC += 2; } } if (alpha != 1.f) @@ -4983,8 +5741,6 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, __lsx_vst((__m128i)_f2, p0 + M, 0); __lsx_vst((__m128i)_f3, p0 + M + 4, 0); p0 += M * 2; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 2; } for (; jj < max_jj; jj++) { @@ -5020,12 +5776,14 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _f0 = __lsx_vfmadd_s((__m128)_ci0, _beta, _f0); _f1 = __lsx_vfmadd_s((__m128)_ci1, _beta, _f1); } + pC++; } if (broadcast_type_C == 4) { - const __m128 _cc = __lsx_vreplfr2vr_s(pC[0] * beta); + __m128 _cc = __lsx_vreplfr2vr_s(pC[0] * beta); _f0 = __lsx_vfadd_s(_f0, _cc); _f1 = __lsx_vfadd_s(_f1, _cc); + pC++; } } if (alpha != 1.f) @@ -5036,16 +5794,10 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, __lsx_vst((__m128i)_f0, p0, 0); __lsx_vst((__m128i)_f1, p0 + 4, 0); p0 += M; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC++; } } for (; ii + 3 < max_ii; ii += 4) { - const float* pp0 = pp + (ii + 0) * max_jj; - const float* pp1 = pp + (ii + 1) * max_jj; - const float* pp2 = pp + (ii + 2) * max_jj; - const float* pp3 = pp + (ii + 3) * max_jj; float* p0 = outptr + (size_t)j * M + i + ii; const float* pC = pC_base; @@ -5064,27 +5816,28 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, if (pC && broadcast_type_C == 4) pC += j; - const __m128 _alpha = __lsx_vreplfr2vr_s(alpha); - const __m128 _beta = __lsx_vreplfr2vr_s(beta); + __m128 _alpha = __lsx_vreplfr2vr_s(alpha); + __m128 _beta = __lsx_vreplfr2vr_s(beta); int jj = 0; #if __loongarch_asx - const __m256 _alpha256 = (__m256)__lasx_xvreplfr2vr_s(alpha); - const __m256 _beta256 = (__m256)__lasx_xvreplfr2vr_s(beta); + __m256 _alpha256 = (__m256)__lasx_xvreplfr2vr_s(alpha); + __m256 _beta256 = (__m256)__lasx_xvreplfr2vr_s(beta); for (; jj + 15 < max_jj; jj += 16) { - __m256 _f00 = (__m256)__lasx_xvld(pp0 + jj, 0); - __m256 _f01 = (__m256)__lasx_xvld(pp0 + jj + 8, 0); - __m256 _f10 = (__m256)__lasx_xvld(pp1 + jj, 0); - __m256 _f11 = (__m256)__lasx_xvld(pp1 + jj + 8, 0); - __m256 _f20 = (__m256)__lasx_xvld(pp2 + jj, 0); - __m256 _f21 = (__m256)__lasx_xvld(pp2 + jj + 8, 0); - __m256 _f30 = (__m256)__lasx_xvld(pp3 + jj, 0); - __m256 _f31 = (__m256)__lasx_xvld(pp3 + jj + 8, 0); + __m256 _f00 = (__m256)__lasx_xvld(pp, 0); + __m256 _f01 = (__m256)__lasx_xvld(pp + 8, 0); + __m256 _f10 = (__m256)__lasx_xvld(pp + 16, 0); + __m256 _f11 = (__m256)__lasx_xvld(pp + 24, 0); + __m256 _f20 = (__m256)__lasx_xvld(pp + 32, 0); + __m256 _f21 = (__m256)__lasx_xvld(pp + 40, 0); + __m256 _f30 = (__m256)__lasx_xvld(pp + 48, 0); + __m256 _f31 = (__m256)__lasx_xvld(pp + 56, 0); + pp += 64; if (pC) { if (broadcast_type_C == 0) { - const __m256 _cc = (__m256)__lasx_xvreplfr2vr_s(pC[0] * beta); + __m256 _cc = (__m256)__lasx_xvreplfr2vr_s(pC[0] * beta); _f00 = __lasx_xvfadd_s(_f00, _cc); _f01 = __lasx_xvfadd_s(_f01, _cc); _f10 = __lasx_xvfadd_s(_f10, _cc); @@ -5096,10 +5849,10 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, } if (broadcast_type_C == 1 || broadcast_type_C == 2) { - const __m256 _c0 = (__m256)__lasx_xvreplfr2vr_s(pC[i + ii] * beta); - const __m256 _c1 = (__m256)__lasx_xvreplfr2vr_s(pC[i + ii + 1] * beta); - const __m256 _c2 = (__m256)__lasx_xvreplfr2vr_s(pC[i + ii + 2] * beta); - const __m256 _c3 = (__m256)__lasx_xvreplfr2vr_s(pC[i + ii + 3] * beta); + __m256 _c0 = (__m256)__lasx_xvreplfr2vr_s(pC[i + ii] * beta); + __m256 _c1 = (__m256)__lasx_xvreplfr2vr_s(pC[i + ii + 1] * beta); + __m256 _c2 = (__m256)__lasx_xvreplfr2vr_s(pC[i + ii + 2] * beta); + __m256 _c3 = (__m256)__lasx_xvreplfr2vr_s(pC[i + ii + 3] * beta); _f00 = __lasx_xvfadd_s(_f00, _c0); _f01 = __lasx_xvfadd_s(_f01, _c0); _f10 = __lasx_xvfadd_s(_f10, _c1); @@ -5111,14 +5864,14 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, } if (broadcast_type_C == 3) { - const __m256 _c00 = (__m256)__lasx_xvld(pC, 0); - const __m256 _c01 = (__m256)__lasx_xvld(pC + 8, 0); - const __m256 _c10 = (__m256)__lasx_xvld(pC + c_hstep, 0); - const __m256 _c11 = (__m256)__lasx_xvld(pC + c_hstep + 8, 0); - const __m256 _c20 = (__m256)__lasx_xvld(pC + c_hstep * 2, 0); - const __m256 _c21 = (__m256)__lasx_xvld(pC + c_hstep * 2 + 8, 0); - const __m256 _c30 = (__m256)__lasx_xvld(pC + c_hstep * 3, 0); - const __m256 _c31 = (__m256)__lasx_xvld(pC + c_hstep * 3 + 8, 0); + __m256 _c00 = (__m256)__lasx_xvld(pC, 0); + __m256 _c01 = (__m256)__lasx_xvld(pC + 8, 0); + __m256 _c10 = (__m256)__lasx_xvld(pC + c_hstep, 0); + __m256 _c11 = (__m256)__lasx_xvld(pC + c_hstep + 8, 0); + __m256 _c20 = (__m256)__lasx_xvld(pC + c_hstep * 2, 0); + __m256 _c21 = (__m256)__lasx_xvld(pC + c_hstep * 2 + 8, 0); + __m256 _c30 = (__m256)__lasx_xvld(pC + c_hstep * 3, 0); + __m256 _c31 = (__m256)__lasx_xvld(pC + c_hstep * 3 + 8, 0); if (beta == 1.f) { _f00 = __lasx_xvfadd_s(_f00, _c00); @@ -5141,6 +5894,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _f30 = __lasx_xvfmadd_s(_c30, _beta256, _f30); _f31 = __lasx_xvfmadd_s(_c31, _beta256, _f31); } + pC += 16; } if (broadcast_type_C == 4) { @@ -5159,6 +5913,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _f21 = __lasx_xvfadd_s(_f21, _c1); _f30 = __lasx_xvfadd_s(_f30, _c0); _f31 = __lasx_xvfadd_s(_f31, _c1); + pC += 16; } } if (alpha != 1.f) @@ -5207,15 +5962,14 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, __lsx_vst(__lasx_extract_128_hi(_r2), p0 + M * 14, 0); __lsx_vst(__lasx_extract_128_hi(_r3), p0 + M * 15, 0); p0 += M * 16; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 16; } for (; jj + 7 < max_jj; jj += 8) { - __m256 _f0 = (__m256)__lasx_xvld(pp0 + jj, 0); - __m256 _f1 = (__m256)__lasx_xvld(pp1 + jj, 0); - __m256 _f2 = (__m256)__lasx_xvld(pp2 + jj, 0); - __m256 _f3 = (__m256)__lasx_xvld(pp3 + jj, 0); + __m256 _f0 = (__m256)__lasx_xvld(pp, 0); + __m256 _f1 = (__m256)__lasx_xvld(pp + 8, 0); + __m256 _f2 = (__m256)__lasx_xvld(pp + 16, 0); + __m256 _f3 = (__m256)__lasx_xvld(pp + 24, 0); + pp += 32; if (pC) { if (broadcast_type_C == 0) @@ -5249,6 +6003,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _f2 = __lasx_xvfmadd_s((__m256)__lasx_xvld(pC + c_hstep * 2, 0), _beta256, _f2); _f3 = __lasx_xvfmadd_s((__m256)__lasx_xvld(pC + c_hstep * 3, 0), _beta256, _f3); } + pC += 8; } if (broadcast_type_C == 4) { @@ -5259,6 +6014,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _f1 = __lasx_xvfadd_s(_f1, _c4); _f2 = __lasx_xvfadd_s(_f2, _c4); _f3 = __lasx_xvfadd_s(_f3, _c4); + pC += 8; } } if (alpha != 1.f) @@ -5287,20 +6043,19 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, __lsx_vst(__lasx_extract_128_hi(_r2), p0 + M * 6, 0); __lsx_vst(__lasx_extract_128_hi(_r3), p0 + M * 7, 0); p0 += M * 8; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 8; } -#endif +#endif // __loongarch_asx for (; jj + 7 < max_jj; jj += 8) { - __m128 _f00 = (__m128)__lsx_vld(pp0 + jj, 0); - __m128 _f01 = (__m128)__lsx_vld(pp0 + jj + 4, 0); - __m128 _f10 = (__m128)__lsx_vld(pp1 + jj, 0); - __m128 _f11 = (__m128)__lsx_vld(pp1 + jj + 4, 0); - __m128 _f20 = (__m128)__lsx_vld(pp2 + jj, 0); - __m128 _f21 = (__m128)__lsx_vld(pp2 + jj + 4, 0); - __m128 _f30 = (__m128)__lsx_vld(pp3 + jj, 0); - __m128 _f31 = (__m128)__lsx_vld(pp3 + jj + 4, 0); + __m128 _f00 = (__m128)__lsx_vld(pp, 0); + __m128 _f01 = (__m128)__lsx_vld(pp + 4, 0); + __m128 _f10 = (__m128)__lsx_vld(pp + 8, 0); + __m128 _f11 = (__m128)__lsx_vld(pp + 12, 0); + __m128 _f20 = (__m128)__lsx_vld(pp + 16, 0); + __m128 _f21 = (__m128)__lsx_vld(pp + 20, 0); + __m128 _f30 = (__m128)__lsx_vld(pp + 24, 0); + __m128 _f31 = (__m128)__lsx_vld(pp + 28, 0); + pp += 32; transpose4x4_ps(_f00, _f10, _f20, _f30); transpose4x4_ps(_f01, _f11, _f21, _f31); if (pC) @@ -5350,6 +6105,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _f21 = __lsx_vfmadd_s(_c21, _beta, _f21); _f31 = __lsx_vfmadd_s(_c31, _beta, _f31); } + pC += 8; } if (broadcast_type_C == 4) { @@ -5368,6 +6124,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _f11 = __lsx_vfadd_s(_f11, (__m128)__lsx_vreplvei_w((__m128i)_cc1, 1)); _f21 = __lsx_vfadd_s(_f21, (__m128)__lsx_vreplvei_w((__m128i)_cc1, 2)); _f31 = __lsx_vfadd_s(_f31, (__m128)__lsx_vreplvei_w((__m128i)_cc1, 3)); + pC += 8; } } if (alpha != 1.f) @@ -5390,15 +6147,14 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, __lsx_vst((__m128i)_f21, p0 + M * 6, 0); __lsx_vst((__m128i)_f31, p0 + M * 7, 0); p0 += M * 8; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 8; } for (; jj + 3 < max_jj; jj += 4) { - __m128 _f0 = (__m128)__lsx_vld(pp0 + jj, 0); - __m128 _f1 = (__m128)__lsx_vld(pp1 + jj, 0); - __m128 _f2 = (__m128)__lsx_vld(pp2 + jj, 0); - __m128 _f3 = (__m128)__lsx_vld(pp3 + jj, 0); + __m128 _f0 = (__m128)__lsx_vld(pp, 0); + __m128 _f1 = (__m128)__lsx_vld(pp + 4, 0); + __m128 _f2 = (__m128)__lsx_vld(pp + 8, 0); + __m128 _f3 = (__m128)__lsx_vld(pp + 12, 0); + pp += 16; transpose4x4_ps(_f0, _f1, _f2, _f3); if (pC) { @@ -5430,6 +6186,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _f2 = __lsx_vfmadd_s(_c2, _beta, _f2); _f3 = __lsx_vfmadd_s(_c3, _beta, _f3); } + pC += 4; } if (broadcast_type_C == 4) { @@ -5440,6 +6197,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _f1 = __lsx_vfadd_s(_f1, (__m128)__lsx_vreplvei_w((__m128i)_c4, 1)); _f2 = __lsx_vfadd_s(_f2, (__m128)__lsx_vreplvei_w((__m128i)_c4, 2)); _f3 = __lsx_vfadd_s(_f3, (__m128)__lsx_vreplvei_w((__m128i)_c4, 3)); + pC += 4; } } if (alpha != 1.f) @@ -5454,15 +6212,14 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, __lsx_vst((__m128i)_f2, p0 + M * 2, 0); __lsx_vst((__m128i)_f3, p0 + M * 3, 0); p0 += M * 4; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 4; } for (; jj + 1 < max_jj; jj += 2) { - __m128i _r0 = __lsx_vldrepl_d(pp0 + jj, 0); - __m128i _r1 = __lsx_vldrepl_d(pp1 + jj, 0); - __m128i _r2 = __lsx_vldrepl_d(pp2 + jj, 0); - __m128i _r3 = __lsx_vldrepl_d(pp3 + jj, 0); + __m128i _r0 = __lsx_vldrepl_d(pp, 0); + __m128i _r1 = __lsx_vldrepl_d(pp + 2, 0); + __m128i _r2 = __lsx_vldrepl_d(pp + 4, 0); + __m128i _r3 = __lsx_vldrepl_d(pp + 6, 0); + pp += 8; __m128i _t0 = __lsx_vilvl_w(_r1, _r0); __m128i _t1 = __lsx_vilvl_w(_r3, _r2); __m128 _f0 = (__m128)__lsx_vilvl_d(_t1, _t0); @@ -5482,8 +6239,8 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _r3 = __lsx_vldrepl_d(pC + c_hstep * 3, 0); _t0 = __lsx_vilvl_w(_r1, _r0); _t1 = __lsx_vilvl_w(_r3, _r2); - const __m128 _cc0 = (__m128)__lsx_vilvl_d(_t1, _t0); - const __m128 _cc1 = (__m128)__lsx_vilvh_d(_t1, _t0); + __m128 _cc0 = (__m128)__lsx_vilvl_d(_t1, _t0); + __m128 _cc1 = (__m128)__lsx_vilvh_d(_t1, _t0); if (beta == 1.f) { _f0 = __lsx_vfadd_s(_f0, _cc0); @@ -5494,11 +6251,13 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _f0 = __lsx_vfmadd_s(_cc0, _beta, _f0); _f1 = __lsx_vfmadd_s(_cc1, _beta, _f1); } + pC += 2; } if (broadcast_type_C == 4) { _f0 = __lsx_vfadd_s(_f0, __lsx_vreplfr2vr_s(pC[0] * beta)); _f1 = __lsx_vfadd_s(_f1, __lsx_vreplfr2vr_s(pC[1] * beta)); + pC += 2; } } if (alpha != 1.f) @@ -5509,15 +6268,11 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, __lsx_vst((__m128i)_f0, p0, 0); __lsx_vst((__m128i)_f1, p0 + M, 0); p0 += M * 2; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 2; } for (; jj < max_jj; jj++) { - __m128i _fi = __lsx_vldrepl_w(pp0 + jj, 0); - _fi = __lsx_vinsgr2vr_w(_fi, ((const int*)(pp1 + jj))[0], 1); - _fi = __lsx_vinsgr2vr_w(_fi, ((const int*)(pp2 + jj))[0], 2); - _fi = __lsx_vinsgr2vr_w(_fi, ((const int*)(pp3 + jj))[0], 3); + __m128i _fi = __lsx_vld(pp, 0); + pp += 4; __m128 _f = (__m128)_fi; if (pC) { @@ -5533,6 +6288,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _f = __lsx_vfadd_s(_f, (__m128)_ci); else _f = __lsx_vfmadd_s((__m128)_ci, _beta, _f); + pC++; } if (broadcast_type_C == 4) { @@ -5541,21 +6297,18 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _f = __lsx_vfadd_s(_f, _c4); else _f = __lsx_vfmadd_s(_c4, _beta, _f); + pC++; } } if (alpha != 1.f) _f = __lsx_vfmul_s(_f, _alpha); __lsx_vst((__m128i)_f, p0, 0); p0 += M; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC++; } } -#endif +#endif // __loongarch_sx for (; ii + 1 < max_ii; ii += 2) { - const float* pp0 = pp + (ii + 0) * max_jj; - const float* pp1 = pp + (ii + 1) * max_jj; float* p0 = outptr + (size_t)j * M + i + ii; const float* pC = pC_base; @@ -5592,22 +6345,22 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, if (beta != 1.f) _c = __lsx_vfmul_s(_c, __lsx_vreplfr2vr_s(beta)); } - const __m128 _alpha = __lsx_vreplfr2vr_s(alpha); - const __m128 _beta = __lsx_vreplfr2vr_s(beta); + __m128 _alpha = __lsx_vreplfr2vr_s(alpha); + __m128 _beta = __lsx_vreplfr2vr_s(beta); #if __loongarch_asx - const __m256 _alpha256 = (__m256)__lasx_xvreplfr2vr_s(alpha); - const __m256 _beta256 = (__m256)__lasx_xvreplfr2vr_s(beta); + __m256 _alpha256 = (__m256)__lasx_xvreplfr2vr_s(alpha); + __m256 _beta256 = (__m256)__lasx_xvreplfr2vr_s(beta); for (; jj + 15 < max_jj; jj += 16) { - __m256 _f00 = (__m256)__lasx_xvld(pp0 + jj, 0); - __m256 _f01 = (__m256)__lasx_xvld(pp0 + jj + 8, 0); - __m256 _f10 = (__m256)__lasx_xvld(pp1 + jj, 0); - __m256 _f11 = (__m256)__lasx_xvld(pp1 + jj + 8, 0); + __m256 _f00 = (__m256)__lasx_xvld(pp, 0); + __m256 _f01 = (__m256)__lasx_xvld(pp + 8, 0); + __m256 _f10 = (__m256)__lasx_xvld(pp + 16, 0); + __m256 _f11 = (__m256)__lasx_xvld(pp + 24, 0); if (pC) { if (broadcast_type_C == 0) { - const __m256 _cc = (__m256)__lasx_xvreplfr2vr_s(c0); + __m256 _cc = (__m256)__lasx_xvreplfr2vr_s(c0); _f00 = __lasx_xvfadd_s(_f00, _cc); _f01 = __lasx_xvfadd_s(_f01, _cc); _f10 = __lasx_xvfadd_s(_f10, _cc); @@ -5615,8 +6368,8 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, } if (broadcast_type_C == 1 || broadcast_type_C == 2) { - const __m256 _c0 = (__m256)__lasx_xvreplfr2vr_s(c0); - const __m256 _c1 = (__m256)__lasx_xvreplfr2vr_s(c1); + __m256 _c0 = (__m256)__lasx_xvreplfr2vr_s(c0); + __m256 _c1 = (__m256)__lasx_xvreplfr2vr_s(c1); _f00 = __lasx_xvfadd_s(_f00, _c0); _f01 = __lasx_xvfadd_s(_f01, _c0); _f10 = __lasx_xvfadd_s(_f10, _c1); @@ -5624,10 +6377,10 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, } if (broadcast_type_C == 3) { - const __m256 _c00 = (__m256)__lasx_xvld(pC, 0); - const __m256 _c01 = (__m256)__lasx_xvld(pC + 8, 0); - const __m256 _c10 = (__m256)__lasx_xvld(pC + c_hstep, 0); - const __m256 _c11 = (__m256)__lasx_xvld(pC + c_hstep + 8, 0); + __m256 _c00 = (__m256)__lasx_xvld(pC, 0); + __m256 _c01 = (__m256)__lasx_xvld(pC + 8, 0); + __m256 _c10 = (__m256)__lasx_xvld(pC + c_hstep, 0); + __m256 _c11 = (__m256)__lasx_xvld(pC + c_hstep + 8, 0); if (beta == 1.f) { _f00 = __lasx_xvfadd_s(_f00, _c00); @@ -5642,6 +6395,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _f10 = __lasx_xvfmadd_s(_c10, _beta256, _f10); _f11 = __lasx_xvfmadd_s(_c11, _beta256, _f11); } + pC += 16; } if (broadcast_type_C == 4) { @@ -5656,6 +6410,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _f01 = __lasx_xvfadd_s(_f01, _c1); _f10 = __lasx_xvfadd_s(_f10, _c0); _f11 = __lasx_xvfadd_s(_f11, _c1); + pC += 16; } } if (alpha != 1.f) @@ -5686,13 +6441,12 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, __lasx_xvstelm_d(_tmp1, p0 + M * 14, 0, 2); __lasx_xvstelm_d(_tmp1, p0 + M * 15, 0, 3); p0 += M * 16; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 16; + pp += 32; } for (; jj + 7 < max_jj; jj += 8) { - __m256 _f0 = (__m256)__lasx_xvld(pp0 + jj, 0); - __m256 _f1 = (__m256)__lasx_xvld(pp1 + jj, 0); + __m256 _f0 = (__m256)__lasx_xvld(pp, 0); + __m256 _f1 = (__m256)__lasx_xvld(pp + 8, 0); if (pC) { if (broadcast_type_C == 0) @@ -5718,6 +6472,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _f0 = __lasx_xvfmadd_s((__m256)__lasx_xvld(pC, 0), _beta256, _f0); _f1 = __lasx_xvfmadd_s((__m256)__lasx_xvld(pC + c_hstep, 0), _beta256, _f1); } + pC += 8; } if (broadcast_type_C == 4) { @@ -5726,6 +6481,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _c4 = __lasx_xvfmul_s(_c4, _beta256); _f0 = __lasx_xvfadd_s(_f0, _c4); _f1 = __lasx_xvfadd_s(_f1, _c4); + pC += 8; } } if (alpha != 1.f) @@ -5745,21 +6501,20 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, __lasx_xvstelm_d(_tmp1, p0 + M * 6, 0, 2); __lasx_xvstelm_d(_tmp1, p0 + M * 7, 0, 3); p0 += M * 8; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 8; + pp += 16; } -#endif +#endif // __loongarch_asx for (; jj + 7 < max_jj; jj += 8) { - __m128 _f00 = (__m128)__lsx_vld(pp0 + jj, 0); - __m128 _f01 = (__m128)__lsx_vld(pp0 + jj + 4, 0); - __m128 _f10 = (__m128)__lsx_vld(pp1 + jj, 0); - __m128 _f11 = (__m128)__lsx_vld(pp1 + jj + 4, 0); + __m128 _f00 = (__m128)__lsx_vld(pp, 0); + __m128 _f01 = (__m128)__lsx_vld(pp + 4, 0); + __m128 _f10 = (__m128)__lsx_vld(pp + 8, 0); + __m128 _f11 = (__m128)__lsx_vld(pp + 12, 0); if (pC) { if (broadcast_type_C == 0) { - const __m128 _cc = __lsx_vreplfr2vr_s(c0); + __m128 _cc = __lsx_vreplfr2vr_s(c0); _f00 = __lsx_vfadd_s(_f00, _cc); _f01 = __lsx_vfadd_s(_f01, _cc); _f10 = __lsx_vfadd_s(_f10, _cc); @@ -5767,8 +6522,8 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, } if (broadcast_type_C == 1 || broadcast_type_C == 2) { - const __m128 _c0 = __lsx_vreplfr2vr_s(c0); - const __m128 _c1 = __lsx_vreplfr2vr_s(c1); + __m128 _c0 = __lsx_vreplfr2vr_s(c0); + __m128 _c1 = __lsx_vreplfr2vr_s(c1); _f00 = __lsx_vfadd_s(_f00, _c0); _f01 = __lsx_vfadd_s(_f01, _c0); _f10 = __lsx_vfadd_s(_f10, _c1); @@ -5776,10 +6531,10 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, } if (broadcast_type_C == 3) { - const __m128 _c00 = (__m128)__lsx_vld(pC, 0); - const __m128 _c01 = (__m128)__lsx_vld(pC + 4, 0); - const __m128 _c10 = (__m128)__lsx_vld(pC + c_hstep, 0); - const __m128 _c11 = (__m128)__lsx_vld(pC + c_hstep + 4, 0); + __m128 _c00 = (__m128)__lsx_vld(pC, 0); + __m128 _c01 = (__m128)__lsx_vld(pC + 4, 0); + __m128 _c10 = (__m128)__lsx_vld(pC + c_hstep, 0); + __m128 _c11 = (__m128)__lsx_vld(pC + c_hstep + 4, 0); if (beta == 1.f) { _f00 = __lsx_vfadd_s(_f00, _c00); @@ -5794,6 +6549,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _f10 = __lsx_vfmadd_s(_c10, _beta, _f10); _f11 = __lsx_vfmadd_s(_c11, _beta, _f11); } + pC += 8; } if (broadcast_type_C == 4) { @@ -5808,6 +6564,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _f01 = __lsx_vfadd_s(_f01, _c1); _f10 = __lsx_vfadd_s(_f10, _c0); _f11 = __lsx_vfadd_s(_f11, _c1); + pC += 8; } } if (alpha != 1.f) @@ -5830,13 +6587,12 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, __lsx_vstelm_d(_tmp1, p0 + M * 6, 0, 0); __lsx_vstelm_d(_tmp1, p0 + M * 7, 0, 1); p0 += M * 8; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 8; + pp += 16; } for (; jj + 3 < max_jj; jj += 4) { - __m128 _f0 = (__m128)__lsx_vld(pp0 + jj, 0); - __m128 _f1 = (__m128)__lsx_vld(pp1 + jj, 0); + __m128 _f0 = (__m128)__lsx_vld(pp, 0); + __m128 _f1 = (__m128)__lsx_vld(pp + 4, 0); if (pC) { if (broadcast_type_C == 0) @@ -5864,6 +6620,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _f0 = __lsx_vfmadd_s(_c0, _beta, _f0); _f1 = __lsx_vfmadd_s(_c1, _beta, _f1); } + pC += 4; } if (broadcast_type_C == 4) { @@ -5872,6 +6629,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _c4 = __lsx_vfmul_s(_c4, _beta); _f0 = __lsx_vfadd_s(_f0, _c4); _f1 = __lsx_vfadd_s(_f1, _c4); + pC += 4; } } if (alpha != 1.f) @@ -5886,14 +6644,13 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, __lsx_vstelm_d(_tmp1, p0 + M * 2, 0, 0); __lsx_vstelm_d(_tmp1, p0 + M * 3, 0, 1); p0 += M * 4; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 4; + pp += 8; } for (; jj + 1 < max_jj; jj += 2) { - __m128i _r0 = __lsx_vldrepl_d(pp0 + jj, 0); - __m128i _r1 = __lsx_vldrepl_d(pp1 + jj, 0); - __m128 _f = (__m128)__lsx_vilvl_w(_r1, _r0); + __m128 _f = (__m128)__lsx_vshuf4i_w(__lsx_vld(pp, 0), _LSX_SHUFFLE(3, 1, 2, 0)); + __m128i _r0; + __m128i _r1; if (pC) { if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) @@ -5902,11 +6659,12 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, { _r0 = __lsx_vldrepl_d(pC, 0); _r1 = __lsx_vldrepl_d(pC + c_hstep, 0); - const __m128 _cc = (__m128)__lsx_vilvl_w(_r1, _r0); + __m128 _cc = (__m128)__lsx_vilvl_w(_r1, _r0); if (beta == 1.f) _f = __lsx_vfadd_s(_f, _cc); else _f = __lsx_vfmadd_s(_cc, _beta, _f); + pC += 2; } if (broadcast_type_C == 4) { @@ -5916,6 +6674,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _f = __lsx_vfadd_s(_f, (__m128)_cc); else _f = __lsx_vfmadd_s((__m128)_cc, _beta, _f); + pC += 2; } } if (alpha != 1.f) @@ -5923,13 +6682,11 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, __lsx_vstelm_d((__m128i)_f, p0, 0, 0); __lsx_vstelm_d((__m128i)_f, p0 + M, 0, 1); p0 += M * 2; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 2; + pp += 4; } for (; jj < max_jj; jj++) { - __m128i _fi = __lsx_vldrepl_w(pp0 + jj, 0); - _fi = __lsx_vinsgr2vr_w(_fi, ((const int*)(pp1 + jj))[0], 1); + __m128i _fi = __lsx_vldrepl_d(pp, 0); __m128 _f = (__m128)_fi; if (pC) { @@ -5943,6 +6700,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _f = __lsx_vfadd_s(_f, (__m128)_ci); else _f = __lsx_vfmadd_s((__m128)_ci, _beta, _f); + pC++; } if (broadcast_type_C == 4) { @@ -5951,20 +6709,85 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _f = __lsx_vfadd_s(_f, _c4); else _f = __lsx_vfmadd_s(_c4, _beta, _f); + pC++; } } if (alpha != 1.f) _f = __lsx_vfmul_s(_f, _alpha); __lsx_vstelm_d((__m128i)_f, p0, 0, 0); p0 += M; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC++; + pp += 2; + } +#endif // __loongarch_sx + for (; jj + 1 < max_jj; jj += 2) + { + float f00 = pp[0]; + float f01 = pp[1]; + float f10 = pp[2]; + float f11 = pp[3]; + if (pC) + { + if (broadcast_type_C == 0) + { + f00 += c0; + f01 += c0; + f10 += c0; + f11 += c0; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + f00 += c0; + f01 += c0; + f10 += c1; + f11 += c1; + } + if (broadcast_type_C == 3) + { + if (beta == 1.f) + { + f00 += pC[0]; + f01 += pC[1]; + f10 += pC[c_hstep]; + f11 += pC[c_hstep + 1]; + } + else + { + f00 += pC[0] * beta; + f01 += pC[1] * beta; + f10 += pC[c_hstep] * beta; + f11 += pC[c_hstep + 1] * beta; + } + pC += 2; + } + if (broadcast_type_C == 4) + { + const float cc0 = beta == 1.f ? pC[0] : pC[0] * beta; + const float cc1 = beta == 1.f ? pC[1] : pC[1] * beta; + f00 += cc0; + f01 += cc1; + f10 += cc0; + f11 += cc1; + pC += 2; + } + } + if (alpha != 1.f) + { + f00 *= alpha; + f01 *= alpha; + f10 *= alpha; + f11 *= alpha; + } + p0[0] = f00; + p0[1] = f10; + p0[M] = f01; + p0[M + 1] = f11; + p0 += M * 2; + pp += 4; } -#endif for (; jj < max_jj; jj++) { - float f0 = pp0[jj]; - float f1 = pp1[jj]; + float f0 = pp[0]; + float f1 = pp[1]; if (pC) { if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) @@ -5984,12 +6807,14 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, f0 += pC[0] * beta; f1 += pC[c_hstep] * beta; } + pC++; } if (broadcast_type_C == 4) { float c = beta == 1.f ? pC[0] : pC[0] * beta; f0 += c; f1 += c; + pC++; } } if (alpha != 1.f) @@ -6000,13 +6825,11 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, p0[0] = f0; p0[1] = f1; p0 += M; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC++; + pp += 2; } } for (; ii < max_ii; ii++) { - const float* pp0 = pp + ii * max_jj; float* p0 = outptr + (size_t)j * M + i + ii; const float* pC = pC_base; @@ -6024,14 +6847,15 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, } int jj = 0; +#if __loongarch_sx #if __loongarch_asx - const __m256 _alpha256 = (__m256)__lasx_xvreplfr2vr_s(alpha); - const __m256 _beta256 = (__m256)__lasx_xvreplfr2vr_s(beta); - const __m256 _c256 = (__m256)__lasx_xvreplfr2vr_s(c0); + __m256 _alpha256 = (__m256)__lasx_xvreplfr2vr_s(alpha); + __m256 _beta256 = (__m256)__lasx_xvreplfr2vr_s(beta); + __m256 _c256 = (__m256)__lasx_xvreplfr2vr_s(c0); for (; jj + 15 < max_jj; jj += 16) { - __m256 _f0 = (__m256)__lasx_xvld(pp0 + jj, 0); - __m256 _f1 = (__m256)__lasx_xvld(pp0 + jj + 8, 0); + __m256 _f0 = (__m256)__lasx_xvld(pp, 0); + __m256 _f1 = (__m256)__lasx_xvld(pp + 8, 0); if (pC) { if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) @@ -6041,8 +6865,8 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, } if (broadcast_type_C == 3 || broadcast_type_C == 4) { - const __m256 _c0 = (__m256)__lasx_xvld(pC, 0); - const __m256 _c1 = (__m256)__lasx_xvld(pC + 8, 0); + __m256 _c0 = (__m256)__lasx_xvld(pC, 0); + __m256 _c1 = (__m256)__lasx_xvld(pC + 8, 0); if (beta == 1.f) { _f0 = __lasx_xvfadd_s(_f0, _c0); @@ -6053,6 +6877,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _f0 = __lasx_xvfmadd_s(_c0, _beta256, _f0); _f1 = __lasx_xvfmadd_s(_c1, _beta256, _f1); } + pC += 16; } } if (alpha != 1.f) @@ -6085,12 +6910,11 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, __lasx_xvstelm_w((__m256i)_f1, p0 + M * 15, 0, 7); } p0 += M * 16; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 16; + pp += 16; } for (; jj + 7 < max_jj; jj += 8) { - __m256 _f0 = (__m256)__lasx_xvld(pp0 + jj, 0); + __m256 _f0 = (__m256)__lasx_xvld(pp, 0); if (pC) { if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) @@ -6102,6 +6926,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _f0 = __lasx_xvfadd_s(_f0, _c0); else _f0 = __lasx_xvfmadd_s(_c0, _beta256, _f0); + pC += 8; } } if (alpha != 1.f) _f0 = __lasx_xvfmul_s(_f0, _alpha256); @@ -6119,18 +6944,16 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, __lasx_xvstelm_w((__m256i)_f0, p0 + M * 7, 0, 7); } p0 += M * 8; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 8; + pp += 8; } -#endif -#if __loongarch_sx - const __m128 _alpha128 = __lsx_vreplfr2vr_s(alpha); - const __m128 _beta128 = __lsx_vreplfr2vr_s(beta); - const __m128 _c128 = __lsx_vreplfr2vr_s(c0); +#endif // __loongarch_asx + __m128 _alpha128 = __lsx_vreplfr2vr_s(alpha); + __m128 _beta128 = __lsx_vreplfr2vr_s(beta); + __m128 _c128 = __lsx_vreplfr2vr_s(c0); for (; jj + 7 < max_jj; jj += 8) { - __m128 _f0 = (__m128)__lsx_vld(pp0 + jj, 0); - __m128 _f1 = (__m128)__lsx_vld(pp0 + jj + 4, 0); + __m128 _f0 = (__m128)__lsx_vld(pp, 0); + __m128 _f1 = (__m128)__lsx_vld(pp + 4, 0); if (pC) { if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) @@ -6140,8 +6963,8 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, } if (broadcast_type_C == 3 || broadcast_type_C == 4) { - const __m128 _c0 = (__m128)__lsx_vld(pC, 0); - const __m128 _c1 = (__m128)__lsx_vld(pC + 4, 0); + __m128 _c0 = (__m128)__lsx_vld(pC, 0); + __m128 _c1 = (__m128)__lsx_vld(pC + 4, 0); if (beta == 1.f) { _f0 = __lsx_vfadd_s(_f0, _c0); @@ -6152,6 +6975,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _f0 = __lsx_vfmadd_s(_c0, _beta128, _f0); _f1 = __lsx_vfmadd_s(_c1, _beta128, _f1); } + pC += 8; } } if (alpha != 1.f) @@ -6176,12 +7000,11 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, __lsx_vstelm_w((__m128i)_f1, p0 + M * 7, 0, 3); } p0 += M * 8; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 8; + pp += 8; } for (; jj + 3 < max_jj; jj += 4) { - __m128 _f0 = (__m128)__lsx_vld(pp0 + jj, 0); + __m128 _f0 = (__m128)__lsx_vld(pp, 0); if (pC) { if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) @@ -6193,6 +7016,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _f0 = __lsx_vfadd_s(_f0, _c0); else _f0 = __lsx_vfmadd_s(_c0, _beta128, _f0); + pC += 4; } } if (alpha != 1.f) _f0 = __lsx_vfmul_s(_f0, _alpha128); @@ -6206,23 +7030,23 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, __lsx_vstelm_w((__m128i)_f0, p0 + M * 3, 0, 3); } p0 += M * 4; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 4; + pp += 4; } for (; jj + 1 < max_jj; jj += 2) { - __m128 _f0 = (__m128)__lsx_vldrepl_d(pp0 + jj, 0); + __m128 _f0 = (__m128)__lsx_vldrepl_d(pp, 0); if (pC) { if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) _f0 = __lsx_vfadd_s(_f0, _c128); if (broadcast_type_C == 3 || broadcast_type_C == 4) { - const __m128 _cc = (__m128)__lsx_vldrepl_d(pC, 0); + __m128 _cc = (__m128)__lsx_vldrepl_d(pC, 0); if (beta == 1.f) _f0 = __lsx_vfadd_s(_f0, _cc); else _f0 = __lsx_vfmadd_s(_cc, _beta128, _f0); + pC += 2; } } if (alpha != 1.f) @@ -6235,31 +7059,31 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, __lsx_vstelm_w((__m128i)_f0, p0 + M, 0, 1); } p0 += M * 2; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 2; + pp += 2; } -#endif +#endif // __loongarch_sx for (; jj < max_jj; jj++) { - float f0 = pp0[jj]; + float f0 = *pp++; if (pC) { if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) f0 += c0; if (broadcast_type_C == 3 || broadcast_type_C == 4) + { f0 += beta == 1.f ? pC[0] : pC[0] * beta; + pC++; + } } if (alpha != 1.f) f0 *= alpha; p0[0] = f0; p0 += M; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC++; } } } -static void get_optimal_tile_mnk_wq_int8(int M, int N, int K, int constant_TILE_M, int constant_TILE_N, int constant_TILE_K, int& TILE_M, int& TILE_N, int& TILE_K, int nT) +static void get_optimal_tile_mnk_wq_int8(int M, int N, int K, int block_size, int constant_TILE_M, int constant_TILE_N, int constant_TILE_K, int& TILE_M, int& TILE_N, int& TILE_K, int nT) { // resolve optimal tile size from cache size const size_t l2_cache_size = get_cpu_level2_cache_size(); @@ -6267,7 +7091,24 @@ static void get_optimal_tile_mnk_wq_int8(int M, int N, int K, int constant_TILE_ if (nT == 0) nT = get_physical_big_cpu_count(); - const int tile_size = std::max(1, (int)((float)l2_cache_size / 2 / sizeof(signed char) / std::max(1, K))); + int tile_size = (int)sqrtf((float)l2_cache_size / (2 * sizeof(signed char) + sizeof(float))); + + TILE_K = std::max(block_size, tile_size / block_size * block_size); + if (K > 0) + { + if (TILE_K >= K) + { + TILE_K = K; + } + else + { + const int nn_K = (K + TILE_K - 1) / TILE_K; + const int tile_k = (K + nn_K - 1) / nn_K; + TILE_K = std::max(block_size, tile_k / block_size * block_size); + } + } + + tile_size = std::max(1, (int)((float)l2_cache_size / 2 / sizeof(signed char) / std::max(1, TILE_K))); #if __loongarch_sx const int tile_m_align = M >= nT * 8 ? 8 : M >= nT * 4 ? 4 : M >= nT * 2 ? 2 : 1; @@ -6277,12 +7118,11 @@ static void get_optimal_tile_mnk_wq_int8(int M, int N, int K, int constant_TILE_ const int tile_n_align = tile_m_align == 8 ? 4 : 8; #endif #else - const int tile_m_align = M >= nT * 4 ? 4 : M >= nT * 2 ? 2 : 1; + const int tile_m_align = M >= nT * 2 ? 2 : 1; const int tile_n_align = 2; #endif TILE_M = tile_m_align; TILE_N = std::max(tile_n_align, tile_size / tile_n_align * tile_n_align); - TILE_K = K; if (N > 0) { @@ -6294,6 +7134,12 @@ static void get_optimal_tile_mnk_wq_int8(int M, int N, int K, int constant_TILE_ if (constant_TILE_N > 0) TILE_N = (constant_TILE_N + tile_n_align - 1) / tile_n_align * tile_n_align; + if (constant_TILE_K > 0) + { + TILE_K = std::max(block_size, constant_TILE_K / block_size * block_size); + if (K > 0) + TILE_K = std::min(TILE_K, K); + } + (void)constant_TILE_M; - (void)constant_TILE_K; } diff --git a/src/layer/mips/gemm_mips.cpp b/src/layer/mips/gemm_mips.cpp index 84a72f800a5e..6716ec693def 100644 --- a/src/layer/mips/gemm_mips.cpp +++ b/src/layer/mips/gemm_mips.cpp @@ -4480,13 +4480,13 @@ static int gemm_BT_mips_wq_int8(const Mat& A, const Mat& packed_B, const Mat& pa const int M = transA ? A.w : (A.dims == 3 ? A.c : A.h) * A.elempack; const int block_count = (K + block_size - 1) / block_size; int TILE_M, TILE_N, TILE_K; - get_optimal_tile_mnk_wq_int8(M, N, K, constant_TILE_M, constant_TILE_N, constant_TILE_K, TILE_M, TILE_N, TILE_K, nT); + get_optimal_tile_mnk_wq_int8(M, N, K, block_size, constant_TILE_M, constant_TILE_N, constant_TILE_K, TILE_M, TILE_N, TILE_K, nT); - (void)TILE_K; const int mr = std::min(M, TILE_M); const int nr = std::min(N, TILE_N); const int nn_M = (M + TILE_M - 1) / TILE_M; const int nn_N = (N + TILE_N - 1) / TILE_N; + const int nn_K = (K + TILE_K - 1) / TILE_K; const float* input_scale_ptr = input_scales; Mat BT = packed_B.reshape(K, N); Mat BT_descales = packed_B_descales.reshape(block_count, N); @@ -4502,19 +4502,27 @@ static int gemm_BT_mips_wq_int8(const Mat& A, const Mat& packed_B, const Mat& pa if (AT.empty() || AT_descales.empty()) return -100; + const int nn_MK = nn_M * nn_K; #pragma omp parallel for num_threads(nT) - for (int ppi = 0; ppi < nn_M; ppi++) + for (int ppik = 0; ppik < nn_MK; ppik++) { + const int ppi = ppik / nn_K; + const int ppk = ppik % nn_K; + const int i = ppi * TILE_M; + const int k = ppk * TILE_K; const int max_ii = std::min(M - i, TILE_M); + const int max_kk = std::min(K - k, TILE_K); - Mat AT_tile = AT.channel(i / TILE_M).row_range(0, max_ii); - Mat AT_descales_tile = AT_descales.channel(i / TILE_M).row_range(0, max_ii); + Mat AT_channel = AT.channel(i / TILE_M).reshape(K * mr); + Mat AT_descales_channel = AT_descales.channel(i / TILE_M).reshape(block_count * mr); + Mat AT_tile = AT_channel.range(k * mr, max_kk * mr); + Mat AT_descales_tile = AT_descales_channel.range(k / block_size * mr, (max_kk + block_size - 1) / block_size * mr); if (transA) - transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, block_size, input_scale_ptr); + transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, k, max_kk, block_size, input_scale_ptr); else - quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, block_size, input_scale_ptr); + quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, k, max_kk, block_size, input_scale_ptr); } const int nn_MN = nn_M * nn_N; @@ -4529,13 +4537,20 @@ static int gemm_BT_mips_wq_int8(const Mat& A, const Mat& packed_B, const Mat& pa const int max_ii = std::min(M - i, TILE_M); const int max_jj = std::min(N - j, TILE_N); - Mat AT_tile = AT.channel(i / TILE_M).row_range(0, max_ii); - Mat AT_descales_tile = AT_descales.channel(i / TILE_M).row_range(0, max_ii); Mat BT_tile = BT.row_range(j, max_jj); Mat BT_descales_tile = BT_descales.row_range(j, max_jj); Mat topT_tile = topT.channel(get_omp_thread_num()); + Mat AT_channel = AT.channel(i / TILE_M).reshape(K * mr); + Mat AT_descales_channel = AT_descales.channel(i / TILE_M).reshape(block_count * mr); - gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, block_size); + for (int k = 0; k < K; k += TILE_K) + { + const int max_kk = std::min(K - k, TILE_K); + Mat AT_tile = AT_channel.range(k * mr, max_kk * mr); + Mat AT_descales_tile = AT_descales_channel.range(k / block_size * mr, (max_kk + block_size - 1) / block_size * mr); + + gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, k, max_kk, K, block_size); + } if (output_transpose) transpose_unpack_output_tile_wq_int8(topT_tile, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, alpha, beta); @@ -4556,21 +4571,32 @@ static int gemm_BT_mips_wq_int8(const Mat& A, const Mat& packed_B, const Mat& pa const int i = ppi * TILE_M; const int max_ii = std::min(M - i, TILE_M); - Mat AT_tile = ATX.channel(get_omp_thread_num()).row_range(0, max_ii); - Mat AT_descales_tile = ATX_descales.channel(get_omp_thread_num()).row_range(0, max_ii); Mat topT_tile = topT.channel(get_omp_thread_num()); - - if (transA) - transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, block_size, input_scale_ptr); - else - quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, block_size, input_scale_ptr); + Mat ATX_channel = ATX.channel(get_omp_thread_num()).reshape(K * mr); + Mat ATX_descales_channel = ATX_descales.channel(get_omp_thread_num()).reshape(block_count * mr); for (int j = 0; j < N; j += TILE_N) { const int max_jj = std::min(N - j, TILE_N); Mat BT_tile = BT.row_range(j, max_jj); Mat BT_descales_tile = BT_descales.row_range(j, max_jj); - gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, block_size); + + for (int k = 0; k < K; k += TILE_K) + { + const int max_kk = std::min(K - k, TILE_K); + Mat AT_tile = ATX_channel.range(k * mr, max_kk * mr); + Mat AT_descales_tile = ATX_descales_channel.range(k / block_size * mr, (max_kk + block_size - 1) / block_size * mr); + + if (j == 0) + { + if (transA) + transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, k, max_kk, block_size, input_scale_ptr); + else + quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, k, max_kk, block_size, input_scale_ptr); + } + + gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, k, max_kk, K, block_size); + } if (output_transpose) transpose_unpack_output_tile_wq_int8(topT_tile, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, alpha, beta); diff --git a/src/layer/mips/gemm_mips_mmi.cpp b/src/layer/mips/gemm_mips_mmi.cpp index 85fdf18c06f0..06696bf0cfbc 100644 --- a/src/layer/mips/gemm_mips_mmi.cpp +++ b/src/layer/mips/gemm_mips_mmi.cpp @@ -66,19 +66,19 @@ int pack_B_wq_int8_loongson_mmi(const Mat& B, const Mat& B_scales, Mat& packed_B return pack_B_wq_int8(B, B_scales, packed_B, packed_B_descales, N, K, block_size, opt); } -void quantize_A_tile_wq_int8_loongson_mmi(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int block_size, const float* input_scale_ptr) +void quantize_A_tile_wq_int8_loongson_mmi(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int k, int max_kk, int block_size, const float* input_scale_ptr) { - quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, block_size, input_scale_ptr); + quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, k, max_kk, block_size, input_scale_ptr); } -void transpose_quantize_A_tile_wq_int8_loongson_mmi(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int block_size, const float* input_scale_ptr) +void transpose_quantize_A_tile_wq_int8_loongson_mmi(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int k, int max_kk, int block_size, const float* input_scale_ptr) { - transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, block_size, input_scale_ptr); + transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, k, max_kk, block_size, input_scale_ptr); } -void gemm_transB_packed_tile_wq_int8_loongson_mmi(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int block_size) +void gemm_transB_packed_tile_wq_int8_loongson_mmi(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int k, int max_kk, int K, int block_size) { - gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, block_size); + gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, k, max_kk, K, block_size); } #endif // NCNN_WEIGHT_QUANT diff --git a/src/layer/mips/gemm_wq_int8.h b/src/layer/mips/gemm_wq_int8.h index a93737caf9ea..864b72f8c626 100644 --- a/src/layer/mips/gemm_wq_int8.h +++ b/src/layer/mips/gemm_wq_int8.h @@ -3,9 +3,9 @@ #if NCNN_RUNTIME_CPU && NCNN_MMI && !__mips_msa && !__mips_loongson_mmi int pack_B_wq_int8_loongson_mmi(const Mat& B, const Mat& B_scales, Mat& packed_B, Mat& packed_B_descales, int N, int K, int block_size, const Option& opt); -void quantize_A_tile_wq_int8_loongson_mmi(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int block_size, const float* input_scale_ptr); -void transpose_quantize_A_tile_wq_int8_loongson_mmi(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int block_size, const float* input_scale_ptr); -void gemm_transB_packed_tile_wq_int8_loongson_mmi(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int block_size); +void quantize_A_tile_wq_int8_loongson_mmi(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int k, int max_kk, int block_size, const float* input_scale_ptr); +void transpose_quantize_A_tile_wq_int8_loongson_mmi(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int k, int max_kk, int block_size, const float* input_scale_ptr); +void gemm_transB_packed_tile_wq_int8_loongson_mmi(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int k, int max_kk, int K, int block_size); #endif // group-major, output-major within each K4/K2/K1 fragment @@ -40,41 +40,65 @@ static int pack_B_wq_int8(const Mat& B, const Mat& B_scales, Mat& packed_B, Mat& for (int jj = 0; jj < 8; jj += 4) { + const signed char* p0 = B.row(j + jj); + const signed char* p1 = B.row(j + jj + 1); + const signed char* p2 = B.row(j + jj + 2); + const signed char* p3 = B.row(j + jj + 3); + const float* s0 = B_scales.row(j + jj); + const float* s1 = B_scales.row(j + jj + 1); + const float* s2 = B_scales.row(j + jj + 2); + const float* s3 = B_scales.row(j + jj + 3); + for (int g = 0; g < block_count; g++) { const int k0 = g * block_size; const int max_kk = std::min(K - k0, block_size); - int kk = 0; for (; kk + 3 < max_kk; kk += 4) { - const signed char* p0 = B.row(j + jj) + k0 + kk; - const signed char* p1 = B.row(j + jj + 1) + k0 + kk; - const signed char* p2 = B.row(j + jj + 2) + k0 + kk; - const signed char* p3 = B.row(j + jj + 3) + k0 + kk; - const v16i8 _p = (v16i8)__msa_set_w(__msa_load_w(p0), __msa_load_w(p1), __msa_load_w(p2), __msa_load_w(p3)); + v16i8 _p = (v16i8)__msa_set_w(__msa_load_w(p0), __msa_load_w(p1), __msa_load_w(p2), __msa_load_w(p3)); __msa_st_b(_p, pp, 0); pp += 16; + p0 += 4; + p1 += 4; + p2 += 4; + p3 += 4; } if (kk + 1 < max_kk) { - for (int n = 0; n < 4; n++) - { - const signed char* p0 = B.row(j + jj + n) + k0 + kk; - pp[0] = p0[0]; - pp[1] = p0[1]; - pp += 2; - } + pp[0] = p0[0]; + pp[1] = p0[1]; + pp[2] = p1[0]; + pp[3] = p1[1]; + pp[4] = p2[0]; + pp[5] = p2[1]; + pp[6] = p3[0]; + pp[7] = p3[1]; + pp += 8; + p0 += 2; + p1 += 2; + p2 += 2; + p3 += 2; kk += 2; } if (kk < max_kk) { - for (int n = 0; n < 4; n++) - *pp++ = B.row(j + jj + n)[k0 + kk]; - } - - for (int n = 0; n < 4; n++) - *pd++ = 1.f / B_scales.row(j + jj + n)[g]; + pp[0] = p0[0]; + pp[1] = p1[0]; + pp[2] = p2[0]; + pp[3] = p3[0]; + pp += 4; + p0++; + p1++; + p2++; + p3++; + } + + pd[0] = 1.f / *s0++; + pd[1] = 1.f / *s1++; + pd[2] = 1.f / *s2++; + pd[3] = 1.f / *s3++; + pd += 4; } } } @@ -87,42 +111,65 @@ static int pack_B_wq_int8(const Mat& B, const Mat& B_scales, Mat& packed_B, Mat& const int j = nn8 * 8 + ppj * 4; signed char* pp = (signed char*)packed_B + (size_t)j * K; float* pd = (float*)packed_B_descales + (size_t)j * block_count; + const signed char* p0 = B.row(j); + const signed char* p1 = B.row(j + 1); + const signed char* p2 = B.row(j + 2); + const signed char* p3 = B.row(j + 3); + const float* s0 = B_scales.row(j); + const float* s1 = B_scales.row(j + 1); + const float* s2 = B_scales.row(j + 2); + const float* s3 = B_scales.row(j + 3); for (int g = 0; g < block_count; g++) { const int k0 = g * block_size; const int max_kk = std::min(K - k0, block_size); - int kk = 0; for (; kk + 3 < max_kk; kk += 4) { - const signed char* p0 = B.row(j) + k0 + kk; - const signed char* p1 = B.row(j + 1) + k0 + kk; - const signed char* p2 = B.row(j + 2) + k0 + kk; - const signed char* p3 = B.row(j + 3) + k0 + kk; - const v16i8 _p = (v16i8)__msa_set_w(__msa_load_w(p0), __msa_load_w(p1), __msa_load_w(p2), __msa_load_w(p3)); + v16i8 _p = (v16i8)__msa_set_w(__msa_load_w(p0), __msa_load_w(p1), __msa_load_w(p2), __msa_load_w(p3)); __msa_st_b(_p, pp, 0); pp += 16; + p0 += 4; + p1 += 4; + p2 += 4; + p3 += 4; } if (kk + 1 < max_kk) { - for (int n = 0; n < 4; n++) - { - const signed char* p0 = B.row(j + n) + k0 + kk; - pp[0] = p0[0]; - pp[1] = p0[1]; - pp += 2; - } + pp[0] = p0[0]; + pp[1] = p0[1]; + pp[2] = p1[0]; + pp[3] = p1[1]; + pp[4] = p2[0]; + pp[5] = p2[1]; + pp[6] = p3[0]; + pp[7] = p3[1]; + pp += 8; + p0 += 2; + p1 += 2; + p2 += 2; + p3 += 2; kk += 2; } if (kk < max_kk) { - for (int n = 0; n < 4; n++) - *pp++ = B.row(j + n)[k0 + kk]; + pp[0] = p0[0]; + pp[1] = p1[0]; + pp[2] = p2[0]; + pp[3] = p3[0]; + pp += 4; + p0++; + p1++; + p2++; + p3++; } - for (int n = 0; n < 4; n++) - *pd++ = 1.f / B_scales.row(j + n)[g]; + pd[0] = 1.f / *s0++; + pd[1] = 1.f / *s1++; + pd[2] = 1.f / *s2++; + pd[3] = 1.f / *s3++; + pd += 4; } } j += nn4 * 4; @@ -136,17 +183,18 @@ static int pack_B_wq_int8(const Mat& B, const Mat& B_scales, Mat& packed_B, Mat& const int j = j2 + ppj * 2; signed char* pp = (signed char*)packed_B + (size_t)j * K; float* pd = (float*)packed_B_descales + (size_t)j * block_count; + const signed char* p0 = B.row(j); + const signed char* p1 = B.row(j + 1); + const float* s0 = B_scales.row(j); + const float* s1 = B_scales.row(j + 1); for (int g = 0; g < block_count; g++) { const int k0 = g * block_size; const int max_kk = std::min(K - k0, block_size); - int kk = 0; for (; kk + 3 < max_kk; kk += 4) { - const signed char* p0 = B.row(j) + k0 + kk; - const signed char* p1 = B.row(j + 1) + k0 + kk; pp[0] = p0[0]; pp[1] = p0[1]; pp[2] = p0[2]; @@ -156,27 +204,31 @@ static int pack_B_wq_int8(const Mat& B, const Mat& B_scales, Mat& packed_B, Mat& pp[6] = p1[2]; pp[7] = p1[3]; pp += 8; + p0 += 4; + p1 += 4; } if (kk + 1 < max_kk) { - const signed char* p0 = B.row(j) + k0 + kk; - const signed char* p1 = B.row(j + 1) + k0 + kk; pp[0] = p0[0]; pp[1] = p0[1]; pp[2] = p1[0]; pp[3] = p1[1]; pp += 4; + p0 += 2; + p1 += 2; kk += 2; } if (kk < max_kk) { - pp[0] = B.row(j)[k0 + kk]; - pp[1] = B.row(j + 1)[k0 + kk]; + pp[0] = p0[0]; + pp[1] = p1[0]; pp += 2; + p0++; + p1++; } - *pd++ = 1.f / B_scales.row(j)[g]; - *pd++ = 1.f / B_scales.row(j + 1)[g]; + *pd++ = 1.f / *s0++; + *pd++ = 1.f / *s1++; } } j += nn2 * 2; @@ -185,33 +237,37 @@ static int pack_B_wq_int8(const Mat& B, const Mat& B_scales, Mat& packed_B, Mat& { signed char* pp = (signed char*)packed_B + (size_t)j * K; float* pd = (float*)packed_B_descales + (size_t)j * block_count; + const signed char* p0 = B.row(j); + const float* s0 = B_scales.row(j); for (int g = 0; g < block_count; g++) { const int k0 = g * block_size; const int max_kk = std::min(K - k0, block_size); - const signed char* p0 = B.row(j) + k0; - int kk = 0; for (; kk + 3 < max_kk; kk += 4) { - pp[0] = p0[kk]; - pp[1] = p0[kk + 1]; - pp[2] = p0[kk + 2]; - pp[3] = p0[kk + 3]; + pp[0] = p0[0]; + pp[1] = p0[1]; + pp[2] = p0[2]; + pp[3] = p0[3]; pp += 4; + p0 += 4; } if (kk + 1 < max_kk) { - pp[0] = p0[kk]; - pp[1] = p0[kk + 1]; + pp[0] = p0[0]; + pp[1] = p0[1]; pp += 2; + p0 += 2; kk += 2; } if (kk < max_kk) - *pp++ = p0[kk]; + { + *pp++ = *p0++; + } - *pd++ = 1.f / B_scales.row(j)[g]; + *pd++ = 1.f / *s0++; } } @@ -219,39 +275,50 @@ static int pack_B_wq_int8(const Mat& B, const Mat& B_scales, Mat& packed_B, Mat& } // group-major, row-major within each K4/K2/K1 fragment -static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int block_size, const float* input_scale_ptr) +static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int k, int max_kk, int block_size, const float* input_scale_ptr) { #if NCNN_RUNTIME_CPU && NCNN_MMI && !__mips_msa && !__mips_loongson_mmi if (ncnn::cpu_support_loongson_mmi()) { - quantize_A_tile_wq_int8_loongson_mmi(A, AT_tile, AT_descales_tile, i, max_ii, block_size, input_scale_ptr); + quantize_A_tile_wq_int8_loongson_mmi(A, AT_tile, AT_descales_tile, i, max_ii, k, max_kk, block_size, input_scale_ptr); return; } #endif signed char* pp = AT_tile; float* pd = AT_descales_tile; - const int K = AT_tile.w; - const int block_count = AT_descales_tile.w; + const int K = max_kk; + const int block_count = (max_kk + block_size - 1) / block_size; const size_t A_hstep = A.dims == 3 ? A.cstep : (size_t)A.w; + const float* A_data = (const float*)A + k; + input_scale_ptr = input_scale_ptr ? input_scale_ptr + k : 0; int ii = 0; #if __mips_msa for (; ii + 7 < max_ii; ii += 8) { - const float* p0 = (const float*)A + (size_t)(i + ii) * A_hstep; - const float* p1 = (const float*)A + (size_t)(i + ii + 1) * A_hstep; - const float* p2 = (const float*)A + (size_t)(i + ii + 2) * A_hstep; - const float* p3 = (const float*)A + (size_t)(i + ii + 3) * A_hstep; - const float* p4 = (const float*)A + (size_t)(i + ii + 4) * A_hstep; - const float* p5 = (const float*)A + (size_t)(i + ii + 5) * A_hstep; - const float* p6 = (const float*)A + (size_t)(i + ii + 6) * A_hstep; - const float* p7 = (const float*)A + (size_t)(i + ii + 7) * A_hstep; + const float* p0 = A_data + (size_t)(i + ii) * A_hstep; + const float* p1 = A_data + (size_t)(i + ii + 1) * A_hstep; + const float* p2 = A_data + (size_t)(i + ii + 2) * A_hstep; + const float* p3 = A_data + (size_t)(i + ii + 3) * A_hstep; + const float* p4 = A_data + (size_t)(i + ii + 4) * A_hstep; + const float* p5 = A_data + (size_t)(i + ii + 5) * A_hstep; + const float* p6 = A_data + (size_t)(i + ii + 6) * A_hstep; + const float* p7 = A_data + (size_t)(i + ii + 7) * A_hstep; for (int g = 0; g < block_count; g++) { const int k0 = g * block_size; const int max_kk = std::min(K - k0, block_size); + const float* p0g = p0 + k0; + const float* p1g = p1 + k0; + const float* p2g = p2 + k0; + const float* p3g = p3 + k0; + const float* p4g = p4 + k0; + const float* p5g = p5 + k0; + const float* p6g = p6 + k0; + const float* p7g = p7 + k0; + const float* sg = input_scale_ptr ? input_scale_ptr + k0 : 0; const v16u8 _abs_mask = (v16u8)__msa_fill_w(0x7fffffff); v4f32 _absmax0 = (v4f32)__msa_fill_w(0); v4f32 _absmax1 = (v4f32)__msa_fill_w(0); @@ -262,20 +329,29 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales v4f32 _absmax6 = (v4f32)__msa_fill_w(0); v4f32 _absmax7 = (v4f32)__msa_fill_w(0); + const float* p0a = p0g; + const float* p1a = p1g; + const float* p2a = p2g; + const float* p3a = p3g; + const float* p4a = p4g; + const float* p5a = p5g; + const float* p6a = p6g; + const float* p7a = p7g; + const float* psa = sg; int kk = 0; for (; kk + 3 < max_kk; kk += 4) { - v4f32 _p0 = (v4f32)__msa_ld_w(p0 + k0 + kk, 0); - v4f32 _p1 = (v4f32)__msa_ld_w(p1 + k0 + kk, 0); - v4f32 _p2 = (v4f32)__msa_ld_w(p2 + k0 + kk, 0); - v4f32 _p3 = (v4f32)__msa_ld_w(p3 + k0 + kk, 0); - v4f32 _p4 = (v4f32)__msa_ld_w(p4 + k0 + kk, 0); - v4f32 _p5 = (v4f32)__msa_ld_w(p5 + k0 + kk, 0); - v4f32 _p6 = (v4f32)__msa_ld_w(p6 + k0 + kk, 0); - v4f32 _p7 = (v4f32)__msa_ld_w(p7 + k0 + kk, 0); - if (input_scale_ptr) - { - const v4f32 _s = (v4f32)__msa_ld_w(input_scale_ptr + k0 + kk, 0); + v4f32 _p0 = (v4f32)__msa_ld_w(p0a, 0); + v4f32 _p1 = (v4f32)__msa_ld_w(p1a, 0); + v4f32 _p2 = (v4f32)__msa_ld_w(p2a, 0); + v4f32 _p3 = (v4f32)__msa_ld_w(p3a, 0); + v4f32 _p4 = (v4f32)__msa_ld_w(p4a, 0); + v4f32 _p5 = (v4f32)__msa_ld_w(p5a, 0); + v4f32 _p6 = (v4f32)__msa_ld_w(p6a, 0); + v4f32 _p7 = (v4f32)__msa_ld_w(p7a, 0); + if (psa) + { + v4f32 _s = (v4f32)__msa_ld_w(psa, 0); _p0 = __msa_fmul_w(_p0, _s); _p1 = __msa_fmul_w(_p1, _s); _p2 = __msa_fmul_w(_p2, _s); @@ -293,6 +369,16 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales _absmax5 = __msa_fmax_w(_absmax5, (v4f32)__msa_and_v((v16u8)_p5, _abs_mask)); _absmax6 = __msa_fmax_w(_absmax6, (v4f32)__msa_and_v((v16u8)_p6, _abs_mask)); _absmax7 = __msa_fmax_w(_absmax7, (v4f32)__msa_and_v((v16u8)_p7, _abs_mask)); + p0a += 4; + p1a += 4; + p2a += 4; + p3a += 4; + p4a += 4; + p5a += 4; + p6a += 4; + p7a += 4; + if (psa) + psa += 4; } float absmax0 = __msa_reduce_fmax_w(_absmax0); @@ -306,34 +392,25 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales for (; kk < max_kk; kk++) { - const int k = k0 + kk; - const float s = input_scale_ptr ? input_scale_ptr[k] : 1.f; - absmax0 = std::max(absmax0, fabsf(p0[k] * s)); - absmax1 = std::max(absmax1, fabsf(p1[k] * s)); - absmax2 = std::max(absmax2, fabsf(p2[k] * s)); - absmax3 = std::max(absmax3, fabsf(p3[k] * s)); - absmax4 = std::max(absmax4, fabsf(p4[k] * s)); - absmax5 = std::max(absmax5, fabsf(p5[k] * s)); - absmax6 = std::max(absmax6, fabsf(p6[k] * s)); - absmax7 = std::max(absmax7, fabsf(p7[k] * s)); - } - - volatile double scale0_fp64 = absmax0 == 0.f ? 1.0 : 127.0 / (double)absmax0; - volatile double scale1_fp64 = absmax1 == 0.f ? 1.0 : 127.0 / (double)absmax1; - volatile double scale2_fp64 = absmax2 == 0.f ? 1.0 : 127.0 / (double)absmax2; - volatile double scale3_fp64 = absmax3 == 0.f ? 1.0 : 127.0 / (double)absmax3; - volatile double scale4_fp64 = absmax4 == 0.f ? 1.0 : 127.0 / (double)absmax4; - volatile double scale5_fp64 = absmax5 == 0.f ? 1.0 : 127.0 / (double)absmax5; - volatile double scale6_fp64 = absmax6 == 0.f ? 1.0 : 127.0 / (double)absmax6; - volatile double scale7_fp64 = absmax7 == 0.f ? 1.0 : 127.0 / (double)absmax7; - const float scale0 = (float)scale0_fp64; - const float scale1 = (float)scale1_fp64; - const float scale2 = (float)scale2_fp64; - const float scale3 = (float)scale3_fp64; - const float scale4 = (float)scale4_fp64; - const float scale5 = (float)scale5_fp64; - const float scale6 = (float)scale6_fp64; - const float scale7 = (float)scale7_fp64; + const float s = psa ? *psa++ : 1.f; + absmax0 = std::max(absmax0, fabsf(*p0a++ * s)); + absmax1 = std::max(absmax1, fabsf(*p1a++ * s)); + absmax2 = std::max(absmax2, fabsf(*p2a++ * s)); + absmax3 = std::max(absmax3, fabsf(*p3a++ * s)); + absmax4 = std::max(absmax4, fabsf(*p4a++ * s)); + absmax5 = std::max(absmax5, fabsf(*p5a++ * s)); + absmax6 = std::max(absmax6, fabsf(*p6a++ * s)); + absmax7 = std::max(absmax7, fabsf(*p7a++ * s)); + } + + const float scale0 = absmax0 == 0.f ? 1.f : 127.f / absmax0; + const float scale1 = absmax1 == 0.f ? 1.f : 127.f / absmax1; + const float scale2 = absmax2 == 0.f ? 1.f : 127.f / absmax2; + const float scale3 = absmax3 == 0.f ? 1.f : 127.f / absmax3; + const float scale4 = absmax4 == 0.f ? 1.f : 127.f / absmax4; + const float scale5 = absmax5 == 0.f ? 1.f : 127.f / absmax5; + const float scale6 = absmax6 == 0.f ? 1.f : 127.f / absmax6; + const float scale7 = absmax7 == 0.f ? 1.f : 127.f / absmax7; pd[0] = absmax0 / 127.f; pd[1] = absmax1 / 127.f; pd[2] = absmax2 / 127.f; @@ -344,72 +421,99 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales pd[7] = absmax7 / 127.f; pd += 8; + const float* p0q = p0g; + const float* p1q = p1g; + const float* p2q = p2g; + const float* p3q = p3g; + const float* p4q = p4g; + const float* p5q = p5g; + const float* p6q = p6g; + const float* p7q = p7g; + const float* psq = sg; kk = 0; for (; kk + 3 < max_kk; kk += 4) { - const v4f32 _s = input_scale_ptr ? (v4f32)__msa_ld_w(input_scale_ptr + k0 + kk, 0) : __msa_fill_w_f32(1.f); - v4f32 _p = (v4f32)__msa_ld_w(p0 + k0 + kk, 0); + v4f32 _s = psq ? (v4f32)__msa_ld_w(psq, 0) : __msa_fill_w_f32(1.f); + v4f32 _p = (v4f32)__msa_ld_w(p0q, 0); _p = __msa_fmul_w(__msa_fmul_w(_p, _s), __msa_fill_w_f32(scale0)); ((int*)pp)[0] = __msa_copy_s_w((v4i32)float2int8(_p), 0); - _p = (v4f32)__msa_ld_w(p1 + k0 + kk, 0); + _p = (v4f32)__msa_ld_w(p1q, 0); _p = __msa_fmul_w(__msa_fmul_w(_p, _s), __msa_fill_w_f32(scale1)); ((int*)pp)[1] = __msa_copy_s_w((v4i32)float2int8(_p), 0); - _p = (v4f32)__msa_ld_w(p2 + k0 + kk, 0); + _p = (v4f32)__msa_ld_w(p2q, 0); _p = __msa_fmul_w(__msa_fmul_w(_p, _s), __msa_fill_w_f32(scale2)); ((int*)pp)[2] = __msa_copy_s_w((v4i32)float2int8(_p), 0); - _p = (v4f32)__msa_ld_w(p3 + k0 + kk, 0); + _p = (v4f32)__msa_ld_w(p3q, 0); _p = __msa_fmul_w(__msa_fmul_w(_p, _s), __msa_fill_w_f32(scale3)); ((int*)pp)[3] = __msa_copy_s_w((v4i32)float2int8(_p), 0); - _p = (v4f32)__msa_ld_w(p4 + k0 + kk, 0); + _p = (v4f32)__msa_ld_w(p4q, 0); _p = __msa_fmul_w(__msa_fmul_w(_p, _s), __msa_fill_w_f32(scale4)); ((int*)pp)[4] = __msa_copy_s_w((v4i32)float2int8(_p), 0); - _p = (v4f32)__msa_ld_w(p5 + k0 + kk, 0); + _p = (v4f32)__msa_ld_w(p5q, 0); _p = __msa_fmul_w(__msa_fmul_w(_p, _s), __msa_fill_w_f32(scale5)); ((int*)pp)[5] = __msa_copy_s_w((v4i32)float2int8(_p), 0); - _p = (v4f32)__msa_ld_w(p6 + k0 + kk, 0); + _p = (v4f32)__msa_ld_w(p6q, 0); _p = __msa_fmul_w(__msa_fmul_w(_p, _s), __msa_fill_w_f32(scale6)); ((int*)pp)[6] = __msa_copy_s_w((v4i32)float2int8(_p), 0); - _p = (v4f32)__msa_ld_w(p7 + k0 + kk, 0); + _p = (v4f32)__msa_ld_w(p7q, 0); _p = __msa_fmul_w(__msa_fmul_w(_p, _s), __msa_fill_w_f32(scale7)); ((int*)pp)[7] = __msa_copy_s_w((v4i32)float2int8(_p), 0); pp += 32; + p0q += 4; + p1q += 4; + p2q += 4; + p3q += 4; + p4q += 4; + p5q += 4; + p6q += 4; + p7q += 4; + if (psq) + psq += 4; } if (kk + 1 < max_kk) { - const int k = k0 + kk; - const float s0 = input_scale_ptr ? input_scale_ptr[k] : 1.f; - const float s1 = input_scale_ptr ? input_scale_ptr[k + 1] : 1.f; - pp[0] = float2int8(p0[k] * s0 * scale0); - pp[1] = float2int8(p0[k + 1] * s1 * scale0); - pp[2] = float2int8(p1[k] * s0 * scale1); - pp[3] = float2int8(p1[k + 1] * s1 * scale1); - pp[4] = float2int8(p2[k] * s0 * scale2); - pp[5] = float2int8(p2[k + 1] * s1 * scale2); - pp[6] = float2int8(p3[k] * s0 * scale3); - pp[7] = float2int8(p3[k + 1] * s1 * scale3); - pp[8] = float2int8(p4[k] * s0 * scale4); - pp[9] = float2int8(p4[k + 1] * s1 * scale4); - pp[10] = float2int8(p5[k] * s0 * scale5); - pp[11] = float2int8(p5[k + 1] * s1 * scale5); - pp[12] = float2int8(p6[k] * s0 * scale6); - pp[13] = float2int8(p6[k + 1] * s1 * scale6); - pp[14] = float2int8(p7[k] * s0 * scale7); - pp[15] = float2int8(p7[k + 1] * s1 * scale7); + const float s0 = psq ? psq[0] : 1.f; + const float s1 = psq ? psq[1] : 1.f; + pp[0] = float2int8(p0q[0] * s0 * scale0); + pp[1] = float2int8(p0q[1] * s1 * scale0); + pp[2] = float2int8(p1q[0] * s0 * scale1); + pp[3] = float2int8(p1q[1] * s1 * scale1); + pp[4] = float2int8(p2q[0] * s0 * scale2); + pp[5] = float2int8(p2q[1] * s1 * scale2); + pp[6] = float2int8(p3q[0] * s0 * scale3); + pp[7] = float2int8(p3q[1] * s1 * scale3); + pp[8] = float2int8(p4q[0] * s0 * scale4); + pp[9] = float2int8(p4q[1] * s1 * scale4); + pp[10] = float2int8(p5q[0] * s0 * scale5); + pp[11] = float2int8(p5q[1] * s1 * scale5); + pp[12] = float2int8(p6q[0] * s0 * scale6); + pp[13] = float2int8(p6q[1] * s1 * scale6); + pp[14] = float2int8(p7q[0] * s0 * scale7); + pp[15] = float2int8(p7q[1] * s1 * scale7); pp += 16; + p0q += 2; + p1q += 2; + p2q += 2; + p3q += 2; + p4q += 2; + p5q += 2; + p6q += 2; + p7q += 2; + if (psq) + psq += 2; kk += 2; } if (kk < max_kk) { - const int k = k0 + kk; - const float s = input_scale_ptr ? input_scale_ptr[k] : 1.f; - pp[0] = float2int8(p0[k] * s * scale0); - pp[1] = float2int8(p1[k] * s * scale1); - pp[2] = float2int8(p2[k] * s * scale2); - pp[3] = float2int8(p3[k] * s * scale3); - pp[4] = float2int8(p4[k] * s * scale4); - pp[5] = float2int8(p5[k] * s * scale5); - pp[6] = float2int8(p6[k] * s * scale6); - pp[7] = float2int8(p7[k] * s * scale7); + const float s = psq ? *psq : 1.f; + pp[0] = float2int8(*p0q * s * scale0); + pp[1] = float2int8(*p1q * s * scale1); + pp[2] = float2int8(*p2q * s * scale2); + pp[3] = float2int8(*p3q * s * scale3); + pp[4] = float2int8(*p4q * s * scale4); + pp[5] = float2int8(*p5q * s * scale5); + pp[6] = float2int8(*p6q * s * scale6); + pp[7] = float2int8(*p7q * s * scale7); pp += 8; } } @@ -420,15 +524,20 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales const int i1 = i + ii + 1; const int i2 = i + ii + 2; const int i3 = i + ii + 3; - const float* p0 = (const float*)A + (size_t)i0 * A_hstep; - const float* p1 = (const float*)A + (size_t)i1 * A_hstep; - const float* p2 = (const float*)A + (size_t)i2 * A_hstep; - const float* p3 = (const float*)A + (size_t)i3 * A_hstep; + const float* p0 = A_data + (size_t)i0 * A_hstep; + const float* p1 = A_data + (size_t)i1 * A_hstep; + const float* p2 = A_data + (size_t)i2 * A_hstep; + const float* p3 = A_data + (size_t)i3 * A_hstep; for (int g = 0; g < block_count; g++) { const int k0 = g * block_size; const int max_kk = std::min(K - k0, block_size); + const float* p0g = p0 + k0; + const float* p1g = p1 + k0; + const float* p2g = p2 + k0; + const float* p3g = p3 + k0; + const float* sg = input_scale_ptr ? input_scale_ptr + k0 : 0; float absmax0 = 0.f; float absmax1 = 0.f; float absmax2 = 0.f; @@ -440,16 +549,21 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales v4f32 _absmax2 = (v4f32)__msa_fill_w(0); v4f32 _absmax3 = (v4f32)__msa_fill_w(0); + const float* p0a = p0g; + const float* p1a = p1g; + const float* p2a = p2g; + const float* p3a = p3g; + const float* psa = sg; int kk = 0; for (; kk + 3 < max_kk; kk += 4) { - v4f32 _p0 = (v4f32)__msa_ld_w(p0 + k0 + kk, 0); - v4f32 _p1 = (v4f32)__msa_ld_w(p1 + k0 + kk, 0); - v4f32 _p2 = (v4f32)__msa_ld_w(p2 + k0 + kk, 0); - v4f32 _p3 = (v4f32)__msa_ld_w(p3 + k0 + kk, 0); - if (input_scale_ptr) + v4f32 _p0 = (v4f32)__msa_ld_w(p0a, 0); + v4f32 _p1 = (v4f32)__msa_ld_w(p1a, 0); + v4f32 _p2 = (v4f32)__msa_ld_w(p2a, 0); + v4f32 _p3 = (v4f32)__msa_ld_w(p3a, 0); + if (psa) { - const v4f32 _s = (v4f32)__msa_ld_w(input_scale_ptr + k0 + kk, 0); + v4f32 _s = (v4f32)__msa_ld_w(psa, 0); _p0 = __msa_fmul_w(_p0, _s); _p1 = __msa_fmul_w(_p1, _s); _p2 = __msa_fmul_w(_p2, _s); @@ -459,6 +573,12 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales _absmax1 = __msa_fmax_w(_absmax1, (v4f32)__msa_and_v((v16u8)_p1, _abs_mask)); _absmax2 = __msa_fmax_w(_absmax2, (v4f32)__msa_and_v((v16u8)_p2, _abs_mask)); _absmax3 = __msa_fmax_w(_absmax3, (v4f32)__msa_and_v((v16u8)_p3, _abs_mask)); + p0a += 4; + p1a += 4; + p2a += 4; + p3a += 4; + if (psa) + psa += 4; } absmax0 = __msa_reduce_fmax_w(_absmax0); absmax1 = __msa_reduce_fmax_w(_absmax1); @@ -467,13 +587,13 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales for (; kk < max_kk; kk++) { - float v0 = p0[k0 + kk]; - float v1 = p1[k0 + kk]; - float v2 = p2[k0 + kk]; - float v3 = p3[k0 + kk]; - if (input_scale_ptr) + float v0 = *p0a++; + float v1 = *p1a++; + float v2 = *p2a++; + float v3 = *p3a++; + if (psa) { - const float s = input_scale_ptr[k0 + kk]; + const float s = *psa++; v0 *= s; v1 *= s; v2 *= s; @@ -485,45 +605,46 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales absmax3 = std::max(absmax3, fabsf(v3)); } - volatile double scale0_fp64 = absmax0 == 0.f ? 1.0 : 127.0 / (double)absmax0; - volatile double scale1_fp64 = absmax1 == 0.f ? 1.0 : 127.0 / (double)absmax1; - volatile double scale2_fp64 = absmax2 == 0.f ? 1.0 : 127.0 / (double)absmax2; - volatile double scale3_fp64 = absmax3 == 0.f ? 1.0 : 127.0 / (double)absmax3; - const float scale0 = (float)scale0_fp64; - const float scale1 = (float)scale1_fp64; - const float scale2 = (float)scale2_fp64; - const float scale3 = (float)scale3_fp64; + const float scale0 = absmax0 == 0.f ? 1.f : 127.f / absmax0; + const float scale1 = absmax1 == 0.f ? 1.f : 127.f / absmax1; + const float scale2 = absmax2 == 0.f ? 1.f : 127.f / absmax2; + const float scale3 = absmax3 == 0.f ? 1.f : 127.f / absmax3; pd[0] = absmax0 / 127.f; pd[1] = absmax1 / 127.f; pd[2] = absmax2 / 127.f; pd[3] = absmax3 / 127.f; pd += 4; + const float* p0q = p0g; + const float* p1q = p1g; + const float* p2q = p2g; + const float* p3q = p3g; + const float* psq = sg; kk = 0; for (; kk + 3 < max_kk; kk += 4) { - float v00 = p0[k0 + kk]; - float v01 = p0[k0 + kk + 1]; - float v02 = p0[k0 + kk + 2]; - float v03 = p0[k0 + kk + 3]; - float v10 = p1[k0 + kk]; - float v11 = p1[k0 + kk + 1]; - float v12 = p1[k0 + kk + 2]; - float v13 = p1[k0 + kk + 3]; - float v20 = p2[k0 + kk]; - float v21 = p2[k0 + kk + 1]; - float v22 = p2[k0 + kk + 2]; - float v23 = p2[k0 + kk + 3]; - float v30 = p3[k0 + kk]; - float v31 = p3[k0 + kk + 1]; - float v32 = p3[k0 + kk + 2]; - float v33 = p3[k0 + kk + 3]; - if (input_scale_ptr) - { - const float s0 = input_scale_ptr[k0 + kk]; - const float s1 = input_scale_ptr[k0 + kk + 1]; - const float s2 = input_scale_ptr[k0 + kk + 2]; - const float s3 = input_scale_ptr[k0 + kk + 3]; + float v00 = p0q[0]; + float v01 = p0q[1]; + float v02 = p0q[2]; + float v03 = p0q[3]; + float v10 = p1q[0]; + float v11 = p1q[1]; + float v12 = p1q[2]; + float v13 = p1q[3]; + float v20 = p2q[0]; + float v21 = p2q[1]; + float v22 = p2q[2]; + float v23 = p2q[3]; + float v30 = p3q[0]; + float v31 = p3q[1]; + float v32 = p3q[2]; + float v33 = p3q[3]; + if (psq) + { + const float s0 = psq[0]; + const float s1 = psq[1]; + const float s2 = psq[2]; + const float s3 = psq[3]; v00 *= s0; v01 *= s1; v02 *= s2; @@ -540,14 +661,6 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales v31 *= s1; v32 *= s2; v33 *= s3; - asm volatile("" - : "+f"(v00), "+f"(v01), "+f"(v02), "+f"(v03)); - asm volatile("" - : "+f"(v10), "+f"(v11), "+f"(v12), "+f"(v13)); - asm volatile("" - : "+f"(v20), "+f"(v21), "+f"(v22), "+f"(v23)); - asm volatile("" - : "+f"(v30), "+f"(v31), "+f"(v32), "+f"(v33)); } pp[0] = float2int8(v00 * scale0); pp[1] = float2int8(v01 * scale0); @@ -566,21 +679,27 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales pp[14] = float2int8(v32 * scale3); pp[15] = float2int8(v33 * scale3); pp += 16; + p0q += 4; + p1q += 4; + p2q += 4; + p3q += 4; + if (psq) + psq += 4; } if (kk + 1 < max_kk) { - float v00 = p0[k0 + kk]; - float v01 = p0[k0 + kk + 1]; - float v10 = p1[k0 + kk]; - float v11 = p1[k0 + kk + 1]; - float v20 = p2[k0 + kk]; - float v21 = p2[k0 + kk + 1]; - float v30 = p3[k0 + kk]; - float v31 = p3[k0 + kk + 1]; - if (input_scale_ptr) - { - const float s0 = input_scale_ptr[k0 + kk]; - const float s1 = input_scale_ptr[k0 + kk + 1]; + float v00 = p0q[0]; + float v01 = p0q[1]; + float v10 = p1q[0]; + float v11 = p1q[1]; + float v20 = p2q[0]; + float v21 = p2q[1]; + float v30 = p3q[0]; + float v31 = p3q[1]; + if (psq) + { + const float s0 = psq[0]; + const float s1 = psq[1]; v00 *= s0; v01 *= s1; v10 *= s0; @@ -589,10 +708,6 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales v21 *= s1; v30 *= s0; v31 *= s1; - asm volatile("" - : "+f"(v00), "+f"(v01), "+f"(v10), "+f"(v11)); - asm volatile("" - : "+f"(v20), "+f"(v21), "+f"(v30), "+f"(v31)); } pp[0] = float2int8(v00 * scale0); pp[1] = float2int8(v01 * scale0); @@ -603,24 +718,27 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales pp[6] = float2int8(v30 * scale3); pp[7] = float2int8(v31 * scale3); pp += 8; + p0q += 2; + p1q += 2; + p2q += 2; + p3q += 2; + if (psq) + psq += 2; kk += 2; } for (; kk < max_kk; kk++) { - const int k = k0 + kk; - float v0 = p0[k]; - float v1 = p1[k]; - float v2 = p2[k]; - float v3 = p3[k]; - if (input_scale_ptr) + float v0 = *p0q++; + float v1 = *p1q++; + float v2 = *p2q++; + float v3 = *p3q++; + if (psq) { - const float s = input_scale_ptr[k]; + const float s = *psq++; v0 *= s; v1 *= s; v2 *= s; v3 *= s; - asm volatile("" - : "+f"(v0), "+f"(v1), "+f"(v2), "+f"(v3)); } pp[0] = float2int8(v0 * scale0); pp[1] = float2int8(v1 * scale1); @@ -635,23 +753,28 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales { const int i0 = i + ii; const int i1 = i + ii + 1; - const float* p0 = (const float*)A + (size_t)i0 * A_hstep; - const float* p1 = (const float*)A + (size_t)i1 * A_hstep; + const float* p0 = A_data + (size_t)i0 * A_hstep; + const float* p1 = A_data + (size_t)i1 * A_hstep; for (int g = 0; g < block_count; g++) { const int k0 = g * block_size; const int max_kk = std::min(K - k0, block_size); + const float* p0g = p0 + k0; + const float* p1g = p1 + k0; + const float* sg = input_scale_ptr ? input_scale_ptr + k0 : 0; float absmax0 = 0.f; float absmax1 = 0.f; + const float* p0a = p0g; + const float* p1a = p1g; + const float* psa = sg; for (int kk = 0; kk < max_kk; kk++) { - const int k = k0 + kk; - float v0 = p0[k]; - float v1 = p1[k]; - if (input_scale_ptr) + float v0 = *p0a++; + float v1 = *p1a++; + if (psa) { - const float s = input_scale_ptr[k]; + const float s = *psa++; v0 *= s; v1 *= s; } @@ -659,65 +782,84 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales absmax1 = std::max(absmax1, fabsf(v1)); } - volatile double scale0_fp64 = absmax0 == 0.f ? 1.0 : 127.0 / (double)absmax0; - volatile double scale1_fp64 = absmax1 == 0.f ? 1.0 : 127.0 / (double)absmax1; - const float scale0 = (float)scale0_fp64; - const float scale1 = (float)scale1_fp64; + const float scale0 = absmax0 == 0.f ? 1.f : 127.f / absmax0; + const float scale1 = absmax1 == 0.f ? 1.f : 127.f / absmax1; pd[0] = absmax0 / 127.f; pd[1] = absmax1 / 127.f; pd += 2; + const float* p0q = p0g; + const float* p1q = p1g; + const float* psq = sg; int kk = 0; for (; kk + 3 < max_kk; kk += 4) { - for (int r = 0; r < 4; r++) - { - const int k = k0 + kk + r; - float v0 = p0[k]; - float v1 = p1[k]; - if (input_scale_ptr) - { - v0 *= input_scale_ptr[k]; - v1 *= input_scale_ptr[k]; - asm volatile("" - : "+f"(v0), "+f"(v1)); - } - pp[r] = float2int8(v0 * scale0); - pp[4 + r] = float2int8(v1 * scale1); + float v00 = p0q[0]; + float v01 = p0q[1]; + float v02 = p0q[2]; + float v03 = p0q[3]; + float v10 = p1q[0]; + float v11 = p1q[1]; + float v12 = p1q[2]; + float v13 = p1q[3]; + if (psq) + { + v00 *= psq[0]; + v01 *= psq[1]; + v02 *= psq[2]; + v03 *= psq[3]; + v10 *= psq[0]; + v11 *= psq[1]; + v12 *= psq[2]; + v13 *= psq[3]; } + pp[0] = float2int8(v00 * scale0); + pp[1] = float2int8(v01 * scale0); + pp[2] = float2int8(v02 * scale0); + pp[3] = float2int8(v03 * scale0); + pp[4] = float2int8(v10 * scale1); + pp[5] = float2int8(v11 * scale1); + pp[6] = float2int8(v12 * scale1); + pp[7] = float2int8(v13 * scale1); pp += 8; + p0q += 4; + p1q += 4; + if (psq) + psq += 4; } if (kk + 1 < max_kk) { - for (int r = 0; r < 2; r++) + float v00 = p0q[0]; + float v01 = p0q[1]; + float v10 = p1q[0]; + float v11 = p1q[1]; + if (psq) { - const int k = k0 + kk + r; - float v0 = p0[k]; - float v1 = p1[k]; - if (input_scale_ptr) - { - v0 *= input_scale_ptr[k]; - v1 *= input_scale_ptr[k]; - asm volatile("" - : "+f"(v0), "+f"(v1)); - } - pp[r] = float2int8(v0 * scale0); - pp[2 + r] = float2int8(v1 * scale1); + v00 *= psq[0]; + v01 *= psq[1]; + v10 *= psq[0]; + v11 *= psq[1]; } + pp[0] = float2int8(v00 * scale0); + pp[1] = float2int8(v01 * scale0); + pp[2] = float2int8(v10 * scale1); + pp[3] = float2int8(v11 * scale1); pp += 4; + p0q += 2; + p1q += 2; + if (psq) + psq += 2; kk += 2; } for (; kk < max_kk; kk++) { - const int k = k0 + kk; - float v0 = p0[k]; - float v1 = p1[k]; - if (input_scale_ptr) + float v0 = *p0q++; + float v1 = *p1q++; + if (psq) { - v0 *= input_scale_ptr[k]; - v1 *= input_scale_ptr[k]; - asm volatile("" - : "+f"(v0), "+f"(v1)); + const float s = *psq++; + v0 *= s; + v1 *= s; } pp[0] = float2int8(v0 * scale0); pp[1] = float2int8(v1 * scale1); @@ -728,70 +870,75 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales for (; ii < max_ii; ii++) { const int i0 = i + ii; - const float* p0 = (const float*)A + (size_t)i0 * A_hstep; + const float* p0 = A_data + (size_t)i0 * A_hstep; for (int g = 0; g < block_count; g++) { const int k0 = g * block_size; const int max_kk = std::min(K - k0, block_size); + const float* p0g = p0 + k0; + const float* sg = input_scale_ptr ? input_scale_ptr + k0 : 0; float absmax0 = 0.f; + const float* p0a = p0g; + const float* psa = sg; for (int kk = 0; kk < max_kk; kk++) { - const int k = k0 + kk; - float v0 = p0[k]; - if (input_scale_ptr) - v0 *= input_scale_ptr[k]; + float v0 = *p0a++; + if (psa) + v0 *= *psa++; absmax0 = std::max(absmax0, fabsf(v0)); } - volatile double scale0_fp64 = absmax0 == 0.f ? 1.0 : 127.0 / (double)absmax0; - const float scale0 = (float)scale0_fp64; + const float scale0 = absmax0 == 0.f ? 1.f : 127.f / absmax0; *pd++ = absmax0 / 127.f; + const float* p0q = p0g; + const float* psq = sg; int kk = 0; for (; kk + 3 < max_kk; kk += 4) { - for (int r = 0; r < 4; r++) + float v0 = p0q[0]; + float v1 = p0q[1]; + float v2 = p0q[2]; + float v3 = p0q[3]; + if (psq) { - const int k = k0 + kk + r; - float v0 = p0[k]; - if (input_scale_ptr) - { - v0 *= input_scale_ptr[k]; - asm volatile("" - : "+f"(v0)); - } - pp[r] = float2int8(v0 * scale0); + v0 *= psq[0]; + v1 *= psq[1]; + v2 *= psq[2]; + v3 *= psq[3]; } + pp[0] = float2int8(v0 * scale0); + pp[1] = float2int8(v1 * scale0); + pp[2] = float2int8(v2 * scale0); + pp[3] = float2int8(v3 * scale0); pp += 4; + p0q += 4; + if (psq) + psq += 4; } if (kk + 1 < max_kk) { - for (int r = 0; r < 2; r++) + float v0 = p0q[0]; + float v1 = p0q[1]; + if (psq) { - const int k = k0 + kk + r; - float v0 = p0[k]; - if (input_scale_ptr) - { - v0 *= input_scale_ptr[k]; - asm volatile("" - : "+f"(v0)); - } - pp[r] = float2int8(v0 * scale0); + v0 *= psq[0]; + v1 *= psq[1]; } + pp[0] = float2int8(v0 * scale0); + pp[1] = float2int8(v1 * scale0); pp += 2; + p0q += 2; + if (psq) + psq += 2; kk += 2; } for (; kk < max_kk; kk++) { - const int k = k0 + kk; - float v0 = p0[k]; - if (input_scale_ptr) - { - v0 *= input_scale_ptr[k]; - asm volatile("" - : "+f"(v0)); - } + float v0 = *p0q++; + if (psq) + v0 *= *psq++; *pp++ = float2int8(v0 * scale0); } } @@ -799,21 +946,23 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales } // group-major, row-major within each K4/K2/K1 fragment -static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int block_size, const float* input_scale_ptr) +static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int k, int max_kk, int block_size, const float* input_scale_ptr) { #if NCNN_RUNTIME_CPU && NCNN_MMI && !__mips_msa && !__mips_loongson_mmi if (ncnn::cpu_support_loongson_mmi()) { - transpose_quantize_A_tile_wq_int8_loongson_mmi(A, AT_tile, AT_descales_tile, i, max_ii, block_size, input_scale_ptr); + transpose_quantize_A_tile_wq_int8_loongson_mmi(A, AT_tile, AT_descales_tile, i, max_ii, k, max_kk, block_size, input_scale_ptr); return; } #endif signed char* pp = AT_tile; float* pd = AT_descales_tile; - const int K = AT_tile.w; - const int block_count = AT_descales_tile.w; + const int K = max_kk; + const int block_count = (max_kk + block_size - 1) / block_size; const size_t A_hstep = A.dims == 3 ? A.cstep : (size_t)A.w; + const float* A_data = (const float*)A + (size_t)k * A_hstep; + input_scale_ptr = input_scale_ptr ? input_scale_ptr + k : 0; int ii = 0; #if __mips_msa @@ -825,45 +974,41 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A { const int k0 = g * block_size; const int max_kk = std::min(K - k0, block_size); + const float* p0g = A_data + (size_t)k0 * A_hstep + i0; + const float* sg = input_scale_ptr ? input_scale_ptr + k0 : 0; const v16u8 _abs_mask = (v16u8)__msa_fill_w(0x7fffffff); v4f32 _absmax0 = (v4f32)__msa_fill_w(0); v4f32 _absmax1 = (v4f32)__msa_fill_w(0); + const float* p0a = p0g; + const float* psa = sg; int kk = 0; for (; kk < max_kk; kk++) { - const int k = k0 + kk; - v4f32 _p0 = (v4f32)__msa_ld_w((const float*)A + (size_t)k * A_hstep + i0, 0); - v4f32 _p1 = (v4f32)__msa_ld_w((const float*)A + (size_t)k * A_hstep + i0 + 4, 0); - if (input_scale_ptr) + v4f32 _p0 = (v4f32)__msa_ld_w(p0a, 0); + v4f32 _p1 = (v4f32)__msa_ld_w(p0a + 4, 0); + if (psa) { - const v4f32 _s = __msa_fill_w_f32(input_scale_ptr[k]); + v4f32 _s = __msa_fill_w_f32(*psa++); _p0 = __msa_fmul_w(_p0, _s); _p1 = __msa_fmul_w(_p1, _s); } _absmax0 = __msa_fmax_w(_absmax0, (v4f32)__msa_and_v((v16u8)_p0, _abs_mask)); _absmax1 = __msa_fmax_w(_absmax1, (v4f32)__msa_and_v((v16u8)_p1, _abs_mask)); + p0a += A_hstep; } float absmax[8]; __msa_st_w((v4i32)_absmax0, absmax, 0); __msa_st_w((v4i32)_absmax1, absmax + 4, 0); - volatile double scale0_fp64 = absmax[0] == 0.f ? 1.0 : 127.0 / (double)absmax[0]; - volatile double scale1_fp64 = absmax[1] == 0.f ? 1.0 : 127.0 / (double)absmax[1]; - volatile double scale2_fp64 = absmax[2] == 0.f ? 1.0 : 127.0 / (double)absmax[2]; - volatile double scale3_fp64 = absmax[3] == 0.f ? 1.0 : 127.0 / (double)absmax[3]; - volatile double scale4_fp64 = absmax[4] == 0.f ? 1.0 : 127.0 / (double)absmax[4]; - volatile double scale5_fp64 = absmax[5] == 0.f ? 1.0 : 127.0 / (double)absmax[5]; - volatile double scale6_fp64 = absmax[6] == 0.f ? 1.0 : 127.0 / (double)absmax[6]; - volatile double scale7_fp64 = absmax[7] == 0.f ? 1.0 : 127.0 / (double)absmax[7]; - const float scale0 = (float)scale0_fp64; - const float scale1 = (float)scale1_fp64; - const float scale2 = (float)scale2_fp64; - const float scale3 = (float)scale3_fp64; - const float scale4 = (float)scale4_fp64; - const float scale5 = (float)scale5_fp64; - const float scale6 = (float)scale6_fp64; - const float scale7 = (float)scale7_fp64; + const float scale0 = absmax[0] == 0.f ? 1.f : 127.f / absmax[0]; + const float scale1 = absmax[1] == 0.f ? 1.f : 127.f / absmax[1]; + const float scale2 = absmax[2] == 0.f ? 1.f : 127.f / absmax[2]; + const float scale3 = absmax[3] == 0.f ? 1.f : 127.f / absmax[3]; + const float scale4 = absmax[4] == 0.f ? 1.f : 127.f / absmax[4]; + const float scale5 = absmax[5] == 0.f ? 1.f : 127.f / absmax[5]; + const float scale6 = absmax[6] == 0.f ? 1.f : 127.f / absmax[6]; + const float scale7 = absmax[7] == 0.f ? 1.f : 127.f / absmax[7]; pd[0] = absmax[0] / 127.f; pd[1] = absmax[1] / 127.f; pd[2] = absmax[2] / 127.f; @@ -874,10 +1019,12 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A pd[7] = absmax[7] / 127.f; pd += 8; + const float* p0q = p0g; + const float* psq = sg; kk = 0; for (; kk + 3 < max_kk; kk += 4) { - const float* p0 = (const float*)A + (size_t)(k0 + kk) * A_hstep + i0; + const float* p0 = p0q; const float* p1 = p0 + A_hstep; const float* p2 = p1 + A_hstep; const float* p3 = p2 + A_hstep; @@ -885,12 +1032,12 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A v4f32 _p1 = (v4f32)__msa_ld_w(p1, 0); v4f32 _p2 = (v4f32)__msa_ld_w(p2, 0); v4f32 _p3 = (v4f32)__msa_ld_w(p3, 0); - if (input_scale_ptr) + if (psq) { - _p0 = __msa_fmul_w(_p0, __msa_fill_w_f32(input_scale_ptr[k0 + kk])); - _p1 = __msa_fmul_w(_p1, __msa_fill_w_f32(input_scale_ptr[k0 + kk + 1])); - _p2 = __msa_fmul_w(_p2, __msa_fill_w_f32(input_scale_ptr[k0 + kk + 2])); - _p3 = __msa_fmul_w(_p3, __msa_fill_w_f32(input_scale_ptr[k0 + kk + 3])); + _p0 = __msa_fmul_w(_p0, __msa_fill_w_f32(psq[0])); + _p1 = __msa_fmul_w(_p1, __msa_fill_w_f32(psq[1])); + _p2 = __msa_fmul_w(_p2, __msa_fill_w_f32(psq[2])); + _p3 = __msa_fmul_w(_p3, __msa_fill_w_f32(psq[3])); } transpose4x4_ps(_p0, _p1, _p2, _p3); ((int*)pp)[0] = __msa_copy_s_w((v4i32)float2int8(__msa_fmul_w(_p0, __msa_fill_w_f32(scale0))), 0); @@ -902,12 +1049,12 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A _p1 = (v4f32)__msa_ld_w(p1 + 4, 0); _p2 = (v4f32)__msa_ld_w(p2 + 4, 0); _p3 = (v4f32)__msa_ld_w(p3 + 4, 0); - if (input_scale_ptr) + if (psq) { - _p0 = __msa_fmul_w(_p0, __msa_fill_w_f32(input_scale_ptr[k0 + kk])); - _p1 = __msa_fmul_w(_p1, __msa_fill_w_f32(input_scale_ptr[k0 + kk + 1])); - _p2 = __msa_fmul_w(_p2, __msa_fill_w_f32(input_scale_ptr[k0 + kk + 2])); - _p3 = __msa_fmul_w(_p3, __msa_fill_w_f32(input_scale_ptr[k0 + kk + 3])); + _p0 = __msa_fmul_w(_p0, __msa_fill_w_f32(psq[0])); + _p1 = __msa_fmul_w(_p1, __msa_fill_w_f32(psq[1])); + _p2 = __msa_fmul_w(_p2, __msa_fill_w_f32(psq[2])); + _p3 = __msa_fmul_w(_p3, __msa_fill_w_f32(psq[3])); } transpose4x4_ps(_p0, _p1, _p2, _p3); ((int*)pp)[4] = __msa_copy_s_w((v4i32)float2int8(__msa_fmul_w(_p0, __msa_fill_w_f32(scale4))), 0); @@ -915,13 +1062,16 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A ((int*)pp)[6] = __msa_copy_s_w((v4i32)float2int8(__msa_fmul_w(_p2, __msa_fill_w_f32(scale6))), 0); ((int*)pp)[7] = __msa_copy_s_w((v4i32)float2int8(__msa_fmul_w(_p3, __msa_fill_w_f32(scale7))), 0); pp += 32; + p0q = p3 + A_hstep; + if (psq) + psq += 4; } if (kk + 1 < max_kk) { - const float* p0 = (const float*)A + (size_t)(k0 + kk) * A_hstep + i0; + const float* p0 = p0q; const float* p1 = p0 + A_hstep; - const float s0 = input_scale_ptr ? input_scale_ptr[k0 + kk] : 1.f; - const float s1 = input_scale_ptr ? input_scale_ptr[k0 + kk + 1] : 1.f; + const float s0 = psq ? psq[0] : 1.f; + const float s1 = psq ? psq[1] : 1.f; v4f32 _p0 = __msa_fmul_w((v4f32)__msa_ld_w(p0, 0), __msa_fill_w_f32(s0)); v4f32 _p1 = __msa_fmul_w((v4f32)__msa_ld_w(p1, 0), __msa_fill_w_f32(s1)); v16i8 _q0 = float2int8(__msa_fmul_w(_p0, (v4f32) { @@ -955,19 +1105,20 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A pp[14] = __msa_copy_s_b(_q0, 3); pp[15] = __msa_copy_s_b(_q1, 3); pp += 16; + p0q = p1 + A_hstep; + if (psq) + psq += 2; kk += 2; } if (kk < max_kk) { - const int k = k0 + kk; - const float* p0 = (const float*)A + (size_t)k * A_hstep + i0; - const float s = input_scale_ptr ? input_scale_ptr[k] : 1.f; - v4f32 _p0 = __msa_fmul_w((v4f32)__msa_ld_w(p0, 0), __msa_fill_w_f32(s)); - v4f32 _p1 = __msa_fmul_w((v4f32)__msa_ld_w(p0 + 4, 0), __msa_fill_w_f32(s)); - const v16i8 _q0 = float2int8(__msa_fmul_w(_p0, (v4f32) { + const float s = psq ? *psq : 1.f; + v4f32 _p0 = __msa_fmul_w((v4f32)__msa_ld_w(p0q, 0), __msa_fill_w_f32(s)); + v4f32 _p1 = __msa_fmul_w((v4f32)__msa_ld_w(p0q + 4, 0), __msa_fill_w_f32(s)); + v16i8 _q0 = float2int8(__msa_fmul_w(_p0, (v4f32) { scale0, scale1, scale2, scale3 })); - const v16i8 _q1 = float2int8(__msa_fmul_w(_p1, (v4f32) { + v16i8 _q1 = float2int8(__msa_fmul_w(_p1, (v4f32) { scale4, scale5, scale6, scale7 })); ((int*)pp)[0] = __msa_copy_s_w((v4i32)_q0, 0); @@ -984,43 +1135,45 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A { const int k0 = g * block_size; const int max_kk = std::min(K - k0, block_size); + const float* p0g = A_data + (size_t)k0 * A_hstep + i0; + const float* sg = input_scale_ptr ? input_scale_ptr + k0 : 0; const v16u8 _abs_mask = (v16u8)__msa_fill_w(0x7fffffff); v4f32 _absmax = (v4f32)__msa_fill_w(0); + const float* p0a = p0g; + const float* psa = sg; int kk = 0; for (; kk < max_kk; kk++) { - const int k = k0 + kk; - v4f32 _p = (v4f32)__msa_ld_w((const float*)A + (size_t)k * A_hstep + i0, 0); - if (input_scale_ptr) - _p = __msa_fmul_w(_p, __msa_fill_w_f32(input_scale_ptr[k])); + v4f32 _p = (v4f32)__msa_ld_w(p0a, 0); + if (psa) + _p = __msa_fmul_w(_p, __msa_fill_w_f32(*psa++)); _absmax = __msa_fmax_w(_absmax, (v4f32)__msa_and_v((v16u8)_p, _abs_mask)); + p0a += A_hstep; } float absmax[4]; __msa_st_w((v4i32)_absmax, absmax, 0); - volatile double scale0_fp64 = absmax[0] == 0.f ? 1.0 : 127.0 / (double)absmax[0]; - volatile double scale1_fp64 = absmax[1] == 0.f ? 1.0 : 127.0 / (double)absmax[1]; - volatile double scale2_fp64 = absmax[2] == 0.f ? 1.0 : 127.0 / (double)absmax[2]; - volatile double scale3_fp64 = absmax[3] == 0.f ? 1.0 : 127.0 / (double)absmax[3]; - const float scale0 = (float)scale0_fp64; - const float scale1 = (float)scale1_fp64; - const float scale2 = (float)scale2_fp64; - const float scale3 = (float)scale3_fp64; + const float scale0 = absmax[0] == 0.f ? 1.f : 127.f / absmax[0]; + const float scale1 = absmax[1] == 0.f ? 1.f : 127.f / absmax[1]; + const float scale2 = absmax[2] == 0.f ? 1.f : 127.f / absmax[2]; + const float scale3 = absmax[3] == 0.f ? 1.f : 127.f / absmax[3]; pd[0] = absmax[0] / 127.f; pd[1] = absmax[1] / 127.f; pd[2] = absmax[2] / 127.f; pd[3] = absmax[3] / 127.f; pd += 4; - const v4f32 _scale0 = __msa_fill_w_f32(scale0); - const v4f32 _scale1 = __msa_fill_w_f32(scale1); - const v4f32 _scale2 = __msa_fill_w_f32(scale2); - const v4f32 _scale3 = __msa_fill_w_f32(scale3); + v4f32 _scale0 = __msa_fill_w_f32(scale0); + v4f32 _scale1 = __msa_fill_w_f32(scale1); + v4f32 _scale2 = __msa_fill_w_f32(scale2); + v4f32 _scale3 = __msa_fill_w_f32(scale3); + const float* p0q = p0g; + const float* psq = sg; kk = 0; for (; kk + 3 < max_kk; kk += 4) { - const float* p0 = (const float*)A + (size_t)(k0 + kk) * A_hstep + i0; + const float* p0 = p0q; const float* p1 = p0 + A_hstep; const float* p2 = p1 + A_hstep; const float* p3 = p2 + A_hstep; @@ -1028,12 +1181,12 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A v4f32 _p1 = (v4f32)__msa_ld_w(p1, 0); v4f32 _p2 = (v4f32)__msa_ld_w(p2, 0); v4f32 _p3 = (v4f32)__msa_ld_w(p3, 0); - if (input_scale_ptr) + if (psq) { - _p0 = __msa_fmul_w(_p0, __msa_fill_w_f32(input_scale_ptr[k0 + kk])); - _p1 = __msa_fmul_w(_p1, __msa_fill_w_f32(input_scale_ptr[k0 + kk + 1])); - _p2 = __msa_fmul_w(_p2, __msa_fill_w_f32(input_scale_ptr[k0 + kk + 2])); - _p3 = __msa_fmul_w(_p3, __msa_fill_w_f32(input_scale_ptr[k0 + kk + 3])); + _p0 = __msa_fmul_w(_p0, __msa_fill_w_f32(psq[0])); + _p1 = __msa_fmul_w(_p1, __msa_fill_w_f32(psq[1])); + _p2 = __msa_fmul_w(_p2, __msa_fill_w_f32(psq[2])); + _p3 = __msa_fmul_w(_p3, __msa_fill_w_f32(psq[3])); } transpose4x4_ps(_p0, _p1, _p2, _p3); _p0 = __msa_fmul_w(_p0, _scale0); @@ -1045,10 +1198,13 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A ((int*)pp)[2] = __msa_copy_s_w((v4i32)float2int8(_p2), 0); ((int*)pp)[3] = __msa_copy_s_w((v4i32)float2int8(_p3), 0); pp += 16; + p0q = p3 + A_hstep; + if (psq) + psq += 4; } if (kk + 1 < max_kk) { - const float* p0 = (const float*)A + (size_t)(k0 + kk) * A_hstep + i0; + const float* p0 = p0q; const float* p1 = p0 + A_hstep; float v00 = p0[0]; float v10 = p0[1]; @@ -1058,10 +1214,10 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A float v11 = p1[1]; float v21 = p1[2]; float v31 = p1[3]; - if (input_scale_ptr) + if (psq) { - const float s0 = input_scale_ptr[k0 + kk]; - const float s1 = input_scale_ptr[k0 + kk + 1]; + const float s0 = psq[0]; + const float s1 = psq[1]; v00 *= s0; v10 *= s0; v20 *= s0; @@ -1070,10 +1226,6 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A v11 *= s1; v21 *= s1; v31 *= s1; - asm volatile("" - : "+f"(v00), "+f"(v01), "+f"(v10), "+f"(v11)); - asm volatile("" - : "+f"(v20), "+f"(v21), "+f"(v30), "+f"(v31)); } pp[0] = float2int8(v00 * scale0); pp[1] = float2int8(v01 * scale0); @@ -1084,31 +1236,31 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A pp[6] = float2int8(v30 * scale3); pp[7] = float2int8(v31 * scale3); pp += 8; + p0q = p1 + A_hstep; + if (psq) + psq += 2; kk += 2; } for (; kk < max_kk; kk++) { - const int k = k0 + kk; - const float* p0 = (const float*)A + (size_t)k * A_hstep + i0; - float v0 = p0[0]; - float v1 = p0[1]; - float v2 = p0[2]; - float v3 = p0[3]; - if (input_scale_ptr) + float v0 = p0q[0]; + float v1 = p0q[1]; + float v2 = p0q[2]; + float v3 = p0q[3]; + if (psq) { - const float s = input_scale_ptr[k]; + const float s = *psq++; v0 *= s; v1 *= s; v2 *= s; v3 *= s; - asm volatile("" - : "+f"(v0), "+f"(v1), "+f"(v2), "+f"(v3)); } pp[0] = float2int8(v0 * scale0); pp[1] = float2int8(v1 * scale1); pp[2] = float2int8(v2 * scale2); pp[3] = float2int8(v3 * scale3); pp += 4; + p0q += A_hstep; } } } @@ -1120,51 +1272,51 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A { const int k0 = g * block_size; const int max_kk = std::min(K - k0, block_size); + const float* p0g = A_data + (size_t)k0 * A_hstep + i0; + const float* sg = input_scale_ptr ? input_scale_ptr + k0 : 0; float absmax0 = 0.f; float absmax1 = 0.f; + const float* p0a = p0g; + const float* psa = sg; for (int kk = 0; kk < max_kk; kk++) { - const int k = k0 + kk; - const float* p0 = (const float*)A + (size_t)k * A_hstep + i0; - float v0 = p0[0]; - float v1 = p0[1]; - if (input_scale_ptr) + float v0 = p0a[0]; + float v1 = p0a[1]; + if (psa) { - const float s = input_scale_ptr[k]; + const float s = *psa++; v0 *= s; v1 *= s; } absmax0 = std::max(absmax0, fabsf(v0)); absmax1 = std::max(absmax1, fabsf(v1)); + p0a += A_hstep; } - volatile double scale0_fp64 = absmax0 == 0.f ? 1.0 : 127.0 / (double)absmax0; - volatile double scale1_fp64 = absmax1 == 0.f ? 1.0 : 127.0 / (double)absmax1; - const float scale0 = (float)scale0_fp64; - const float scale1 = (float)scale1_fp64; + const float scale0 = absmax0 == 0.f ? 1.f : 127.f / absmax0; + const float scale1 = absmax1 == 0.f ? 1.f : 127.f / absmax1; pd[0] = absmax0 / 127.f; pd[1] = absmax1 / 127.f; pd += 2; + const float* p0q = p0g; + const float* psq = sg; int kk = 0; for (; kk + 3 < max_kk; kk += 4) { for (int r = 0; r < 4; r++) { - const int k = k0 + kk + r; - const float* p0 = (const float*)A + (size_t)k * A_hstep + i0; - float v0 = p0[0]; - float v1 = p0[1]; - if (input_scale_ptr) + float v0 = p0q[0]; + float v1 = p0q[1]; + if (psq) { - const float s = input_scale_ptr[k]; + const float s = *psq++; v0 *= s; v1 *= s; - asm volatile("" - : "+f"(v0), "+f"(v1)); } pp[r] = float2int8(v0 * scale0); pp[4 + r] = float2int8(v1 * scale1); + p0q += A_hstep; } pp += 8; } @@ -1172,41 +1324,35 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A { for (int r = 0; r < 2; r++) { - const int k = k0 + kk + r; - const float* p0 = (const float*)A + (size_t)k * A_hstep + i0; - float v0 = p0[0]; - float v1 = p0[1]; - if (input_scale_ptr) + float v0 = p0q[0]; + float v1 = p0q[1]; + if (psq) { - const float s = input_scale_ptr[k]; + const float s = *psq++; v0 *= s; v1 *= s; - asm volatile("" - : "+f"(v0), "+f"(v1)); } pp[r] = float2int8(v0 * scale0); pp[2 + r] = float2int8(v1 * scale1); + p0q += A_hstep; } pp += 4; kk += 2; } for (; kk < max_kk; kk++) { - const int k = k0 + kk; - const float* p0 = (const float*)A + (size_t)k * A_hstep + i0; - float v0 = p0[0]; - float v1 = p0[1]; - if (input_scale_ptr) + float v0 = p0q[0]; + float v1 = p0q[1]; + if (psq) { - const float s = input_scale_ptr[k]; + const float s = *psq++; v0 *= s; v1 *= s; - asm volatile("" - : "+f"(v0), "+f"(v1)); } pp[0] = float2int8(v0 * scale0); pp[1] = float2int8(v1 * scale1); pp += 2; + p0q += A_hstep; } } } @@ -1218,34 +1364,35 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A { const int k0 = g * block_size; const int max_kk = std::min(K - k0, block_size); + const float* p0g = A_data + (size_t)k0 * A_hstep + i0; + const float* sg = input_scale_ptr ? input_scale_ptr + k0 : 0; float absmax0 = 0.f; + const float* p0a = p0g; + const float* psa = sg; for (int kk = 0; kk < max_kk; kk++) { - const int k = k0 + kk; - float v0 = ((const float*)A)[(size_t)k * A_hstep + i0]; - if (input_scale_ptr) - v0 *= input_scale_ptr[k]; + float v0 = *p0a; + if (psa) + v0 *= *psa++; absmax0 = std::max(absmax0, fabsf(v0)); + p0a += A_hstep; } - volatile double scale0_fp64 = absmax0 == 0.f ? 1.0 : 127.0 / (double)absmax0; - const float scale0 = (float)scale0_fp64; + const float scale0 = absmax0 == 0.f ? 1.f : 127.f / absmax0; *pd++ = absmax0 / 127.f; + const float* p0q = p0g; + const float* psq = sg; int kk = 0; for (; kk + 3 < max_kk; kk += 4) { for (int r = 0; r < 4; r++) { - const int k = k0 + kk + r; - float v0 = ((const float*)A)[(size_t)k * A_hstep + i0]; - if (input_scale_ptr) - { - v0 *= input_scale_ptr[k]; - asm volatile("" - : "+f"(v0)); - } + float v0 = *p0q; + if (psq) + v0 *= *psq++; pp[r] = float2int8(v0 * scale0); + p0q += A_hstep; } pp += 4; } @@ -1253,54 +1400,48 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A { for (int r = 0; r < 2; r++) { - const int k = k0 + kk + r; - float v0 = ((const float*)A)[(size_t)k * A_hstep + i0]; - if (input_scale_ptr) - { - v0 *= input_scale_ptr[k]; - asm volatile("" - : "+f"(v0)); - } + float v0 = *p0q; + if (psq) + v0 *= *psq++; pp[r] = float2int8(v0 * scale0); + p0q += A_hstep; } pp += 2; kk += 2; } for (; kk < max_kk; kk++) { - const int k = k0 + kk; - float v0 = ((const float*)A)[(size_t)k * A_hstep + i0]; - if (input_scale_ptr) - { - v0 *= input_scale_ptr[k]; - asm volatile("" - : "+f"(v0)); - } + float v0 = *p0q; + if (psq) + v0 *= *psq++; *pp++ = float2int8(v0 * scale0); + p0q += A_hstep; } } } } -static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int block_size) +static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int k0, int max_kk0, int B_hstep, int block_size) { #if NCNN_RUNTIME_CPU && NCNN_MMI && !__mips_msa && !__mips_loongson_mmi if (ncnn::cpu_support_loongson_mmi()) { - gemm_transB_packed_tile_wq_int8_loongson_mmi(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, block_size); + gemm_transB_packed_tile_wq_int8_loongson_mmi(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, k0, max_kk0, B_hstep, block_size); return; } #endif const signed char* pAT = AT_tile; - const int A_hstep = AT_tile.w; + const int A_hstep = max_kk0; const float* pAT_descales = AT_descales_tile; - const int A_descales_hstep = AT_descales_tile.w; + const int A_descales_hstep = (max_kk0 + block_size - 1) / block_size; const signed char* pBT = BT_tile; const float* pBT_descales = BT_descales_tile; float* outptr = topT_tile; - const int K = AT_tile.w; - const int num_blocks = (K + block_size - 1) / block_size; + const int K = max_kk0; + const int num_blocks = (B_hstep + block_size - 1) / block_size; + const int block_start = k0 / block_size; + const int tile_blocks = (max_kk0 + block_size - 1) / block_size; int ii = 0; #if __mips_msa @@ -1313,6 +1454,8 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de int jj = 0; for (; jj + 3 < max_jj; jj += 4) { + pB += (size_t)4 * k0; + pBD += (size_t)4 * block_start; const signed char* pA = pAT; const float* pAD = pAT_descales; v4f32 _fsum0 = (v4f32)__msa_fill_w(0); @@ -1341,17 +1484,17 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de { __builtin_prefetch(pA + 64); __builtin_prefetch(pB + 64); - const v16i8 _pA0 = __msa_ld_b(pA, 0); - const v16i8 _pA0r = (v16i8)__msa_shf_w((v4i32)_pA0, _MSA_SHUFFLE(1, 0, 3, 2)); - const v16i8 _pB = __msa_ld_b(pB, 0); - const v16i8 _pBr = (v16i8)__msa_shf_w((v4i32)_pB, _MSA_SHUFFLE(0, 3, 2, 1)); + v16i8 _pA0 = __msa_ld_b(pA, 0); + v16i8 _pA0r = (v16i8)__msa_shf_w((v4i32)_pA0, _MSA_SHUFFLE(1, 0, 3, 2)); + v16i8 _pB = __msa_ld_b(pB, 0); + v16i8 _pBr = (v16i8)__msa_shf_w((v4i32)_pB, _MSA_SHUFFLE(0, 3, 2, 1)); _sum0 = __msa_dpadd_s_w(_sum0, __msa_dotp_s_h(_pA0, _pB), _one); _sum1 = __msa_dpadd_s_w(_sum1, __msa_dotp_s_h(_pA0, _pBr), _one); _sum2 = __msa_dpadd_s_w(_sum2, __msa_dotp_s_h(_pA0r, _pB), _one); _sum3 = __msa_dpadd_s_w(_sum3, __msa_dotp_s_h(_pA0r, _pBr), _one); - const v16i8 _pA1 = __msa_ld_b(pA + 16, 0); - const v16i8 _pA1r = (v16i8)__msa_shf_w((v4i32)_pA1, _MSA_SHUFFLE(1, 0, 3, 2)); + v16i8 _pA1 = __msa_ld_b(pA + 16, 0); + v16i8 _pA1r = (v16i8)__msa_shf_w((v4i32)_pA1, _MSA_SHUFFLE(1, 0, 3, 2)); _sum4 = __msa_dpadd_s_w(_sum4, __msa_dotp_s_h(_pA1, _pB), _one); _sum5 = __msa_dpadd_s_w(_sum5, __msa_dotp_s_h(_pA1, _pBr), _one); _sum6 = __msa_dpadd_s_w(_sum6, __msa_dotp_s_h(_pA1r, _pB), _one); @@ -1375,8 +1518,8 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de if (kk + 1 < max_kk) { - const v16i8 _pB = (v16i8)__msa_fill_d_ptr(pB); - const v8i16 _pA = (v8i16)__msa_ld_b(pA, 0); + v16i8 _pB = (v16i8)__msa_fill_d_ptr(pB); + v8i16 _pA = (v8i16)__msa_ld_b(pA, 0); v8i16 _s = __msa_dotp_s_h((v16i8)__msa_splati_h(_pA, 0), _pB); _sum0 = __msa_addv_w(_sum0, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); _s = __msa_dotp_s_h((v16i8)__msa_splati_h(_pA, 1), _pB); @@ -1399,8 +1542,8 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } if (kk < max_kk) { - const v16i8 _pA8 = (v16i8)__msa_fill_d_ptr(pA); - const v8i16 _pA = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pA8, 0), _pA8); + v16i8 _pA8 = (v16i8)__msa_fill_d_ptr(pA); + v8i16 _pA = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pA8, 0), _pA8); v8i16 _pB = (v8i16)__msa_fill_w(*(const int*)pB); _pB = (v8i16)__msa_ilvr_b(__msa_clti_s_b((v16i8)_pB, 0), (v16i8)_pB); v8i16 _s = __msa_mulv_h(__msa_splati_h(_pA, 0), _pB); @@ -1423,23 +1566,42 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB += 4; } - const v4f32 _descaleB = (v4f32)__msa_ld_w(pBD, 0); - const v4f32 _descaleA0 = (v4f32)__msa_ld_w(pAD, 0); - const v4f32 _descaleA1 = (v4f32)__msa_ld_w(pAD + 4, 0); - _fsum0 = __msa_fadd_w(_fsum0, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum0), __msa_fmul_w(_descaleB, (v4f32)__msa_splati_w((v4i32)_descaleA0, 0)))); - _fsum1 = __msa_fadd_w(_fsum1, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum1), __msa_fmul_w(_descaleB, (v4f32)__msa_splati_w((v4i32)_descaleA0, 1)))); - _fsum2 = __msa_fadd_w(_fsum2, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum2), __msa_fmul_w(_descaleB, (v4f32)__msa_splati_w((v4i32)_descaleA0, 2)))); - _fsum3 = __msa_fadd_w(_fsum3, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum3), __msa_fmul_w(_descaleB, (v4f32)__msa_splati_w((v4i32)_descaleA0, 3)))); - _fsum4 = __msa_fadd_w(_fsum4, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum4), __msa_fmul_w(_descaleB, (v4f32)__msa_splati_w((v4i32)_descaleA1, 0)))); - _fsum5 = __msa_fadd_w(_fsum5, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum5), __msa_fmul_w(_descaleB, (v4f32)__msa_splati_w((v4i32)_descaleA1, 1)))); - _fsum6 = __msa_fadd_w(_fsum6, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum6), __msa_fmul_w(_descaleB, (v4f32)__msa_splati_w((v4i32)_descaleA1, 2)))); - _fsum7 = __msa_fadd_w(_fsum7, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum7), __msa_fmul_w(_descaleB, (v4f32)__msa_splati_w((v4i32)_descaleA1, 3)))); + v4f32 _descaleB = (v4f32)__msa_ld_w(pBD, 0); + v4f32 _descaleA0 = (v4f32)__msa_ld_w(pAD, 0); + v4f32 _descaleA1 = (v4f32)__msa_ld_w(pAD + 4, 0); + v4f32 _scale = __msa_fmul_w(_descaleB, (v4f32)__msa_splati_w((v4i32)_descaleA0, 0)); + _fsum0 = __ncnn_msa_fmadd_w(_fsum0, (v4f32)__msa_ffint_s_w(_sum0), _scale); + _scale = __msa_fmul_w(_descaleB, (v4f32)__msa_splati_w((v4i32)_descaleA0, 1)); + _fsum1 = __ncnn_msa_fmadd_w(_fsum1, (v4f32)__msa_ffint_s_w(_sum1), _scale); + _scale = __msa_fmul_w(_descaleB, (v4f32)__msa_splati_w((v4i32)_descaleA0, 2)); + _fsum2 = __ncnn_msa_fmadd_w(_fsum2, (v4f32)__msa_ffint_s_w(_sum2), _scale); + _scale = __msa_fmul_w(_descaleB, (v4f32)__msa_splati_w((v4i32)_descaleA0, 3)); + _fsum3 = __ncnn_msa_fmadd_w(_fsum3, (v4f32)__msa_ffint_s_w(_sum3), _scale); + _scale = __msa_fmul_w(_descaleB, (v4f32)__msa_splati_w((v4i32)_descaleA1, 0)); + _fsum4 = __ncnn_msa_fmadd_w(_fsum4, (v4f32)__msa_ffint_s_w(_sum4), _scale); + _scale = __msa_fmul_w(_descaleB, (v4f32)__msa_splati_w((v4i32)_descaleA1, 1)); + _fsum5 = __ncnn_msa_fmadd_w(_fsum5, (v4f32)__msa_ffint_s_w(_sum5), _scale); + _scale = __msa_fmul_w(_descaleB, (v4f32)__msa_splati_w((v4i32)_descaleA1, 2)); + _fsum6 = __ncnn_msa_fmadd_w(_fsum6, (v4f32)__msa_ffint_s_w(_sum6), _scale); + _scale = __msa_fmul_w(_descaleB, (v4f32)__msa_splati_w((v4i32)_descaleA1, 3)); + _fsum7 = __ncnn_msa_fmadd_w(_fsum7, (v4f32)__msa_ffint_s_w(_sum7), _scale); pAD += 8; pBD += 4; } transpose4x4_ps(_fsum0, _fsum1, _fsum2, _fsum3); transpose4x4_ps(_fsum4, _fsum5, _fsum6, _fsum7); + if (k0 != 0) + { + _fsum0 = __msa_fadd_w(_fsum0, (v4f32)__msa_ld_w(outptr, 0)); + _fsum4 = __msa_fadd_w(_fsum4, (v4f32)__msa_ld_w(outptr + 4, 0)); + _fsum1 = __msa_fadd_w(_fsum1, (v4f32)__msa_ld_w(outptr + 8, 0)); + _fsum5 = __msa_fadd_w(_fsum5, (v4f32)__msa_ld_w(outptr + 12, 0)); + _fsum2 = __msa_fadd_w(_fsum2, (v4f32)__msa_ld_w(outptr + 16, 0)); + _fsum6 = __msa_fadd_w(_fsum6, (v4f32)__msa_ld_w(outptr + 20, 0)); + _fsum3 = __msa_fadd_w(_fsum3, (v4f32)__msa_ld_w(outptr + 24, 0)); + _fsum7 = __msa_fadd_w(_fsum7, (v4f32)__msa_ld_w(outptr + 28, 0)); + } __msa_st_w((v4i32)_fsum0, outptr, 0); __msa_st_w((v4i32)_fsum4, outptr + 4, 0); __msa_st_w((v4i32)_fsum1, outptr + 8, 0); @@ -1449,9 +1611,13 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de __msa_st_w((v4i32)_fsum3, outptr + 24, 0); __msa_st_w((v4i32)_fsum7, outptr + 28, 0); outptr += 32; + pB += (size_t)4 * (B_hstep - k0 - max_kk0); + pBD += (size_t)4 * (num_blocks - block_start - tile_blocks); } for (; jj + 1 < max_jj; jj += 2) { + pB += (size_t)2 * k0; + pBD += (size_t)2 * block_start; const signed char* pA = pAT; const float* pAD = pAT_descales; v4f32 _fsum0 = (v4f32)__msa_fill_w(0); @@ -1469,10 +1635,12 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de int kk = 0; for (; kk + 3 < max_kk; kk += 4) { - const v16i8 _pA0 = __msa_ld_b(pA, 0); - const v16i8 _pA1 = __msa_ld_b(pA + 16, 0); - const v16i8 _pB = (v16i8)__msa_fill_d_ptr(pB); - const v16i8 _pBr = (v16i8)__msa_shf_w((v4i32)_pB, _MSA_SHUFFLE(0, 3, 2, 1)); + __builtin_prefetch(pA + 64); + __builtin_prefetch(pB + 16); + v16i8 _pA0 = __msa_ld_b(pA, 0); + v16i8 _pA1 = __msa_ld_b(pA + 16, 0); + v16i8 _pB = (v16i8)__msa_fill_d_ptr(pB); + v16i8 _pBr = (v16i8)__msa_shf_w((v4i32)_pB, _MSA_SHUFFLE(0, 3, 2, 1)); _sum0 = __msa_dpadd_s_w(_sum0, __msa_dotp_s_h(_pA0, _pB), _one); _sum1 = __msa_dpadd_s_w(_sum1, __msa_dotp_s_h(_pA0, _pBr), _one); _sum2 = __msa_dpadd_s_w(_sum2, __msa_dotp_s_h(_pA1, _pB), _one); @@ -1481,27 +1649,27 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB += 8; } - const v4i32 _sum0e = __msa_shf_w(_sum0, _MSA_SHUFFLE(3, 1, 2, 0)); - const v4i32 _sum0o = __msa_shf_w(_sum0, _MSA_SHUFFLE(2, 0, 3, 1)); - const v4i32 _sum1e = __msa_shf_w(_sum1, _MSA_SHUFFLE(3, 1, 2, 0)); - const v4i32 _sum1o = __msa_shf_w(_sum1, _MSA_SHUFFLE(2, 0, 3, 1)); + v4i32 _sum0e = __msa_shf_w(_sum0, _MSA_SHUFFLE(3, 1, 2, 0)); + v4i32 _sum0o = __msa_shf_w(_sum0, _MSA_SHUFFLE(2, 0, 3, 1)); + v4i32 _sum1e = __msa_shf_w(_sum1, _MSA_SHUFFLE(3, 1, 2, 0)); + v4i32 _sum1o = __msa_shf_w(_sum1, _MSA_SHUFFLE(2, 0, 3, 1)); v4i32 _sum0x = (v4i32)__msa_ilvr_w(_sum1o, _sum0e); v4i32 _sum1x = (v4i32)__msa_ilvr_w(_sum0o, _sum1e); - const v4i32 _sum2e = __msa_shf_w(_sum2, _MSA_SHUFFLE(3, 1, 2, 0)); - const v4i32 _sum2o = __msa_shf_w(_sum2, _MSA_SHUFFLE(2, 0, 3, 1)); - const v4i32 _sum3e = __msa_shf_w(_sum3, _MSA_SHUFFLE(3, 1, 2, 0)); - const v4i32 _sum3o = __msa_shf_w(_sum3, _MSA_SHUFFLE(2, 0, 3, 1)); + v4i32 _sum2e = __msa_shf_w(_sum2, _MSA_SHUFFLE(3, 1, 2, 0)); + v4i32 _sum2o = __msa_shf_w(_sum2, _MSA_SHUFFLE(2, 0, 3, 1)); + v4i32 _sum3e = __msa_shf_w(_sum3, _MSA_SHUFFLE(3, 1, 2, 0)); + v4i32 _sum3o = __msa_shf_w(_sum3, _MSA_SHUFFLE(2, 0, 3, 1)); v4i32 _sum2x = (v4i32)__msa_ilvr_w(_sum3o, _sum2e); v4i32 _sum3x = (v4i32)__msa_ilvr_w(_sum2o, _sum3e); if (kk + 1 < max_kk) { - const v16i8 _pA = __msa_ld_b(pA, 0); - const v8i16 _pB = (v8i16)__msa_fill_w(*(const int*)pB); - const v8i16 _s0 = __msa_dotp_s_h(_pA, (v16i8)__msa_splati_h(_pB, 0)); - const v8i16 _s1 = __msa_dotp_s_h(_pA, (v16i8)__msa_splati_h(_pB, 1)); - const v8i16 _sign0 = __msa_clti_s_h(_s0, 0); - const v8i16 _sign1 = __msa_clti_s_h(_s1, 0); + v16i8 _pA = __msa_ld_b(pA, 0); + v8i16 _pB = (v8i16)__msa_fill_w(*(const int*)pB); + v8i16 _s0 = __msa_dotp_s_h(_pA, (v16i8)__msa_splati_h(_pB, 0)); + v8i16 _s1 = __msa_dotp_s_h(_pA, (v16i8)__msa_splati_h(_pB, 1)); + v8i16 _sign0 = __msa_clti_s_h(_s0, 0); + v8i16 _sign1 = __msa_clti_s_h(_s1, 0); _sum0x = __msa_addv_w(_sum0x, (v4i32)__msa_ilvr_h(_sign0, _s0)); _sum1x = __msa_addv_w(_sum1x, (v4i32)__msa_ilvr_h(_sign1, _s1)); _sum2x = __msa_addv_w(_sum2x, (v4i32)__msa_ilvl_h(_sign0, _s0)); @@ -1512,15 +1680,15 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } if (kk < max_kk) { - const v16i8 _pA8 = (v16i8)__msa_fill_d_ptr(pA); - const v8i16 _pA = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pA8, 0), _pA8); - const v16i8 _pB8 = (v16i8)__msa_fill_h(*(const short*)pB); - const v8i16 _pB0 = (v8i16)__msa_ilvr_b(__msa_clti_s_b(__msa_splati_b(_pB8, 0), 0), __msa_splati_b(_pB8, 0)); - const v8i16 _pB1 = (v8i16)__msa_ilvr_b(__msa_clti_s_b(__msa_splati_b(_pB8, 1), 0), __msa_splati_b(_pB8, 1)); - const v8i16 _s0 = __msa_mulv_h(_pA, _pB0); - const v8i16 _s1 = __msa_mulv_h(_pA, _pB1); - const v8i16 _sign0 = __msa_clti_s_h(_s0, 0); - const v8i16 _sign1 = __msa_clti_s_h(_s1, 0); + v16i8 _pA8 = (v16i8)__msa_fill_d_ptr(pA); + v8i16 _pA = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pA8, 0), _pA8); + v16i8 _pB8 = (v16i8)__msa_fill_h(*(const short*)pB); + v8i16 _pB0 = (v8i16)__msa_ilvr_b(__msa_clti_s_b(__msa_splati_b(_pB8, 0), 0), __msa_splati_b(_pB8, 0)); + v8i16 _pB1 = (v8i16)__msa_ilvr_b(__msa_clti_s_b(__msa_splati_b(_pB8, 1), 0), __msa_splati_b(_pB8, 1)); + v8i16 _s0 = __msa_mulv_h(_pA, _pB0); + v8i16 _s1 = __msa_mulv_h(_pA, _pB1); + v8i16 _sign0 = __msa_clti_s_h(_s0, 0); + v8i16 _sign1 = __msa_clti_s_h(_s1, 0); _sum0x = __msa_addv_w(_sum0x, (v4i32)__msa_ilvr_h(_sign0, _s0)); _sum1x = __msa_addv_w(_sum1x, (v4i32)__msa_ilvr_h(_sign1, _s1)); _sum2x = __msa_addv_w(_sum2x, (v4i32)__msa_ilvl_h(_sign0, _s0)); @@ -1529,24 +1697,39 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB += 2; } - const v4f32 _descaleA0 = (v4f32)__msa_ld_w(pAD, 0); - const v4f32 _descaleA1 = (v4f32)__msa_ld_w(pAD + 4, 0); - _fsum0 = __msa_fadd_w(_fsum0, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum0x), __msa_fmul_w(_descaleA0, __msa_fill_w_f32(pBD[0])))); - _fsum1 = __msa_fadd_w(_fsum1, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum2x), __msa_fmul_w(_descaleA1, __msa_fill_w_f32(pBD[0])))); - _fsum2 = __msa_fadd_w(_fsum2, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum1x), __msa_fmul_w(_descaleA0, __msa_fill_w_f32(pBD[1])))); - _fsum3 = __msa_fadd_w(_fsum3, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum3x), __msa_fmul_w(_descaleA1, __msa_fill_w_f32(pBD[1])))); + v4f32 _descaleA0 = (v4f32)__msa_ld_w(pAD, 0); + v4f32 _descaleA1 = (v4f32)__msa_ld_w(pAD + 4, 0); + v4f32 _scale = __msa_fmul_w(_descaleA0, __msa_fill_w_f32(pBD[0])); + _fsum0 = __ncnn_msa_fmadd_w(_fsum0, (v4f32)__msa_ffint_s_w(_sum0x), _scale); + _scale = __msa_fmul_w(_descaleA1, __msa_fill_w_f32(pBD[0])); + _fsum1 = __ncnn_msa_fmadd_w(_fsum1, (v4f32)__msa_ffint_s_w(_sum2x), _scale); + _scale = __msa_fmul_w(_descaleA0, __msa_fill_w_f32(pBD[1])); + _fsum2 = __ncnn_msa_fmadd_w(_fsum2, (v4f32)__msa_ffint_s_w(_sum1x), _scale); + _scale = __msa_fmul_w(_descaleA1, __msa_fill_w_f32(pBD[1])); + _fsum3 = __ncnn_msa_fmadd_w(_fsum3, (v4f32)__msa_ffint_s_w(_sum3x), _scale); pAD += 8; pBD += 2; } + if (k0 != 0) + { + _fsum0 = __msa_fadd_w(_fsum0, (v4f32)__msa_ld_w(outptr, 0)); + _fsum1 = __msa_fadd_w(_fsum1, (v4f32)__msa_ld_w(outptr + 4, 0)); + _fsum2 = __msa_fadd_w(_fsum2, (v4f32)__msa_ld_w(outptr + 8, 0)); + _fsum3 = __msa_fadd_w(_fsum3, (v4f32)__msa_ld_w(outptr + 12, 0)); + } __msa_st_w((v4i32)_fsum0, outptr, 0); __msa_st_w((v4i32)_fsum1, outptr + 4, 0); __msa_st_w((v4i32)_fsum2, outptr + 8, 0); __msa_st_w((v4i32)_fsum3, outptr + 12, 0); outptr += 16; + pB += (size_t)2 * (B_hstep - k0 - max_kk0); + pBD += (size_t)2 * (num_blocks - block_start - tile_blocks); } for (; jj < max_jj; jj++) { + pB += k0; + pBD += block_start; const signed char* pA = pAT; const float* pAD = pAT_descales; v4f32 _fsum0 = (v4f32)__msa_fill_w(0); @@ -1560,9 +1743,11 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de int kk = 0; for (; kk + 3 < max_kk; kk += 4) { - const v16i8 _pA0 = __msa_ld_b(pA, 0); - const v16i8 _pA1 = __msa_ld_b(pA + 16, 0); - const v16i8 _pB = (v16i8)__msa_fill_w(*(const int*)pB); + __builtin_prefetch(pA + 64); + __builtin_prefetch(pB + 16); + v16i8 _pA0 = __msa_ld_b(pA, 0); + v16i8 _pA1 = __msa_ld_b(pA + 16, 0); + v16i8 _pB = (v16i8)__msa_fill_w(*(const int*)pB); _sum0 = __msa_dpadd_s_w(_sum0, __msa_dotp_s_h(_pA0, _pB), _one); _sum1 = __msa_dpadd_s_w(_sum1, __msa_dotp_s_h(_pA1, _pB), _one); pA += 32; @@ -1570,10 +1755,10 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } if (kk + 1 < max_kk) { - const v16i8 _pA = __msa_ld_b(pA, 0); - const v16i8 _pB = (v16i8)__msa_fill_h(*(const short*)pB); - const v8i16 _s = __msa_dotp_s_h(_pA, _pB); - const v8i16 _sign = __msa_clti_s_h(_s, 0); + v16i8 _pA = __msa_ld_b(pA, 0); + v16i8 _pB = (v16i8)__msa_fill_h(*(const short*)pB); + v8i16 _s = __msa_dotp_s_h(_pA, _pB); + v8i16 _sign = __msa_clti_s_h(_s, 0); _sum0 = __msa_addv_w(_sum0, (v4i32)__msa_ilvr_h(_sign, _s)); _sum1 = __msa_addv_w(_sum1, (v4i32)__msa_ilvl_h(_sign, _s)); pA += 16; @@ -1582,27 +1767,36 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } if (kk < max_kk) { - const v16i8 _pA8 = (v16i8)__msa_fill_d_ptr(pA); - const v8i16 _pA = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pA8, 0), _pA8); - const v8i16 _s = __msa_mulv_h(_pA, __msa_fill_h(pB[0])); - const v8i16 _sign = __msa_clti_s_h(_s, 0); + v16i8 _pA8 = (v16i8)__msa_fill_d_ptr(pA); + v8i16 _pA = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pA8, 0), _pA8); + v8i16 _s = __msa_mulv_h(_pA, __msa_fill_h(pB[0])); + v8i16 _sign = __msa_clti_s_h(_s, 0); _sum0 = __msa_addv_w(_sum0, (v4i32)__msa_ilvr_h(_sign, _s)); _sum1 = __msa_addv_w(_sum1, (v4i32)__msa_ilvl_h(_sign, _s)); pA += 8; pB++; } - const v4f32 _descaleA0 = (v4f32)__msa_ld_w(pAD, 0); - const v4f32 _descaleA1 = (v4f32)__msa_ld_w(pAD + 4, 0); - _fsum0 = __msa_fadd_w(_fsum0, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum0), __msa_fmul_w(_descaleA0, __msa_fill_w_f32(pBD[0])))); - _fsum1 = __msa_fadd_w(_fsum1, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum1), __msa_fmul_w(_descaleA1, __msa_fill_w_f32(pBD[0])))); + v4f32 _descaleA0 = (v4f32)__msa_ld_w(pAD, 0); + v4f32 _descaleA1 = (v4f32)__msa_ld_w(pAD + 4, 0); + v4f32 _scale = __msa_fmul_w(_descaleA0, __msa_fill_w_f32(pBD[0])); + _fsum0 = __ncnn_msa_fmadd_w(_fsum0, (v4f32)__msa_ffint_s_w(_sum0), _scale); + _scale = __msa_fmul_w(_descaleA1, __msa_fill_w_f32(pBD[0])); + _fsum1 = __ncnn_msa_fmadd_w(_fsum1, (v4f32)__msa_ffint_s_w(_sum1), _scale); pAD += 8; pBD++; } + if (k0 != 0) + { + _fsum0 = __msa_fadd_w(_fsum0, (v4f32)__msa_ld_w(outptr, 0)); + _fsum1 = __msa_fadd_w(_fsum1, (v4f32)__msa_ld_w(outptr + 4, 0)); + } __msa_st_w((v4i32)_fsum0, outptr, 0); __msa_st_w((v4i32)_fsum1, outptr + 4, 0); outptr += 8; + pB += B_hstep - k0 - max_kk0; + pBD += num_blocks - block_start - tile_blocks; } pAT += (size_t)8 * A_hstep; @@ -1619,10 +1813,10 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de { const signed char* pA = pAT; const float* pAD = pAT_descales; - const signed char* pB0 = pB; - const signed char* pB1 = pB + (size_t)4 * K; - const float* pBD0 = pBD; - const float* pBD1 = pBD + (size_t)4 * num_blocks; + const signed char* pB0 = pB + (size_t)4 * k0; + const signed char* pB1 = pB + (size_t)4 * B_hstep + (size_t)4 * k0; + const float* pBD0 = pBD + (size_t)4 * block_start; + const float* pBD1 = pBD + (size_t)4 * num_blocks + (size_t)4 * block_start; v4f32 _fsum0 = (v4f32)__msa_fill_w(0); v4f32 _fsum1 = (v4f32)__msa_fill_w(0); v4f32 _fsum2 = (v4f32)__msa_fill_w(0); @@ -1646,12 +1840,15 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de int kk = 0; for (; kk + 3 < max_kk; kk += 4) { - const v16i8 _pA = __msa_ld_b(pA, 0); - const v16i8 _pAr = (v16i8)__msa_shf_w((v4i32)_pA, _MSA_SHUFFLE(1, 0, 3, 2)); - const v16i8 _pB0 = __msa_ld_b(pB0, 0); - const v16i8 _pB1 = __msa_ld_b(pB1, 0); - const v16i8 _pB0r = (v16i8)__msa_shf_w((v4i32)_pB0, _MSA_SHUFFLE(0, 3, 2, 1)); - const v16i8 _pB1r = (v16i8)__msa_shf_w((v4i32)_pB1, _MSA_SHUFFLE(0, 3, 2, 1)); + __builtin_prefetch(pA + 64); + __builtin_prefetch(pB0 + 64); + __builtin_prefetch(pB1 + 64); + v16i8 _pA = __msa_ld_b(pA, 0); + v16i8 _pAr = (v16i8)__msa_shf_w((v4i32)_pA, _MSA_SHUFFLE(1, 0, 3, 2)); + v16i8 _pB0 = __msa_ld_b(pB0, 0); + v16i8 _pB1 = __msa_ld_b(pB1, 0); + v16i8 _pB0r = (v16i8)__msa_shf_w((v4i32)_pB0, _MSA_SHUFFLE(0, 3, 2, 1)); + v16i8 _pB1r = (v16i8)__msa_shf_w((v4i32)_pB1, _MSA_SHUFFLE(0, 3, 2, 1)); _sum0 = __msa_dpadd_s_w(_sum0, __msa_dotp_s_h(_pA, _pB0), _one); _sum1 = __msa_dpadd_s_w(_sum1, __msa_dotp_s_h(_pA, _pB0r), _one); _sum2 = __msa_dpadd_s_w(_sum2, __msa_dotp_s_h(_pAr, _pB0), _one); @@ -1665,49 +1862,6 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB1 += 16; } - const signed char* pA2 = pA; - const signed char* pB02 = pB0; - const signed char* pB12 = pB1; - const bool has_k2 = kk + 1 < max_kk; - if (kk + 1 < max_kk) - { - pA += 8; - pB0 += 8; - pB1 += 8; - kk += 2; - } - for (; kk < max_kk; kk++) - { - v8i16 _pA = (v8i16)__msa_fill_w(*(const int*)pA); - _pA = (v8i16)__msa_ilvr_b(__msa_clti_s_b((v16i8)_pA, 0), (v16i8)_pA); - const v8i16 _pAr = __msa_shf_h(_pA, _MSA_SHUFFLE(1, 0, 3, 2)); - v8i16 _pB0 = (v8i16)__msa_fill_w(*(const int*)pB0); - v8i16 _pB1 = (v8i16)__msa_fill_w(*(const int*)pB1); - _pB0 = (v8i16)__msa_ilvr_b(__msa_clti_s_b((v16i8)_pB0, 0), (v16i8)_pB0); - _pB1 = (v8i16)__msa_ilvr_b(__msa_clti_s_b((v16i8)_pB1, 0), (v16i8)_pB1); - const v8i16 _pB0r = __msa_shf_h(_pB0, _MSA_SHUFFLE(0, 3, 2, 1)); - const v8i16 _pB1r = __msa_shf_h(_pB1, _MSA_SHUFFLE(0, 3, 2, 1)); - const v8i16 _s0 = __msa_mulv_h(_pA, _pB0); - const v8i16 _s1 = __msa_mulv_h(_pA, _pB0r); - const v8i16 _s2 = __msa_mulv_h(_pAr, _pB0); - const v8i16 _s3 = __msa_mulv_h(_pAr, _pB0r); - const v8i16 _s4 = __msa_mulv_h(_pA, _pB1); - const v8i16 _s5 = __msa_mulv_h(_pA, _pB1r); - const v8i16 _s6 = __msa_mulv_h(_pAr, _pB1); - const v8i16 _s7 = __msa_mulv_h(_pAr, _pB1r); - _sum0 = __msa_addv_w(_sum0, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s0, 0), _s0)); - _sum1 = __msa_addv_w(_sum1, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s1, 0), _s1)); - _sum2 = __msa_addv_w(_sum2, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s2, 0), _s2)); - _sum3 = __msa_addv_w(_sum3, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s3, 0), _s3)); - _sum4 = __msa_addv_w(_sum4, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s4, 0), _s4)); - _sum5 = __msa_addv_w(_sum5, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s5, 0), _s5)); - _sum6 = __msa_addv_w(_sum6, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s6, 0), _s6)); - _sum7 = __msa_addv_w(_sum7, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s7, 0), _s7)); - pA += 4; - pB0 += 4; - pB1 += 4; - } - _sum2 = __msa_shf_w(_sum2, _MSA_SHUFFLE(1, 0, 3, 2)); _sum3 = __msa_shf_w(_sum3, _MSA_SHUFFLE(1, 0, 3, 2)); transpose4x4_epi32(_sum0, _sum1, _sum2, _sum3); @@ -1720,11 +1874,12 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de _sum5 = __msa_shf_w(_sum5, _MSA_SHUFFLE(2, 1, 0, 3)); _sum6 = __msa_shf_w(_sum6, _MSA_SHUFFLE(1, 0, 3, 2)); _sum7 = __msa_shf_w(_sum7, _MSA_SHUFFLE(0, 3, 2, 1)); - if (has_k2) + + if (kk + 1 < max_kk) { - const v8i16 _pA = (v8i16)__msa_fill_d_ptr(pA2); - const v16i8 _pB0 = (v16i8)__msa_fill_d_ptr(pB02); - const v16i8 _pB1 = (v16i8)__msa_fill_d_ptr(pB12); + v8i16 _pA = (v8i16)__msa_fill_d_ptr(pA); + v16i8 _pB0 = (v16i8)__msa_fill_d_ptr(pB0); + v16i8 _pB1 = (v16i8)__msa_fill_d_ptr(pB1); v8i16 _s = __msa_dotp_s_h((v16i8)__msa_splati_h(_pA, 0), _pB0); _sum0 = __msa_addv_w(_sum0, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); _s = __msa_dotp_s_h((v16i8)__msa_splati_h(_pA, 1), _pB0); @@ -1741,18 +1896,58 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de _sum6 = __msa_addv_w(_sum6, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); _s = __msa_dotp_s_h((v16i8)__msa_splati_h(_pA, 3), _pB1); _sum7 = __msa_addv_w(_sum7, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + pA += 8; + pB0 += 8; + pB1 += 8; + kk += 2; + } + for (; kk < max_kk; kk++) + { + v16i8 _pA8 = (v16i8)__msa_fill_w(*(const int*)pA); + v8i16 _pA = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pA8, 0), _pA8); + v16i8 _pB08 = (v16i8)__msa_fill_w(*(const int*)pB0); + v16i8 _pB18 = (v16i8)__msa_fill_w(*(const int*)pB1); + v8i16 _pB0 = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pB08, 0), _pB08); + v8i16 _pB1 = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pB18, 0), _pB18); + v8i16 _s = __msa_mulv_h(__msa_splati_h(_pA, 0), _pB0); + _sum0 = __msa_addv_w(_sum0, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + _s = __msa_mulv_h(__msa_splati_h(_pA, 1), _pB0); + _sum1 = __msa_addv_w(_sum1, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + _s = __msa_mulv_h(__msa_splati_h(_pA, 2), _pB0); + _sum2 = __msa_addv_w(_sum2, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + _s = __msa_mulv_h(__msa_splati_h(_pA, 3), _pB0); + _sum3 = __msa_addv_w(_sum3, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + _s = __msa_mulv_h(__msa_splati_h(_pA, 0), _pB1); + _sum4 = __msa_addv_w(_sum4, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + _s = __msa_mulv_h(__msa_splati_h(_pA, 1), _pB1); + _sum5 = __msa_addv_w(_sum5, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + _s = __msa_mulv_h(__msa_splati_h(_pA, 2), _pB1); + _sum6 = __msa_addv_w(_sum6, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + _s = __msa_mulv_h(__msa_splati_h(_pA, 3), _pB1); + _sum7 = __msa_addv_w(_sum7, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + pA += 4; + pB0 += 4; + pB1 += 4; } - const v4f32 _descaleB0 = (v4f32)__msa_ld_w(pBD0, 0); - const v4f32 _descaleB1 = (v4f32)__msa_ld_w(pBD1, 0); - _fsum0 = __msa_fadd_w(_fsum0, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum0), __msa_fmul_w(_descaleB0, __msa_fill_w_f32(pAD[0])))); - _fsum1 = __msa_fadd_w(_fsum1, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum1), __msa_fmul_w(_descaleB0, __msa_fill_w_f32(pAD[1])))); - _fsum2 = __msa_fadd_w(_fsum2, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum2), __msa_fmul_w(_descaleB0, __msa_fill_w_f32(pAD[2])))); - _fsum3 = __msa_fadd_w(_fsum3, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum3), __msa_fmul_w(_descaleB0, __msa_fill_w_f32(pAD[3])))); - _fsum4 = __msa_fadd_w(_fsum4, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum4), __msa_fmul_w(_descaleB1, __msa_fill_w_f32(pAD[0])))); - _fsum5 = __msa_fadd_w(_fsum5, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum5), __msa_fmul_w(_descaleB1, __msa_fill_w_f32(pAD[1])))); - _fsum6 = __msa_fadd_w(_fsum6, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum6), __msa_fmul_w(_descaleB1, __msa_fill_w_f32(pAD[2])))); - _fsum7 = __msa_fadd_w(_fsum7, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum7), __msa_fmul_w(_descaleB1, __msa_fill_w_f32(pAD[3])))); + v4f32 _descaleB0 = (v4f32)__msa_ld_w(pBD0, 0); + v4f32 _descaleB1 = (v4f32)__msa_ld_w(pBD1, 0); + v4f32 _scale = __msa_fmul_w(_descaleB0, __msa_fill_w_f32(pAD[0])); + _fsum0 = __ncnn_msa_fmadd_w(_fsum0, (v4f32)__msa_ffint_s_w(_sum0), _scale); + _scale = __msa_fmul_w(_descaleB0, __msa_fill_w_f32(pAD[1])); + _fsum1 = __ncnn_msa_fmadd_w(_fsum1, (v4f32)__msa_ffint_s_w(_sum1), _scale); + _scale = __msa_fmul_w(_descaleB0, __msa_fill_w_f32(pAD[2])); + _fsum2 = __ncnn_msa_fmadd_w(_fsum2, (v4f32)__msa_ffint_s_w(_sum2), _scale); + _scale = __msa_fmul_w(_descaleB0, __msa_fill_w_f32(pAD[3])); + _fsum3 = __ncnn_msa_fmadd_w(_fsum3, (v4f32)__msa_ffint_s_w(_sum3), _scale); + _scale = __msa_fmul_w(_descaleB1, __msa_fill_w_f32(pAD[0])); + _fsum4 = __ncnn_msa_fmadd_w(_fsum4, (v4f32)__msa_ffint_s_w(_sum4), _scale); + _scale = __msa_fmul_w(_descaleB1, __msa_fill_w_f32(pAD[1])); + _fsum5 = __ncnn_msa_fmadd_w(_fsum5, (v4f32)__msa_ffint_s_w(_sum5), _scale); + _scale = __msa_fmul_w(_descaleB1, __msa_fill_w_f32(pAD[2])); + _fsum6 = __ncnn_msa_fmadd_w(_fsum6, (v4f32)__msa_ffint_s_w(_sum6), _scale); + _scale = __msa_fmul_w(_descaleB1, __msa_fill_w_f32(pAD[3])); + _fsum7 = __ncnn_msa_fmadd_w(_fsum7, (v4f32)__msa_ffint_s_w(_sum7), _scale); pAD += 4; pBD0 += 4; pBD1 += 4; @@ -1760,6 +1955,17 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de transpose4x4_ps(_fsum0, _fsum1, _fsum2, _fsum3); transpose4x4_ps(_fsum4, _fsum5, _fsum6, _fsum7); + if (k0 != 0) + { + _fsum0 = __msa_fadd_w(_fsum0, (v4f32)__msa_ld_w(outptr, 0)); + _fsum1 = __msa_fadd_w(_fsum1, (v4f32)__msa_ld_w(outptr + 4, 0)); + _fsum2 = __msa_fadd_w(_fsum2, (v4f32)__msa_ld_w(outptr + 8, 0)); + _fsum3 = __msa_fadd_w(_fsum3, (v4f32)__msa_ld_w(outptr + 12, 0)); + _fsum4 = __msa_fadd_w(_fsum4, (v4f32)__msa_ld_w(outptr + 16, 0)); + _fsum5 = __msa_fadd_w(_fsum5, (v4f32)__msa_ld_w(outptr + 20, 0)); + _fsum6 = __msa_fadd_w(_fsum6, (v4f32)__msa_ld_w(outptr + 24, 0)); + _fsum7 = __msa_fadd_w(_fsum7, (v4f32)__msa_ld_w(outptr + 28, 0)); + } __msa_st_w((v4i32)_fsum0, outptr, 0); __msa_st_w((v4i32)_fsum1, outptr + 4, 0); __msa_st_w((v4i32)_fsum2, outptr + 8, 0); @@ -1769,11 +1975,13 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de __msa_st_w((v4i32)_fsum6, outptr + 24, 0); __msa_st_w((v4i32)_fsum7, outptr + 28, 0); outptr += 32; - pB = pB1; - pBD = pBD1; + pB += (size_t)8 * B_hstep; + pBD += (size_t)8 * num_blocks; } for (; jj + 3 < max_jj; jj += 4) { + pB += (size_t)4 * k0; + pBD += (size_t)4 * block_start; const signed char* pA = pAT; const float* pAD = pAT_descales; v4f32 _fsum0 = (v4f32)__msa_fill_w(0); @@ -1791,10 +1999,12 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de int kk = 0; for (; kk + 3 < max_kk; kk += 4) { - const v16i8 _pA = __msa_ld_b(pA, 0); - const v16i8 _pAr = (v16i8)__msa_shf_w((v4i32)_pA, _MSA_SHUFFLE(1, 0, 3, 2)); - const v16i8 _pB0 = __msa_ld_b(pB, 0); - const v16i8 _pB0r = (v16i8)__msa_shf_w((v4i32)_pB0, _MSA_SHUFFLE(0, 3, 2, 1)); + __builtin_prefetch(pA + 64); + __builtin_prefetch(pB + 32); + v16i8 _pA = __msa_ld_b(pA, 0); + v16i8 _pAr = (v16i8)__msa_shf_w((v4i32)_pA, _MSA_SHUFFLE(1, 0, 3, 2)); + v16i8 _pB0 = __msa_ld_b(pB, 0); + v16i8 _pB0r = (v16i8)__msa_shf_w((v4i32)_pB0, _MSA_SHUFFLE(0, 3, 2, 1)); _sum0 = __msa_dpadd_s_w(_sum0, __msa_dotp_s_h(_pA, _pB0), _one); _sum1 = __msa_dpadd_s_w(_sum1, __msa_dotp_s_h(_pA, _pB0r), _one); _sum2 = __msa_dpadd_s_w(_sum2, __msa_dotp_s_h(_pAr, _pB0), _one); @@ -1808,12 +2018,12 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de v8i16 _sum2_3 = __msa_fill_h(0); if (kk + 1 < max_kk) { - const v16i8 _pB = (v16i8)__msa_fill_d_ptr(pB); - const v8i16 _pA = (v8i16)__msa_fill_d_ptr(pA); - const v16i8 _pA0 = (v16i8)__msa_splati_h(_pA, 0); - const v16i8 _pA1 = (v16i8)__msa_splati_h(_pA, 1); - const v16i8 _pA2 = (v16i8)__msa_splati_h(_pA, 2); - const v16i8 _pA3 = (v16i8)__msa_splati_h(_pA, 3); + v16i8 _pB = (v16i8)__msa_fill_d_ptr(pB); + v8i16 _pA = (v8i16)__msa_fill_d_ptr(pA); + v16i8 _pA0 = (v16i8)__msa_splati_h(_pA, 0); + v16i8 _pA1 = (v16i8)__msa_splati_h(_pA, 1); + v16i8 _pA2 = (v16i8)__msa_splati_h(_pA, 2); + v16i8 _pA3 = (v16i8)__msa_splati_h(_pA, 3); _sum2_0 = __msa_dotp_s_h(_pA0, _pB); _sum2_1 = __msa_dotp_s_h(_pA1, _pB); _sum2_2 = __msa_dotp_s_h(_pA2, _pB); @@ -1826,14 +2036,14 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de { v8i16 _pA = (v8i16)__msa_fill_w(*(const int*)pA); _pA = (v8i16)__msa_ilvr_b(__msa_clti_s_b((v16i8)_pA, 0), (v16i8)_pA); - const v8i16 _pAr = __msa_shf_h(_pA, _MSA_SHUFFLE(1, 0, 3, 2)); + v8i16 _pAr = __msa_shf_h(_pA, _MSA_SHUFFLE(1, 0, 3, 2)); v8i16 _pB0 = (v8i16)__msa_fill_w(*(const int*)pB); _pB0 = (v8i16)__msa_ilvr_b(__msa_clti_s_b((v16i8)_pB0, 0), (v16i8)_pB0); - const v8i16 _pB0r = __msa_shf_h(_pB0, _MSA_SHUFFLE(0, 3, 2, 1)); - const v8i16 _s0 = __msa_mulv_h(_pA, _pB0); - const v8i16 _s1 = __msa_mulv_h(_pA, _pB0r); - const v8i16 _s2 = __msa_mulv_h(_pAr, _pB0); - const v8i16 _s3 = __msa_mulv_h(_pAr, _pB0r); + v8i16 _pB0r = __msa_shf_h(_pB0, _MSA_SHUFFLE(0, 3, 2, 1)); + v8i16 _s0 = __msa_mulv_h(_pA, _pB0); + v8i16 _s1 = __msa_mulv_h(_pA, _pB0r); + v8i16 _s2 = __msa_mulv_h(_pAr, _pB0); + v8i16 _s3 = __msa_mulv_h(_pAr, _pB0r); _sum0 = __msa_addv_w(_sum0, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s0, 0), _s0)); _sum1 = __msa_addv_w(_sum1, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s1, 0), _s1)); _sum2 = __msa_addv_w(_sum2, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s2, 0), _s2)); @@ -1851,23 +2061,38 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de _sum1 = __msa_addv_w(_sum1, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_sum2_1, 0), _sum2_1)); _sum2 = __msa_addv_w(_sum2, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_sum2_2, 0), _sum2_2)); _sum3 = __msa_addv_w(_sum3, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_sum2_3, 0), _sum2_3)); - const v4f32 _descaleB = (v4f32)__msa_ld_w(pBD, 0); - _fsum0 = __msa_fadd_w(_fsum0, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum0), __msa_fmul_w(_descaleB, __msa_fill_w_f32(pAD[0])))); - _fsum1 = __msa_fadd_w(_fsum1, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum1), __msa_fmul_w(_descaleB, __msa_fill_w_f32(pAD[1])))); - _fsum2 = __msa_fadd_w(_fsum2, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum2), __msa_fmul_w(_descaleB, __msa_fill_w_f32(pAD[2])))); - _fsum3 = __msa_fadd_w(_fsum3, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum3), __msa_fmul_w(_descaleB, __msa_fill_w_f32(pAD[3])))); + v4f32 _descaleB = (v4f32)__msa_ld_w(pBD, 0); + v4f32 _scale = __msa_fmul_w(_descaleB, __msa_fill_w_f32(pAD[0])); + _fsum0 = __ncnn_msa_fmadd_w(_fsum0, (v4f32)__msa_ffint_s_w(_sum0), _scale); + _scale = __msa_fmul_w(_descaleB, __msa_fill_w_f32(pAD[1])); + _fsum1 = __ncnn_msa_fmadd_w(_fsum1, (v4f32)__msa_ffint_s_w(_sum1), _scale); + _scale = __msa_fmul_w(_descaleB, __msa_fill_w_f32(pAD[2])); + _fsum2 = __ncnn_msa_fmadd_w(_fsum2, (v4f32)__msa_ffint_s_w(_sum2), _scale); + _scale = __msa_fmul_w(_descaleB, __msa_fill_w_f32(pAD[3])); + _fsum3 = __ncnn_msa_fmadd_w(_fsum3, (v4f32)__msa_ffint_s_w(_sum3), _scale); pAD += 4; pBD += 4; } transpose4x4_ps(_fsum0, _fsum1, _fsum2, _fsum3); + if (k0 != 0) + { + _fsum0 = __msa_fadd_w(_fsum0, (v4f32)__msa_ld_w(outptr, 0)); + _fsum1 = __msa_fadd_w(_fsum1, (v4f32)__msa_ld_w(outptr + 4, 0)); + _fsum2 = __msa_fadd_w(_fsum2, (v4f32)__msa_ld_w(outptr + 8, 0)); + _fsum3 = __msa_fadd_w(_fsum3, (v4f32)__msa_ld_w(outptr + 12, 0)); + } __msa_st_w((v4i32)_fsum0, outptr, 0); __msa_st_w((v4i32)_fsum1, outptr + 4, 0); __msa_st_w((v4i32)_fsum2, outptr + 8, 0); __msa_st_w((v4i32)_fsum3, outptr + 12, 0); outptr += 16; + pB += (size_t)4 * (B_hstep - k0 - max_kk0); + pBD += (size_t)4 * (num_blocks - block_start - tile_blocks); } for (; jj + 1 < max_jj; jj += 2) { + pB += (size_t)2 * k0; + pBD += (size_t)2 * block_start; const signed char* pA = pAT; const float* pAD = pAT_descales; v4f32 _fsum0 = (v4f32)__msa_fill_w(0); @@ -1880,9 +2105,11 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de int kk = 0; for (; kk + 3 < max_kk; kk += 4) { - const v16i8 _pA = __msa_ld_b(pA, 0); - const v16i8 _pB0 = (v16i8)__msa_fill_d_ptr(pB); - const v16i8 _pB0r = (v16i8)__msa_shf_w((v4i32)_pB0, _MSA_SHUFFLE(0, 3, 2, 1)); + __builtin_prefetch(pA + 32); + __builtin_prefetch(pB + 16); + v16i8 _pA = __msa_ld_b(pA, 0); + v16i8 _pB0 = (v16i8)__msa_fill_d_ptr(pB); + v16i8 _pB0r = (v16i8)__msa_shf_w((v4i32)_pB0, _MSA_SHUFFLE(0, 3, 2, 1)); _sum0 = __msa_dpadd_s_w(_sum0, __msa_dotp_s_h(_pA, _pB0), _one); _sum1 = __msa_dpadd_s_w(_sum1, __msa_dotp_s_h(_pA, _pB0r), _one); pA += 16; @@ -1892,10 +2119,10 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de v8i16 _sum2_1 = __msa_fill_h(0); if (kk + 1 < max_kk) { - const v16i8 _pA = (v16i8)__msa_fill_d_ptr(pA); - const v8i16 _pB = (v8i16)__msa_fill_w(*(const int*)pB); - const v16i8 _pB0 = (v16i8)__msa_splati_h(_pB, 0); - const v16i8 _pB1 = (v16i8)__msa_splati_h(_pB, 1); + v16i8 _pA = (v16i8)__msa_fill_d_ptr(pA); + v8i16 _pB = (v8i16)__msa_fill_w(*(const int*)pB); + v16i8 _pB0 = (v16i8)__msa_splati_h(_pB, 0); + v16i8 _pB1 = (v16i8)__msa_splati_h(_pB, 1); _sum2_0 = __msa_dotp_s_h(_pA, _pB0); _sum2_1 = __msa_dotp_s_h(_pA, _pB1); pA += 8; @@ -1906,36 +2133,47 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de { v8i16 _pA = (v8i16)__msa_fill_w(*(const int*)pA); _pA = (v8i16)__msa_ilvr_b(__msa_clti_s_b((v16i8)_pA, 0), (v16i8)_pA); - const v16i8 _pB8 = (v16i8)__msa_fill_h(*(const short*)pB); - const v8i16 _pB0 = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pB8, 0), _pB8); - const v8i16 _pB0r = __msa_shf_h(_pB0, _MSA_SHUFFLE(0, 3, 2, 1)); - const v8i16 _s0 = __msa_mulv_h(_pA, _pB0); - const v8i16 _s1 = __msa_mulv_h(_pA, _pB0r); + v16i8 _pB8 = (v16i8)__msa_fill_h(*(const short*)pB); + v8i16 _pB0 = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pB8, 0), _pB8); + v8i16 _pB0r = __msa_shf_h(_pB0, _MSA_SHUFFLE(0, 3, 2, 1)); + v8i16 _s0 = __msa_mulv_h(_pA, _pB0); + v8i16 _s1 = __msa_mulv_h(_pA, _pB0r); _sum0 = __msa_addv_w(_sum0, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s0, 0), _s0)); _sum1 = __msa_addv_w(_sum1, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s1, 0), _s1)); pA += 4; pB += 2; } - const v4i32 _sum0e = __msa_shf_w(_sum0, _MSA_SHUFFLE(3, 1, 2, 0)); - const v4i32 _sum0o = __msa_shf_w(_sum0, _MSA_SHUFFLE(2, 0, 3, 1)); - const v4i32 _sum1e = __msa_shf_w(_sum1, _MSA_SHUFFLE(3, 1, 2, 0)); - const v4i32 _sum1o = __msa_shf_w(_sum1, _MSA_SHUFFLE(2, 0, 3, 1)); + v4i32 _sum0e = __msa_shf_w(_sum0, _MSA_SHUFFLE(3, 1, 2, 0)); + v4i32 _sum0o = __msa_shf_w(_sum0, _MSA_SHUFFLE(2, 0, 3, 1)); + v4i32 _sum1e = __msa_shf_w(_sum1, _MSA_SHUFFLE(3, 1, 2, 0)); + v4i32 _sum1o = __msa_shf_w(_sum1, _MSA_SHUFFLE(2, 0, 3, 1)); v4i32 _sum0x = (v4i32)__msa_ilvr_w(_sum1o, _sum0e); v4i32 _sum1x = (v4i32)__msa_ilvr_w(_sum0o, _sum1e); _sum0x = __msa_addv_w(_sum0x, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_sum2_0, 0), _sum2_0)); _sum1x = __msa_addv_w(_sum1x, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_sum2_1, 0), _sum2_1)); - const v4f32 _descaleA = (v4f32)__msa_ld_w(pAD, 0); - _fsum0 = __msa_fadd_w(_fsum0, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum0x), __msa_fmul_w(_descaleA, __msa_fill_w_f32(pBD[0])))); - _fsum1 = __msa_fadd_w(_fsum1, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum1x), __msa_fmul_w(_descaleA, __msa_fill_w_f32(pBD[1])))); + v4f32 _descaleA = (v4f32)__msa_ld_w(pAD, 0); + v4f32 _scale = __msa_fmul_w(_descaleA, __msa_fill_w_f32(pBD[0])); + _fsum0 = __ncnn_msa_fmadd_w(_fsum0, (v4f32)__msa_ffint_s_w(_sum0x), _scale); + _scale = __msa_fmul_w(_descaleA, __msa_fill_w_f32(pBD[1])); + _fsum1 = __ncnn_msa_fmadd_w(_fsum1, (v4f32)__msa_ffint_s_w(_sum1x), _scale); pAD += 4; pBD += 2; } + if (k0 != 0) + { + _fsum0 = __msa_fadd_w(_fsum0, (v4f32)__msa_ld_w(outptr, 0)); + _fsum1 = __msa_fadd_w(_fsum1, (v4f32)__msa_ld_w(outptr + 4, 0)); + } __msa_st_w((v4i32)_fsum0, outptr, 0); __msa_st_w((v4i32)_fsum1, outptr + 4, 0); outptr += 8; + pB += (size_t)2 * (B_hstep - k0 - max_kk0); + pBD += (size_t)2 * (num_blocks - block_start - tile_blocks); } for (; jj < max_jj; jj++) { + pB += k0; + pBD += block_start; const signed char* pA = pAT; const float* pAD = pAT_descales; v4f32 _fsum0 = (v4f32)__msa_fill_w(0); @@ -1946,8 +2184,10 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de int kk = 0; for (; kk + 3 < max_kk; kk += 4) { - const v16i8 _pA = __msa_ld_b(pA, 0); - const v16i8 _pB0 = (v16i8)__msa_fill_w(*(const int*)pB); + __builtin_prefetch(pA + 32); + __builtin_prefetch(pB + 16); + v16i8 _pA = __msa_ld_b(pA, 0); + v16i8 _pB0 = (v16i8)__msa_fill_w(*(const int*)pB); _sum0 = __msa_dpadd_s_w(_sum0, __msa_dotp_s_h(_pA, _pB0), _one); pA += 16; pB += 4; @@ -1955,8 +2195,8 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de v8i16 _sum2_0 = __msa_fill_h(0); if (kk + 1 < max_kk) { - const v16i8 _pA = (v16i8)__msa_fill_d_ptr(pA); - const v16i8 _pB = (v16i8)__msa_fill_h(*(const short*)pB); + v16i8 _pA = (v16i8)__msa_fill_d_ptr(pA); + v16i8 _pB = (v16i8)__msa_fill_h(*(const short*)pB); _sum2_0 = __msa_dotp_s_h(_pA, _pB); pA += 8; pB += 2; @@ -1966,20 +2206,25 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de { v8i16 _pA = (v8i16)__msa_fill_w(*(const int*)pA); _pA = (v8i16)__msa_ilvr_b(__msa_clti_s_b((v16i8)_pA, 0), (v16i8)_pA); - const v8i16 _pB0 = __msa_fill_h(pB[0]); - const v8i16 _s0 = __msa_mulv_h(_pA, _pB0); + v8i16 _pB0 = __msa_fill_h(pB[0]); + v8i16 _s0 = __msa_mulv_h(_pA, _pB0); _sum0 = __msa_addv_w(_sum0, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s0, 0), _s0)); pA += 4; pB++; } _sum0 = __msa_addv_w(_sum0, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_sum2_0, 0), _sum2_0)); - const v4f32 _descaleA = (v4f32)__msa_ld_w(pAD, 0); - _fsum0 = __msa_fadd_w(_fsum0, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum0), __msa_fmul_w(_descaleA, __msa_fill_w_f32(pBD[0])))); + v4f32 _descaleA = (v4f32)__msa_ld_w(pAD, 0); + v4f32 _scale = __msa_fmul_w(_descaleA, __msa_fill_w_f32(pBD[0])); + _fsum0 = __ncnn_msa_fmadd_w(_fsum0, (v4f32)__msa_ffint_s_w(_sum0), _scale); pAD += 4; pBD++; } + if (k0 != 0) + _fsum0 = __msa_fadd_w(_fsum0, (v4f32)__msa_ld_w(outptr, 0)); __msa_st_w((v4i32)_fsum0, outptr, 0); outptr += 4; + pB += B_hstep - k0 - max_kk0; + pBD += num_blocks - block_start - tile_blocks; } pAT += (size_t)4 * A_hstep; @@ -1998,10 +2243,10 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de { const signed char* pA = pAT; const float* pAD = pAT_descales; - const signed char* pB0 = pB; - const signed char* pB1 = pB + (size_t)4 * K; - const float* pBD0 = pBD; - const float* pBD1 = pBD + (size_t)4 * num_blocks; + const signed char* pB0 = pB + (size_t)4 * k0; + const signed char* pB1 = pB + (size_t)4 * B_hstep + (size_t)4 * k0; + const float* pBD0 = pBD + (size_t)4 * block_start; + const float* pBD1 = pBD + (size_t)4 * num_blocks + (size_t)4 * block_start; v4f32 _fsum0 = (v4f32)__msa_fill_w(0); v4f32 _fsum1 = (v4f32)__msa_fill_w(0); v4f32 _fsum2 = (v4f32)__msa_fill_w(0); @@ -2016,13 +2261,16 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de int kk = 0; for (; kk + 3 < max_kk; kk += 4) { - const v16i8 _pA = (v16i8)__msa_fill_d_ptr(pA); - const v16i8 _pB0 = __msa_ld_b(pB0, 0); - const v16i8 _pB00 = (v16i8)__msa_ilvr_w((v4i32)_pB0, (v4i32)_pB0); - const v16i8 _pB01 = (v16i8)__msa_ilvl_w((v4i32)_pB0, (v4i32)_pB0); - const v16i8 _pB1 = __msa_ld_b(pB1, 0); - const v16i8 _pB10 = (v16i8)__msa_ilvr_w((v4i32)_pB1, (v4i32)_pB1); - const v16i8 _pB11 = (v16i8)__msa_ilvl_w((v4i32)_pB1, (v4i32)_pB1); + __builtin_prefetch(pA + 32); + __builtin_prefetch(pB0 + 64); + __builtin_prefetch(pB1 + 64); + v16i8 _pA = (v16i8)__msa_fill_d_ptr(pA); + v16i8 _pB0 = __msa_ld_b(pB0, 0); + v16i8 _pB00 = (v16i8)__msa_ilvr_w((v4i32)_pB0, (v4i32)_pB0); + v16i8 _pB01 = (v16i8)__msa_ilvl_w((v4i32)_pB0, (v4i32)_pB0); + v16i8 _pB1 = __msa_ld_b(pB1, 0); + v16i8 _pB10 = (v16i8)__msa_ilvr_w((v4i32)_pB1, (v4i32)_pB1); + v16i8 _pB11 = (v16i8)__msa_ilvl_w((v4i32)_pB1, (v4i32)_pB1); _sum0 = __msa_dpadd_s_w(_sum0, __msa_dotp_s_h(_pA, _pB00), _one); _sum1 = __msa_dpadd_s_w(_sum1, __msa_dotp_s_h(_pA, _pB01), _one); _sum2 = __msa_dpadd_s_w(_sum2, __msa_dotp_s_h(_pA, _pB10), _one); @@ -2033,19 +2281,19 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } if (kk + 1 < max_kk) { - const v8i16 _pA = (v8i16)__msa_fill_w(*(const int*)pA); - const v16i8 _pA0 = (v16i8)__msa_splati_h(_pA, 0); - const v16i8 _pA1 = (v16i8)__msa_splati_h(_pA, 1); - const v16i8 _pB0 = (v16i8)__msa_fill_d_ptr(pB0); - const v16i8 _pB1 = (v16i8)__msa_fill_d_ptr(pB1); - const v8i16 _s00 = __msa_dotp_s_h(_pA0, _pB0); - const v8i16 _s01 = __msa_dotp_s_h(_pA1, _pB0); - const v8i16 _s10 = __msa_dotp_s_h(_pA0, _pB1); - const v8i16 _s11 = __msa_dotp_s_h(_pA1, _pB1); - const v8i16 _s0 = (v8i16)__msa_ilvr_h(_s01, _s00); - const v8i16 _s2 = (v8i16)__msa_ilvr_h(_s11, _s10); - const v8i16 _sign0 = __msa_clti_s_h(_s0, 0); - const v8i16 _sign2 = __msa_clti_s_h(_s2, 0); + v8i16 _pA = (v8i16)__msa_fill_w(*(const int*)pA); + v16i8 _pA0 = (v16i8)__msa_splati_h(_pA, 0); + v16i8 _pA1 = (v16i8)__msa_splati_h(_pA, 1); + v16i8 _pB0 = (v16i8)__msa_fill_d_ptr(pB0); + v16i8 _pB1 = (v16i8)__msa_fill_d_ptr(pB1); + v8i16 _s00 = __msa_dotp_s_h(_pA0, _pB0); + v8i16 _s01 = __msa_dotp_s_h(_pA1, _pB0); + v8i16 _s10 = __msa_dotp_s_h(_pA0, _pB1); + v8i16 _s11 = __msa_dotp_s_h(_pA1, _pB1); + v8i16 _s0 = (v8i16)__msa_ilvr_h(_s01, _s00); + v8i16 _s2 = (v8i16)__msa_ilvr_h(_s11, _s10); + v8i16 _sign0 = __msa_clti_s_h(_s0, 0); + v8i16 _sign2 = __msa_clti_s_h(_s2, 0); _sum0 = __msa_addv_w(_sum0, (v4i32)__msa_ilvr_h(_sign0, _s0)); _sum1 = __msa_addv_w(_sum1, (v4i32)__msa_ilvl_h(_sign0, _s0)); _sum2 = __msa_addv_w(_sum2, (v4i32)__msa_ilvr_h(_sign2, _s2)); @@ -2057,23 +2305,23 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } if (kk < max_kk) { - const v16i8 _pA8 = (v16i8)__msa_fill_h(*(const short*)pA); - const v16i8 _pA0b = __msa_splati_b(_pA8, 0); - const v16i8 _pA1b = __msa_splati_b(_pA8, 1); - const v8i16 _pA0 = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pA0b, 0), _pA0b); - const v8i16 _pA1 = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pA1b, 0), _pA1b); - const v16i8 _pB08 = (v16i8)__msa_fill_w(*(const int*)pB0); - const v16i8 _pB18 = (v16i8)__msa_fill_w(*(const int*)pB1); - const v8i16 _pB0 = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pB08, 0), _pB08); - const v8i16 _pB1 = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pB18, 0), _pB18); - const v8i16 _s00 = __msa_mulv_h(_pA0, _pB0); - const v8i16 _s01 = __msa_mulv_h(_pA1, _pB0); - const v8i16 _s10 = __msa_mulv_h(_pA0, _pB1); - const v8i16 _s11 = __msa_mulv_h(_pA1, _pB1); - const v8i16 _s0 = (v8i16)__msa_ilvr_h(_s01, _s00); - const v8i16 _s2 = (v8i16)__msa_ilvr_h(_s11, _s10); - const v8i16 _sign0 = __msa_clti_s_h(_s0, 0); - const v8i16 _sign2 = __msa_clti_s_h(_s2, 0); + v16i8 _pA8 = (v16i8)__msa_fill_h(*(const short*)pA); + v16i8 _pA0b = __msa_splati_b(_pA8, 0); + v16i8 _pA1b = __msa_splati_b(_pA8, 1); + v8i16 _pA0 = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pA0b, 0), _pA0b); + v8i16 _pA1 = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pA1b, 0), _pA1b); + v16i8 _pB08 = (v16i8)__msa_fill_w(*(const int*)pB0); + v16i8 _pB18 = (v16i8)__msa_fill_w(*(const int*)pB1); + v8i16 _pB0 = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pB08, 0), _pB08); + v8i16 _pB1 = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pB18, 0), _pB18); + v8i16 _s00 = __msa_mulv_h(_pA0, _pB0); + v8i16 _s01 = __msa_mulv_h(_pA1, _pB0); + v8i16 _s10 = __msa_mulv_h(_pA0, _pB1); + v8i16 _s11 = __msa_mulv_h(_pA1, _pB1); + v8i16 _s0 = (v8i16)__msa_ilvr_h(_s01, _s00); + v8i16 _s2 = (v8i16)__msa_ilvr_h(_s11, _s10); + v8i16 _sign0 = __msa_clti_s_h(_s0, 0); + v8i16 _sign2 = __msa_clti_s_h(_s2, 0); _sum0 = __msa_addv_w(_sum0, (v4i32)__msa_ilvr_h(_sign0, _s0)); _sum1 = __msa_addv_w(_sum1, (v4i32)__msa_ilvl_h(_sign0, _s0)); _sum2 = __msa_addv_w(_sum2, (v4i32)__msa_ilvr_h(_sign2, _s2)); @@ -2082,39 +2330,52 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB0 += 4; pB1 += 4; } - const v4f32 _descaleA = (v4f32) { + v4f32 _descaleA = (v4f32) { pAD[0], pAD[1], pAD[0], pAD[1] }; - const v4f32 _descaleB0 = (v4f32) { + v4f32 _descaleB0 = (v4f32) { pBD0[0], pBD0[0], pBD0[1], pBD0[1] }; - const v4f32 _descaleB1 = (v4f32) { + v4f32 _descaleB1 = (v4f32) { pBD0[2], pBD0[2], pBD0[3], pBD0[3] }; - const v4f32 _descaleB2 = (v4f32) { + v4f32 _descaleB2 = (v4f32) { pBD1[0], pBD1[0], pBD1[1], pBD1[1] }; - const v4f32 _descaleB3 = (v4f32) { + v4f32 _descaleB3 = (v4f32) { pBD1[2], pBD1[2], pBD1[3], pBD1[3] }; - _fsum0 = __msa_fadd_w(_fsum0, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum0), __msa_fmul_w(_descaleA, _descaleB0))); - _fsum1 = __msa_fadd_w(_fsum1, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum1), __msa_fmul_w(_descaleA, _descaleB1))); - _fsum2 = __msa_fadd_w(_fsum2, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum2), __msa_fmul_w(_descaleA, _descaleB2))); - _fsum3 = __msa_fadd_w(_fsum3, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum3), __msa_fmul_w(_descaleA, _descaleB3))); + v4f32 _scale = __msa_fmul_w(_descaleA, _descaleB0); + _fsum0 = __ncnn_msa_fmadd_w(_fsum0, (v4f32)__msa_ffint_s_w(_sum0), _scale); + _scale = __msa_fmul_w(_descaleA, _descaleB1); + _fsum1 = __ncnn_msa_fmadd_w(_fsum1, (v4f32)__msa_ffint_s_w(_sum1), _scale); + _scale = __msa_fmul_w(_descaleA, _descaleB2); + _fsum2 = __ncnn_msa_fmadd_w(_fsum2, (v4f32)__msa_ffint_s_w(_sum2), _scale); + _scale = __msa_fmul_w(_descaleA, _descaleB3); + _fsum3 = __ncnn_msa_fmadd_w(_fsum3, (v4f32)__msa_ffint_s_w(_sum3), _scale); pAD += 2; pBD0 += 4; pBD1 += 4; } + if (k0 != 0) + { + _fsum0 = __msa_fadd_w(_fsum0, (v4f32)__msa_ld_w(outptr, 0)); + _fsum1 = __msa_fadd_w(_fsum1, (v4f32)__msa_ld_w(outptr + 4, 0)); + _fsum2 = __msa_fadd_w(_fsum2, (v4f32)__msa_ld_w(outptr + 8, 0)); + _fsum3 = __msa_fadd_w(_fsum3, (v4f32)__msa_ld_w(outptr + 12, 0)); + } __msa_st_w((v4i32)_fsum0, outptr, 0); __msa_st_w((v4i32)_fsum1, outptr + 4, 0); __msa_st_w((v4i32)_fsum2, outptr + 8, 0); __msa_st_w((v4i32)_fsum3, outptr + 12, 0); outptr += 16; - pB = pB1; - pBD = pBD1; + pB += (size_t)8 * B_hstep; + pBD += (size_t)8 * num_blocks; } for (; jj + 3 < max_jj; jj += 4) { + pB += (size_t)4 * k0; + pBD += (size_t)4 * block_start; const signed char* pA = pAT; const float* pAD = pAT_descales; v4f32 _fsum0 = (v4f32)__msa_fill_w(0); @@ -2127,10 +2388,12 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de int kk = 0; for (; kk + 3 < max_kk; kk += 4) { - const v16i8 _pA = (v16i8)__msa_fill_d_ptr(pA); - const v16i8 _pB0 = __msa_ld_b(pB, 0); - const v16i8 _pB01 = (v16i8)__msa_ilvr_w((v4i32)_pB0, (v4i32)_pB0); - const v16i8 _pB23 = (v16i8)__msa_ilvl_w((v4i32)_pB0, (v4i32)_pB0); + __builtin_prefetch(pA + 32); + __builtin_prefetch(pB + 32); + v16i8 _pA = (v16i8)__msa_fill_d_ptr(pA); + v16i8 _pB0 = __msa_ld_b(pB, 0); + v16i8 _pB01 = (v16i8)__msa_ilvr_w((v4i32)_pB0, (v4i32)_pB0); + v16i8 _pB23 = (v16i8)__msa_ilvl_w((v4i32)_pB0, (v4i32)_pB0); _sum0 = __msa_dpadd_s_w(_sum0, __msa_dotp_s_h(_pA, _pB01), _one); _sum1 = __msa_dpadd_s_w(_sum1, __msa_dotp_s_h(_pA, _pB23), _one); pA += 8; @@ -2138,14 +2401,14 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } if (kk + 1 < max_kk) { - const v8i16 _pA = (v8i16)__msa_fill_w(*(const int*)pA); - const v16i8 _pA0 = (v16i8)__msa_splati_h(_pA, 0); - const v16i8 _pA1 = (v16i8)__msa_splati_h(_pA, 1); - const v16i8 _pB = (v16i8)__msa_fill_d_ptr(pB); - const v8i16 _s00 = __msa_dotp_s_h(_pA0, _pB); - const v8i16 _s01 = __msa_dotp_s_h(_pA1, _pB); - const v8i16 _s0 = (v8i16)__msa_ilvr_h(_s01, _s00); - const v8i16 _sign = __msa_clti_s_h(_s0, 0); + v8i16 _pA = (v8i16)__msa_fill_w(*(const int*)pA); + v16i8 _pA0 = (v16i8)__msa_splati_h(_pA, 0); + v16i8 _pA1 = (v16i8)__msa_splati_h(_pA, 1); + v16i8 _pB = (v16i8)__msa_fill_d_ptr(pB); + v8i16 _s00 = __msa_dotp_s_h(_pA0, _pB); + v8i16 _s01 = __msa_dotp_s_h(_pA1, _pB); + v8i16 _s0 = (v8i16)__msa_ilvr_h(_s01, _s00); + v8i16 _sign = __msa_clti_s_h(_s0, 0); _sum0 = __msa_addv_w(_sum0, (v4i32)__msa_ilvr_h(_sign, _s0)); _sum1 = __msa_addv_w(_sum1, (v4i32)__msa_ilvl_h(_sign, _s0)); pA += 4; @@ -2154,43 +2417,182 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } if (kk < max_kk) { - const v16i8 _pA8 = (v16i8)__msa_fill_h(*(const short*)pA); - const v16i8 _pA0b = __msa_splati_b(_pA8, 0); - const v16i8 _pA1b = __msa_splati_b(_pA8, 1); - const v8i16 _pA0 = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pA0b, 0), _pA0b); - const v8i16 _pA1 = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pA1b, 0), _pA1b); - const v16i8 _pB8 = (v16i8)__msa_fill_w(*(const int*)pB); - const v8i16 _pB = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pB8, 0), _pB8); - const v8i16 _s00 = __msa_mulv_h(_pA0, _pB); - const v8i16 _s01 = __msa_mulv_h(_pA1, _pB); - const v8i16 _s0 = (v8i16)__msa_ilvr_h(_s01, _s00); - const v8i16 _sign = __msa_clti_s_h(_s0, 0); + v16i8 _pA8 = (v16i8)__msa_fill_h(*(const short*)pA); + v16i8 _pA0b = __msa_splati_b(_pA8, 0); + v16i8 _pA1b = __msa_splati_b(_pA8, 1); + v8i16 _pA0 = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pA0b, 0), _pA0b); + v8i16 _pA1 = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pA1b, 0), _pA1b); + v16i8 _pB8 = (v16i8)__msa_fill_w(*(const int*)pB); + v8i16 _pB = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pB8, 0), _pB8); + v8i16 _s00 = __msa_mulv_h(_pA0, _pB); + v8i16 _s01 = __msa_mulv_h(_pA1, _pB); + v8i16 _s0 = (v8i16)__msa_ilvr_h(_s01, _s00); + v8i16 _sign = __msa_clti_s_h(_s0, 0); _sum0 = __msa_addv_w(_sum0, (v4i32)__msa_ilvr_h(_sign, _s0)); _sum1 = __msa_addv_w(_sum1, (v4i32)__msa_ilvl_h(_sign, _s0)); pA += 2; pB += 4; } - const v4f32 _descaleA = (v4f32) { + v4f32 _descaleA = (v4f32) { pAD[0], pAD[1], pAD[0], pAD[1] }; - const v4f32 _descaleB0 = (v4f32) { + v4f32 _descaleB0 = (v4f32) { pBD[0], pBD[0], pBD[1], pBD[1] }; - const v4f32 _descaleB1 = (v4f32) { + v4f32 _descaleB1 = (v4f32) { pBD[2], pBD[2], pBD[3], pBD[3] }; - _fsum0 = __msa_fadd_w(_fsum0, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum0), __msa_fmul_w(_descaleA, _descaleB0))); - _fsum1 = __msa_fadd_w(_fsum1, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum1), __msa_fmul_w(_descaleA, _descaleB1))); + v4f32 _scale = __msa_fmul_w(_descaleA, _descaleB0); + _fsum0 = __ncnn_msa_fmadd_w(_fsum0, (v4f32)__msa_ffint_s_w(_sum0), _scale); + _scale = __msa_fmul_w(_descaleA, _descaleB1); + _fsum1 = __ncnn_msa_fmadd_w(_fsum1, (v4f32)__msa_ffint_s_w(_sum1), _scale); pAD += 2; pBD += 4; } + if (k0 != 0) + { + _fsum0 = __msa_fadd_w(_fsum0, (v4f32)__msa_ld_w(outptr, 0)); + _fsum1 = __msa_fadd_w(_fsum1, (v4f32)__msa_ld_w(outptr + 4, 0)); + } __msa_st_w((v4i32)_fsum0, outptr, 0); __msa_st_w((v4i32)_fsum1, outptr + 4, 0); outptr += 8; + pB += (size_t)4 * (B_hstep - k0 - max_kk0); + pBD += (size_t)4 * (num_blocks - block_start - tile_blocks); + } + for (; jj + 1 < max_jj; jj += 2) + { + pB += (size_t)2 * k0; + pBD += (size_t)2 * block_start; + const signed char* pA = pAT; + const float* pAD = pAT_descales; + v4f32 _fsum = (v4f32)__msa_fill_w(0); + for (int k = 0; k < K; k += block_size) + { + const int max_kk = std::min(K - k, block_size); + v4i32 _sum = __msa_fill_w(0); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __builtin_prefetch(pA + 32); + __builtin_prefetch(pB + 32); + v16i8 _pA = (v16i8)__msa_fill_d_ptr(pA); + v16i8 _pB = (v16i8)__msa_fill_d_ptr(pB); + v16i8 _pB01 = (v16i8)__msa_ilvr_w((v4i32)_pB, (v4i32)_pB); + _sum = __msa_dpadd_s_w(_sum, __msa_dotp_s_h(_pA, _pB01), _one); + pA += 8; + pB += 8; + } + if (kk + 1 < max_kk) + { + v8i16 _pA = (v8i16)__msa_fill_w(*(const int*)pA); + v16i8 _pA0 = (v16i8)__msa_splati_h(_pA, 0); + v16i8 _pA1 = (v16i8)__msa_splati_h(_pA, 1); + v16i8 _pB = (v16i8)__msa_fill_w(*(const int*)pB); + v8i16 _s0 = __msa_dotp_s_h(_pA0, _pB); + v8i16 _s1 = __msa_dotp_s_h(_pA1, _pB); + v8i16 _s = (v8i16)__msa_ilvr_h(_s1, _s0); + _sum = __msa_addv_w(_sum, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + pA += 4; + pB += 4; + kk += 2; + } + if (kk < max_kk) + { + v16i8 _pA8 = (v16i8)__msa_fill_h(*(const short*)pA); + v16i8 _pA0b = __msa_splati_b(_pA8, 0); + v16i8 _pA1b = __msa_splati_b(_pA8, 1); + v8i16 _pA0 = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pA0b, 0), _pA0b); + v8i16 _pA1 = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pA1b, 0), _pA1b); + v16i8 _pB8 = (v16i8)__msa_fill_h(*(const short*)pB); + v8i16 _pB = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pB8, 0), _pB8); + v8i16 _s0 = __msa_mulv_h(_pA0, _pB); + v8i16 _s1 = __msa_mulv_h(_pA1, _pB); + v8i16 _s = (v8i16)__msa_ilvr_h(_s1, _s0); + _sum = __msa_addv_w(_sum, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + pA += 2; + pB += 2; + } + v4f32 _descaleA = (v4f32) { + pAD[0], pAD[1], pAD[0], pAD[1] + }; + v4f32 _descaleB = (v4f32) { + pBD[0], pBD[0], pBD[1], pBD[1] + }; + v4f32 _scale = __msa_fmul_w(_descaleA, _descaleB); + _fsum = __ncnn_msa_fmadd_w(_fsum, (v4f32)__msa_ffint_s_w(_sum), _scale); + pAD += 2; + pBD += 2; + } + if (k0 != 0) + _fsum = __msa_fadd_w(_fsum, (v4f32)__msa_ld_w(outptr, 0)); + __msa_st_w((v4i32)_fsum, outptr, 0); + outptr += 4; + pB += (size_t)2 * (B_hstep - k0 - max_kk0); + pBD += (size_t)2 * (num_blocks - block_start - tile_blocks); + } + for (; jj < max_jj; jj++) + { + pB += k0; + pBD += block_start; + const signed char* pA = pAT; + const float* pAD = pAT_descales; + v4f32 _fsum = (v4f32)__msa_fill_w(0); + for (int k = 0; k < K; k += block_size) + { + const int max_kk = std::min(K - k, block_size); + v4i32 _sum = __msa_fill_w(0); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __builtin_prefetch(pA + 32); + __builtin_prefetch(pB + 16); + v16i8 _pA = (v16i8)__msa_fill_d_ptr(pA); + v16i8 _pB = (v16i8)__msa_fill_w(*(const int*)pB); + _sum = __msa_dpadd_s_w(_sum, __msa_dotp_s_h(_pA, _pB), _one); + pA += 8; + pB += 4; + } + if (kk + 1 < max_kk) + { + v16i8 _pA = (v16i8)__msa_fill_w(*(const int*)pA); + v16i8 _pB = (v16i8)__msa_fill_h(*(const short*)pB); + v8i16 _s = __msa_dotp_s_h(_pA, _pB); + _sum = __msa_addv_w(_sum, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + pA += 4; + pB += 2; + kk += 2; + } + if (kk < max_kk) + { + v16i8 _pA8 = (v16i8)__msa_fill_h(*(const short*)pA); + v8i16 _pA = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pA8, 0), _pA8); + v8i16 _s = __msa_mulv_h(_pA, __msa_fill_h(pB[0])); + _sum = __msa_addv_w(_sum, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + pA += 2; + pB++; + } + v4f32 _descaleA = (v4f32) { + pAD[0], pAD[1], pAD[0], pAD[1] + }; + v4f32 _scale = __msa_fmul_w(_descaleA, __msa_fill_w_f32(pBD[0])); + _fsum = __ncnn_msa_fmadd_w(_fsum, (v4f32)__msa_ffint_s_w(_sum), _scale); + pAD += 2; + pBD++; + } + if (k0 != 0) + _fsum = __msa_fadd_w(_fsum, (v4f32)__msa_loadl_d(outptr)); + __msa_storel_d((v4i32)_fsum, outptr); + outptr += 2; + pB += B_hstep - k0 - max_kk0; + pBD += num_blocks - block_start - tile_blocks; } #endif // __mips_msa +#if !__mips_msa for (; jj + 1 < max_jj; jj += 2) { + pB += (size_t)2 * k0; + pBD += (size_t)2 * block_start; const signed char* pA = pAT; const float* pAD = pAT_descales; float sum00 = 0.f; @@ -2290,8 +2692,8 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de for (; kk + 3 < max_kk; kk += 4) { __builtin_prefetch(pB + 32); - const int8x8_t _pA = __mmi_pldb_s(pA); - const int8x8_t _pB = __mmi_pldb_s(pB); + int8x8_t _pA = __mmi_pldb_s(pA); + int8x8_t _pB = __mmi_pldb_s(pB); int16x4_t _pA0 = (int16x4_t)__mmi_punpcklbh_s(_pA, _zero); int16x4_t _pA1 = (int16x4_t)__mmi_punpckhbh_s(_pA, _zero); int16x4_t _pB0 = (int16x4_t)__mmi_punpcklbh_s(_pB, _zero); @@ -2353,14 +2755,25 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pBD += 2; } + if (k0 != 0) + { + sum00 += outptr[0]; + sum01 += outptr[1]; + sum10 += outptr[2]; + sum11 += outptr[3]; + } outptr[0] = sum00; outptr[1] = sum01; outptr[2] = sum10; outptr[3] = sum11; outptr += 4; + pB += (size_t)2 * (B_hstep - k0 - max_kk0); + pBD += (size_t)2 * (num_blocks - block_start - tile_blocks); } for (; jj < max_jj; jj++) { + pB += k0; + pBD += block_start; const signed char* pA = pAT; const float* pAD = pAT_descales; float sum0 = 0.f; @@ -2442,8 +2855,8 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de for (; kk + 3 < max_kk; kk += 4) { __builtin_prefetch(pB + 16); - const int8x8_t _pA = __mmi_pldb_s(pA); - const int8x8_t _pB = (int8x8_t)__mmi_pfillw_s(*(const int*)pB); + int8x8_t _pA = __mmi_pldb_s(pA); + int8x8_t _pB = (int8x8_t)__mmi_pfillw_s(*(const int*)pB); int16x4_t _pA0 = (int16x4_t)__mmi_punpcklbh_s(_pA, _zero); int16x4_t _pA1 = (int16x4_t)__mmi_punpckhbh_s(_pA, _zero); int16x4_t _pB0 = (int16x4_t)__mmi_punpcklbh_s(_pB, _zero); @@ -2489,10 +2902,18 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pBD++; } + if (k0 != 0) + { + sum0 += outptr[0]; + sum1 += outptr[1]; + } outptr[0] = sum0; outptr[1] = sum1; outptr += 2; + pB += B_hstep - k0 - max_kk0; + pBD += num_blocks - block_start - tile_blocks; } +#endif // !__mips_msa pAT += (size_t)2 * A_hstep; pAT_descales += (size_t)2 * A_descales_hstep; @@ -2509,10 +2930,10 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de { const signed char* pA = pAT; const float* pAD = pAT_descales; - const signed char* pB0 = pB; - const signed char* pB1 = pB + (size_t)4 * K; - const float* pBD0 = pBD; - const float* pBD1 = pBD + (size_t)4 * num_blocks; + const signed char* pB0 = pB + (size_t)4 * k0; + const signed char* pB1 = pB + (size_t)4 * B_hstep + (size_t)4 * k0; + const float* pBD0 = pBD + (size_t)4 * block_start; + const float* pBD1 = pBD + (size_t)4 * num_blocks + (size_t)4 * block_start; v4f32 _fsum0 = (v4f32)__msa_fill_w(0); v4f32 _fsum1 = (v4f32)__msa_fill_w(0); for (int k = 0; k < K; k += block_size) @@ -2521,11 +2942,40 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de v4i32 _sum0 = __msa_fill_w(0); v4i32 _sum1 = __msa_fill_w(0); int kk = 0; + { + v4i32 _sum2 = __msa_fill_w(0); + v4i32 _sum3 = __msa_fill_w(0); + for (; kk + 7 < max_kk; kk += 8) + { + __builtin_prefetch(pA + 32); + __builtin_prefetch(pB0 + 64); + __builtin_prefetch(pB1 + 64); + v16i8 _pA = (v16i8)__msa_fill_w(*(const int*)pA); + v16i8 _pB0 = __msa_ld_b(pB0, 0); + v16i8 _pB1 = __msa_ld_b(pB1, 0); + _sum0 = __msa_dpadd_s_w(_sum0, __msa_dotp_s_h(_pA, _pB0), _one); + _sum1 = __msa_dpadd_s_w(_sum1, __msa_dotp_s_h(_pA, _pB1), _one); + + _pA = (v16i8)__msa_fill_w(*(const int*)(pA + 4)); + _pB0 = __msa_ld_b(pB0 + 16, 0); + _pB1 = __msa_ld_b(pB1 + 16, 0); + _sum2 = __msa_dpadd_s_w(_sum2, __msa_dotp_s_h(_pA, _pB0), _one); + _sum3 = __msa_dpadd_s_w(_sum3, __msa_dotp_s_h(_pA, _pB1), _one); + pA += 8; + pB0 += 32; + pB1 += 32; + } + _sum0 = __msa_addv_w(_sum0, _sum2); + _sum1 = __msa_addv_w(_sum1, _sum3); + } for (; kk + 3 < max_kk; kk += 4) { - const v16i8 _pA = (v16i8)__msa_fill_w(*(const int*)pA); - const v16i8 _pB0 = __msa_ld_b(pB0, 0); - const v16i8 _pB1 = __msa_ld_b(pB1, 0); + __builtin_prefetch(pA + 32); + __builtin_prefetch(pB0 + 64); + __builtin_prefetch(pB1 + 64); + v16i8 _pA = (v16i8)__msa_fill_w(*(const int*)pA); + v16i8 _pB0 = __msa_ld_b(pB0, 0); + v16i8 _pB1 = __msa_ld_b(pB1, 0); _sum0 = __msa_dpadd_s_w(_sum0, __msa_dotp_s_h(_pA, _pB0), _one); _sum1 = __msa_dpadd_s_w(_sum1, __msa_dotp_s_h(_pA, _pB1), _one); pA += 4; @@ -2534,9 +2984,9 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } if (kk + 1 < max_kk) { - const v16i8 _pA = (v16i8)__msa_fill_h(*(const short*)pA); - const v8i16 _s0 = __msa_dotp_s_h(_pA, (v16i8)__msa_fill_d_ptr(pB0)); - const v8i16 _s1 = __msa_dotp_s_h(_pA, (v16i8)__msa_fill_d_ptr(pB1)); + v16i8 _pA = (v16i8)__msa_fill_h(*(const short*)pA); + v8i16 _s0 = __msa_dotp_s_h(_pA, (v16i8)__msa_fill_d_ptr(pB0)); + v8i16 _s1 = __msa_dotp_s_h(_pA, (v16i8)__msa_fill_d_ptr(pB1)); _sum0 = __msa_addv_w(_sum0, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s0, 0), _s0)); _sum1 = __msa_addv_w(_sum1, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s1, 0), _s1)); pA += 2; @@ -2546,36 +2996,45 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } if (kk < max_kk) { - const v8i16 _pA = __msa_fill_h(pA[0]); - const v16i8 _pB08 = (v16i8)__msa_fill_w(*(const int*)pB0); - const v16i8 _pB18 = (v16i8)__msa_fill_w(*(const int*)pB1); - const v8i16 _pB0 = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pB08, 0), _pB08); - const v8i16 _pB1 = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pB18, 0), _pB18); - const v8i16 _s0 = __msa_mulv_h(_pA, _pB0); - const v8i16 _s1 = __msa_mulv_h(_pA, _pB1); + v8i16 _pA = __msa_fill_h(pA[0]); + v16i8 _pB08 = (v16i8)__msa_fill_w(*(const int*)pB0); + v16i8 _pB18 = (v16i8)__msa_fill_w(*(const int*)pB1); + v8i16 _pB0 = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pB08, 0), _pB08); + v8i16 _pB1 = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pB18, 0), _pB18); + v8i16 _s0 = __msa_mulv_h(_pA, _pB0); + v8i16 _s1 = __msa_mulv_h(_pA, _pB1); _sum0 = __msa_addv_w(_sum0, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s0, 0), _s0)); _sum1 = __msa_addv_w(_sum1, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s1, 0), _s1)); pA++; pB0 += 4; pB1 += 4; } - const v4f32 _descaleB0 = (v4f32)__msa_ld_w(pBD0, 0); - const v4f32 _descaleB1 = (v4f32)__msa_ld_w(pBD1, 0); - const v4f32 _descaleA = __msa_fill_w_f32(pAD[0]); - _fsum0 = __msa_fadd_w(_fsum0, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum0), __msa_fmul_w(_descaleA, _descaleB0))); - _fsum1 = __msa_fadd_w(_fsum1, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum1), __msa_fmul_w(_descaleA, _descaleB1))); + v4f32 _descaleB0 = (v4f32)__msa_ld_w(pBD0, 0); + v4f32 _descaleB1 = (v4f32)__msa_ld_w(pBD1, 0); + v4f32 _descaleA = __msa_fill_w_f32(pAD[0]); + v4f32 _scale = __msa_fmul_w(_descaleA, _descaleB0); + _fsum0 = __ncnn_msa_fmadd_w(_fsum0, (v4f32)__msa_ffint_s_w(_sum0), _scale); + _scale = __msa_fmul_w(_descaleA, _descaleB1); + _fsum1 = __ncnn_msa_fmadd_w(_fsum1, (v4f32)__msa_ffint_s_w(_sum1), _scale); pAD++; pBD0 += 4; pBD1 += 4; } + if (k0 != 0) + { + _fsum0 = __msa_fadd_w(_fsum0, (v4f32)__msa_ld_w(outptr, 0)); + _fsum1 = __msa_fadd_w(_fsum1, (v4f32)__msa_ld_w(outptr + 4, 0)); + } __msa_st_w((v4i32)_fsum0, outptr, 0); __msa_st_w((v4i32)_fsum1, outptr + 4, 0); outptr += 8; - pB = pB1; - pBD = pBD1; + pB += (size_t)8 * B_hstep; + pBD += (size_t)8 * num_blocks; } for (; jj + 3 < max_jj; jj += 4) { + pB += (size_t)4 * k0; + pBD += (size_t)4 * block_start; const signed char* pA = pAT; const float* pAD = pAT_descales; v4f32 _fsum0 = (v4f32)__msa_fill_w(0); @@ -2584,18 +3043,38 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de const int max_kk = std::min(K - k, block_size); v4i32 _sum0 = __msa_fill_w(0); int kk = 0; + { + v4i32 _sum1 = __msa_fill_w(0); + for (; kk + 7 < max_kk; kk += 8) + { + __builtin_prefetch(pA + 32); + __builtin_prefetch(pB + 64); + v16i8 _pA = (v16i8)__msa_fill_w(*(const int*)pA); + v16i8 _pB0 = __msa_ld_b(pB, 0); + _sum0 = __msa_dpadd_s_w(_sum0, __msa_dotp_s_h(_pA, _pB0), _one); + + _pA = (v16i8)__msa_fill_w(*(const int*)(pA + 4)); + _pB0 = __msa_ld_b(pB + 16, 0); + _sum1 = __msa_dpadd_s_w(_sum1, __msa_dotp_s_h(_pA, _pB0), _one); + pA += 8; + pB += 32; + } + _sum0 = __msa_addv_w(_sum0, _sum1); + } for (; kk + 3 < max_kk; kk += 4) { - const v16i8 _pA = (v16i8)__msa_fill_w(*(const int*)pA); - const v16i8 _pB0 = __msa_ld_b(pB, 0); + __builtin_prefetch(pA + 32); + __builtin_prefetch(pB + 32); + v16i8 _pA = (v16i8)__msa_fill_w(*(const int*)pA); + v16i8 _pB0 = __msa_ld_b(pB, 0); _sum0 = __msa_dpadd_s_w(_sum0, __msa_dotp_s_h(_pA, _pB0), _one); pA += 4; pB += 16; } if (kk + 1 < max_kk) { - const v16i8 _pA = (v16i8)__msa_fill_h(*(const short*)pA); - const v8i16 _s = __msa_dotp_s_h(_pA, (v16i8)__msa_fill_d_ptr(pB)); + v16i8 _pA = (v16i8)__msa_fill_h(*(const short*)pA); + v8i16 _s = __msa_dotp_s_h(_pA, (v16i8)__msa_fill_d_ptr(pB)); _sum0 = __msa_addv_w(_sum0, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); pA += 2; pB += 8; @@ -2603,24 +3082,138 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } if (kk < max_kk) { - const v16i8 _pB8 = (v16i8)__msa_fill_w(*(const int*)pB); - const v8i16 _pB = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pB8, 0), _pB8); - const v8i16 _s = __msa_mulv_h(__msa_fill_h(pA[0]), _pB); + v16i8 _pB8 = (v16i8)__msa_fill_w(*(const int*)pB); + v8i16 _pB = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pB8, 0), _pB8); + v8i16 _s = __msa_mulv_h(__msa_fill_h(pA[0]), _pB); _sum0 = __msa_addv_w(_sum0, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); pA++; pB += 4; } - const v4f32 _descaleB = (v4f32)__msa_ld_w(pBD, 0); - _fsum0 = __msa_fadd_w(_fsum0, __msa_fmul_w((v4f32)__msa_ffint_s_w(_sum0), __msa_fmul_w(_descaleB, __msa_fill_w_f32(pAD[0])))); + v4f32 _descaleB = (v4f32)__msa_ld_w(pBD, 0); + v4f32 _scale = __msa_fmul_w(_descaleB, __msa_fill_w_f32(pAD[0])); + _fsum0 = __ncnn_msa_fmadd_w(_fsum0, (v4f32)__msa_ffint_s_w(_sum0), _scale); pAD++; pBD += 4; } + if (k0 != 0) + _fsum0 = __msa_fadd_w(_fsum0, (v4f32)__msa_ld_w(outptr, 0)); __msa_st_w((v4i32)_fsum0, outptr, 0); outptr += 4; + pB += (size_t)4 * (B_hstep - k0 - max_kk0); + pBD += (size_t)4 * (num_blocks - block_start - tile_blocks); + } + for (; jj + 1 < max_jj; jj += 2) + { + pB += (size_t)2 * k0; + pBD += (size_t)2 * block_start; + const signed char* pA = pAT; + const float* pAD = pAT_descales; + v4f32 _fsum = (v4f32)__msa_fill_w(0); + for (int k = 0; k < K; k += block_size) + { + const int max_kk = std::min(K - k, block_size); + v4i32 _sum = __msa_fill_w(0); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __builtin_prefetch(pA + 16); + __builtin_prefetch(pB + 32); + v16i8 _pA = (v16i8)__msa_fill_w(*(const int*)pA); + v16i8 _pB = (v16i8)__msa_fill_d_ptr(pB); + _sum = __msa_dpadd_s_w(_sum, __msa_dotp_s_h(_pA, _pB), _one); + pA += 4; + pB += 8; + } + if (kk + 1 < max_kk) + { + v16i8 _pA = (v16i8)__msa_fill_h(*(const short*)pA); + v16i8 _pB = (v16i8)__msa_fill_w(*(const int*)pB); + v8i16 _s = __msa_dotp_s_h(_pA, _pB); + _sum = __msa_addv_w(_sum, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + pA += 2; + pB += 4; + kk += 2; + } + if (kk < max_kk) + { + v16i8 _pB8 = (v16i8)__msa_fill_h(*(const short*)pB); + v8i16 _pB = (v8i16)__msa_ilvr_b(__msa_clti_s_b(_pB8, 0), _pB8); + v8i16 _s = __msa_mulv_h(__msa_fill_h(pA[0]), _pB); + _sum = __msa_addv_w(_sum, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + pA++; + pB += 2; + } + v4f32 _descaleB = (v4f32) { + pBD[0], pBD[1], pBD[0], pBD[1] + }; + v4f32 _scale = __msa_fmul_w(_descaleB, __msa_fill_w_f32(pAD[0])); + _fsum = __ncnn_msa_fmadd_w(_fsum, (v4f32)__msa_ffint_s_w(_sum), _scale); + pAD++; + pBD += 2; + } + if (k0 != 0) + _fsum = __msa_fadd_w(_fsum, (v4f32)__msa_loadl_d(outptr)); + __msa_storel_d((v4i32)_fsum, outptr); + outptr += 2; + pB += (size_t)2 * (B_hstep - k0 - max_kk0); + pBD += (size_t)2 * (num_blocks - block_start - tile_blocks); + } + for (; jj < max_jj; jj++) + { + pB += k0; + pBD += block_start; + const signed char* pA = pAT; + const float* pAD = pAT_descales; + v4f32 _fsum = (v4f32)__msa_fill_w(0); + for (int k = 0; k < K; k += block_size) + { + const int max_kk = std::min(K - k, block_size); + v4i32 _sum = __msa_fill_w(0); + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + __builtin_prefetch(pA + 16); + __builtin_prefetch(pB + 16); + v16i8 _pA = (v16i8)__msa_fill_w(*(const int*)pA); + v16i8 _pB = (v16i8)__msa_fill_w(*(const int*)pB); + _sum = __msa_dpadd_s_w(_sum, __msa_dotp_s_h(_pA, _pB), _one); + pA += 4; + pB += 4; + } + if (kk + 1 < max_kk) + { + v16i8 _pA = (v16i8)__msa_fill_h(*(const short*)pA); + v16i8 _pB = (v16i8)__msa_fill_h(*(const short*)pB); + v8i16 _s = __msa_dotp_s_h(_pA, _pB); + _sum = __msa_addv_w(_sum, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + pA += 2; + pB += 2; + kk += 2; + } + if (kk < max_kk) + { + v8i16 _s = __msa_fill_h(pA[0] * pB[0]); + _sum = __msa_addv_w(_sum, (v4i32)__msa_ilvr_h(__msa_clti_s_h(_s, 0), _s)); + pA++; + pB++; + } + v4f32 _scale = __msa_fill_w_f32(pAD[0] * pBD[0]); + _fsum = __ncnn_msa_fmadd_w(_fsum, (v4f32)__msa_ffint_s_w(_sum), _scale); + pAD++; + pBD++; + } + if (k0 != 0) + _fsum[0] += outptr[0]; + *outptr++ = _fsum[0]; + pB += B_hstep - k0 - max_kk0; + pBD += num_blocks - block_start - tile_blocks; } #endif // __mips_msa +#if !__mips_msa for (; jj + 1 < max_jj; jj += 2) { + pB += (size_t)2 * k0; + pBD += (size_t)2 * block_start; const signed char* pA = pAT; const float* pAD = pAT_descales; float sum0 = 0.f; @@ -2702,8 +3295,8 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de for (; kk + 3 < max_kk; kk += 4) { __builtin_prefetch(pB + 32); - const int8x8_t _pA = (int8x8_t)__mmi_pfillw_s(*(const int*)pA); - const int8x8_t _pB = __mmi_pldb_s(pB); + int8x8_t _pA = (int8x8_t)__mmi_pfillw_s(*(const int*)pA); + int8x8_t _pB = __mmi_pldb_s(pB); int16x4_t _pA0 = (int16x4_t)__mmi_punpcklbh_s(_pA, _zero); int16x4_t _pB0 = (int16x4_t)__mmi_punpcklbh_s(_pB, _zero); int16x4_t _pB1 = (int16x4_t)__mmi_punpckhbh_s(_pB, _zero); @@ -2749,12 +3342,21 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pBD += 2; } + if (k0 != 0) + { + sum0 += outptr[0]; + sum1 += outptr[1]; + } outptr[0] = sum0; outptr[1] = sum1; outptr += 2; + pB += (size_t)2 * (B_hstep - k0 - max_kk0); + pBD += (size_t)2 * (num_blocks - block_start - tile_blocks); } for (; jj < max_jj; jj++) { + pB += k0; + pBD += block_start; const signed char* pA = pAT; const float* pAD = pAT_descales; float sum0 = 0.f; @@ -2824,8 +3426,8 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de const int8x8_t _zero = __mmi_pzerob_s(); for (; kk + 3 < max_kk; kk += 4) { - const int8x8_t _pA = (int8x8_t)__mmi_pfillw_s(*(const int*)pA); - const int8x8_t _pB = (int8x8_t)__mmi_pfillw_s(*(const int*)pB); + int8x8_t _pA = (int8x8_t)__mmi_pfillw_s(*(const int*)pA); + int8x8_t _pB = (int8x8_t)__mmi_pfillw_s(*(const int*)pB); int16x4_t _pA0 = (int16x4_t)__mmi_punpcklbh_s(_pA, _zero); int16x4_t _pB0 = (int16x4_t)__mmi_punpcklbh_s(_pB, _zero); _pA0 = __mmi_psrah_s(__mmi_psllh_s(_pA0, 8), 8); @@ -2858,8 +3460,13 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pBD++; } + if (k0 != 0) + sum0 += outptr[0]; *outptr++ = sum0; + pB += B_hstep - k0 - max_kk0; + pBD += num_blocks - block_start - tile_blocks; } +#endif // !__mips_msa pAT += A_hstep; pAT_descales += A_descales_hstep; @@ -2914,7 +3521,7 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b _c4567 = (v4f32)__msa_ld_w(pC + 4, 0); if (beta != 1.f) { - const v4f32 _beta = __msa_fill_w_f32(beta); + v4f32 _beta = __msa_fill_w_f32(beta); _c0123 = __msa_fmul_w(_c0123, _beta); _c4567 = __msa_fmul_w(_c4567, _beta); } @@ -2960,7 +3567,7 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b } if (broadcast_type_C == 3) { - const v4f32 _beta = __msa_fill_w_f32(beta); + v4f32 _beta = __msa_fill_w_f32(beta); v4f32 _c = (v4f32)__msa_ld_w(pC0, 0); _f0 = __msa_fadd_w(_f0, beta == 1.f ? _c : __msa_fmul_w(_c, _beta)); _c = (v4f32)__msa_ld_w(pC1, 0); @@ -2977,6 +3584,14 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b _f6 = __msa_fadd_w(_f6, beta == 1.f ? _c : __msa_fmul_w(_c, _beta)); _c = (v4f32)__msa_ld_w(pC7, 0); _f7 = __msa_fadd_w(_f7, beta == 1.f ? _c : __msa_fmul_w(_c, _beta)); + pC0 += 4; + pC1 += 4; + pC2 += 4; + pC3 += 4; + pC4 += 4; + pC5 += 4; + pC6 += 4; + pC7 += 4; } if (broadcast_type_C == 4) { @@ -2991,12 +3606,13 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b _f5 = __msa_fadd_w(_f5, _c); _f6 = __msa_fadd_w(_f6, _c); _f7 = __msa_fadd_w(_f7, _c); + pC += 4; } } if (alpha != 1.f) { - const v4f32 _alpha = __msa_fill_w_f32(alpha); + v4f32 _alpha = __msa_fill_w_f32(alpha); _f0 = __msa_fmul_w(_f0, _alpha); _f1 = __msa_fmul_w(_f1, _alpha); _f2 = __msa_fmul_w(_f2, _alpha); @@ -3022,19 +3638,6 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b outptr5 += 4; outptr6 += 4; outptr7 += 4; - if (pC0) - { - pC0 += 4; - pC1 += 4; - pC2 += 4; - pC3 += 4; - pC4 += 4; - pC5 += 4; - pC6 += 4; - pC7 += 4; - } - if (pC && broadcast_type_C == 4) - pC += 4; } for (; jj + 1 < max_jj; jj += 2) { @@ -3069,7 +3672,7 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b }; if (beta != 1.f) { - const v4f32 _beta = __msa_fill_w_f32(beta); + v4f32 _beta = __msa_fill_w_f32(beta); _c0 = __msa_fmul_w(_c0, _beta); _c4 = __msa_fmul_w(_c4, _beta); _c1 = __msa_fmul_w(_c1, _beta); @@ -3079,6 +3682,14 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b _f4 = __msa_fadd_w(_f4, _c4); _f1 = __msa_fadd_w(_f1, _c1); _f5 = __msa_fadd_w(_f5, _c5); + pC0 += 2; + pC1 += 2; + pC2 += 2; + pC3 += 2; + pC4 += 2; + pC5 += 2; + pC6 += 2; + pC7 += 2; } if (broadcast_type_C == 4) { @@ -3093,12 +3704,13 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b _f4 = __msa_fadd_w(_f4, __msa_fill_w_f32(c0)); _f1 = __msa_fadd_w(_f1, __msa_fill_w_f32(c1)); _f5 = __msa_fadd_w(_f5, __msa_fill_w_f32(c1)); + pC += 2; } } if (alpha != 1.f) { - const v4f32 _alpha = __msa_fill_w_f32(alpha); + v4f32 _alpha = __msa_fill_w_f32(alpha); _f0 = __msa_fmul_w(_f0, _alpha); _f4 = __msa_fmul_w(_f4, _alpha); _f1 = __msa_fmul_w(_f1, _alpha); @@ -3128,19 +3740,6 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b outptr5 += 2; outptr6 += 2; outptr7 += 2; - if (pC0) - { - pC0 += 2; - pC1 += 2; - pC2 += 2; - pC3 += 2; - pC4 += 2; - pC5 += 2; - pC6 += 2; - pC7 += 2; - } - if (pC && broadcast_type_C == 4) - pC += 2; } for (; jj < max_jj; jj++) { @@ -3165,12 +3764,20 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b }; if (beta != 1.f) { - const v4f32 _beta = __msa_fill_w_f32(beta); + v4f32 _beta = __msa_fill_w_f32(beta); _c0 = __msa_fmul_w(_c0, _beta); _c4 = __msa_fmul_w(_c4, _beta); } _f0 = __msa_fadd_w(_f0, _c0); _f4 = __msa_fadd_w(_f4, _c4); + pC0++; + pC1++; + pC2++; + pC3++; + pC4++; + pC5++; + pC6++; + pC7++; } if (broadcast_type_C == 4) { @@ -3179,12 +3786,13 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b c *= beta; _f0 = __msa_fadd_w(_f0, __msa_fill_w_f32(c)); _f4 = __msa_fadd_w(_f4, __msa_fill_w_f32(c)); + pC++; } } if (alpha != 1.f) { - const v4f32 _alpha = __msa_fill_w_f32(alpha); + v4f32 _alpha = __msa_fill_w_f32(alpha); _f0 = __msa_fmul_w(_f0, _alpha); _f4 = __msa_fmul_w(_f4, _alpha); } @@ -3306,6 +3914,10 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b _f5 = __msa_fadd_w(_f5, _c5); _f6 = __msa_fadd_w(_f6, _c6); _f7 = __msa_fadd_w(_f7, _c7); + pC0 += 8; + pC1 += 8; + pC2 += 8; + pC3 += 8; } if (broadcast_type_C == 4) { @@ -3325,6 +3937,7 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b _f5 = __msa_fadd_w(_f5, (v4f32)__msa_splati_w((v4i32)_c4, 1)); _f6 = __msa_fadd_w(_f6, (v4f32)__msa_splati_w((v4i32)_c4, 2)); _f7 = __msa_fadd_w(_f7, (v4f32)__msa_splati_w((v4i32)_c4, 3)); + pC += 8; } } @@ -3355,15 +3968,6 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b outptr1 += 8; outptr2 += 8; outptr3 += 8; - if (pC0) - { - pC0 += 8; - pC1 += 8; - pC2 += 8; - pC3 += 8; - } - if (pC && broadcast_type_C == 4) - pC += 8; } for (; jj + 3 < max_jj; jj += 4) { @@ -3401,6 +4005,10 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b _f1 = __msa_fadd_w(_f1, _c1); _f2 = __msa_fadd_w(_f2, _c2); _f3 = __msa_fadd_w(_f3, _c3); + pC0 += 4; + pC1 += 4; + pC2 += 4; + pC3 += 4; } if (broadcast_type_C == 4) { @@ -3411,6 +4019,7 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b _f1 = __msa_fadd_w(_f1, (v4f32)__msa_splati_w((v4i32)_c0, 1)); _f2 = __msa_fadd_w(_f2, (v4f32)__msa_splati_w((v4i32)_c0, 2)); _f3 = __msa_fadd_w(_f3, (v4f32)__msa_splati_w((v4i32)_c0, 3)); + pC += 4; } } @@ -3432,15 +4041,6 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b outptr1 += 4; outptr2 += 4; outptr3 += 4; - if (pC0) - { - pC0 += 4; - pC1 += 4; - pC2 += 4; - pC3 += 4; - } - if (pC && broadcast_type_C == 4) - pC += 4; } for (; jj + 1 < max_jj; jj += 2) { @@ -3471,6 +4071,10 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b } _f0 = __msa_fadd_w(_f0, _c0); _f1 = __msa_fadd_w(_f1, _c1); + pC0 += 2; + pC1 += 2; + pC2 += 2; + pC3 += 2; } if (broadcast_type_C == 4) { @@ -3483,6 +4087,7 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b } _f0 = __msa_fadd_w(_f0, __msa_fill_w_f32(c0)); _f1 = __msa_fadd_w(_f1, __msa_fill_w_f32(c1)); + pC += 2; } } @@ -3504,15 +4109,6 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b outptr1 += 2; outptr2 += 2; outptr3 += 2; - if (pC0) - { - pC0 += 2; - pC1 += 2; - pC2 += 2; - pC3 += 2; - } - if (pC && broadcast_type_C == 4) - pC += 2; } for (; jj < max_jj; jj++) { @@ -3532,6 +4128,10 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b if (beta != 1.f) _c0 = __msa_fmul_w(_c0, __msa_fill_w_f32(beta)); _f0 = __msa_fadd_w(_f0, _c0); + pC0++; + pC1++; + pC2++; + pC3++; } if (broadcast_type_C == 4) { @@ -3539,6 +4139,7 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b if (beta != 1.f) c *= beta; _f0 = __msa_fadd_w(_f0, __msa_fill_w_f32(c)); + pC++; } } if (alpha != 1.f) @@ -3551,15 +4152,6 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b outptr1++; outptr2++; outptr3++; - if (pC0) - { - pC0++; - pC1++; - pC2++; - pC3++; - } - if (pC && broadcast_type_C == 4) - pC++; } outptr += out_hstep * 4; } @@ -3636,6 +4228,8 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b _f1 = __msa_fadd_w(_f1, _c1); _f2 = __msa_fadd_w(_f2, _c2); _f3 = __msa_fadd_w(_f3, _c3); + pC0 += 8; + pC1 += 8; } if (broadcast_type_C == 4) { @@ -3651,6 +4245,7 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b _f1 = __msa_fadd_w(_f1, _c1); _f2 = __msa_fadd_w(_f2, _c0); _f3 = __msa_fadd_w(_f3, _c1); + pC += 8; } } @@ -3669,13 +4264,6 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b __msa_st_w((v4i32)_f3, outptr1 + 4, 0); outptr0 += 8; outptr1 += 8; - if (pC0) - { - pC0 += 8; - pC1 += 8; - } - if (pC && broadcast_type_C == 4) - pC += 8; } for (; jj + 3 < max_jj; jj += 4) { @@ -3705,6 +4293,8 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b } _f0 = __msa_fadd_w(_f0, _c0); _f1 = __msa_fadd_w(_f1, _c1); + pC0 += 4; + pC1 += 4; } if (broadcast_type_C == 4) { @@ -3713,6 +4303,7 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b _c0 = __msa_fmul_w(_c0, __msa_fill_w_f32(beta)); _f0 = __msa_fadd_w(_f0, _c0); _f1 = __msa_fadd_w(_f1, _c0); + pC += 4; } } @@ -3727,13 +4318,6 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b __msa_st_w((v4i32)_f1, outptr1, 0); outptr0 += 4; outptr1 += 4; - if (pC0) - { - pC0 += 4; - pC1 += 4; - } - if (pC && broadcast_type_C == 4) - pC += 4; } for (; jj + 1 < max_jj; jj += 2) { @@ -3754,6 +4338,8 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b if (beta != 1.f) _c = __msa_fmul_w(_c, __msa_fill_w_f32(beta)); _f = __msa_fadd_w(_f, _c); + pC0 += 2; + pC1 += 2; } if (broadcast_type_C == 4) { @@ -3767,6 +4353,7 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b _f = __msa_fadd_w(_f, (v4f32) { cc0, cc0, cc1, cc1 }); + pC += 2; } } @@ -3779,13 +4366,6 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b __msa_storel_d(_f1, outptr1); outptr0 += 2; outptr1 += 2; - if (pC0) - { - pC0 += 2; - pC1 += 2; - } - if (pC && broadcast_type_C == 4) - pC += 2; } #endif // __mips_msa for (; jj + 1 < max_jj; jj += 2) @@ -3829,6 +4409,8 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b sum01 += c01; sum10 += c10; sum11 += c11; + pC0 += 2; + pC1 += 2; } if (broadcast_type_C == 4) { @@ -3843,6 +4425,7 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b sum01 += c0; sum10 += c1; sum11 += c1; + pC += 2; } } @@ -3860,13 +4443,6 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b outptr1[1] = sum11; outptr0 += 2; outptr1 += 2; - if (pC0) - { - pC0 += 2; - pC1 += 2; - } - if (pC && broadcast_type_C == 4) - pC += 2; } for (; jj < max_jj; jj++) { @@ -3896,6 +4472,8 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b } sum0 += c0; sum1 += c1; + pC0++; + pC1++; } if (broadcast_type_C == 4) { @@ -3903,6 +4481,7 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b if (beta != 1.f) c *= beta; sum0 += c; sum1 += c; + pC++; } } if (alpha != 1.f) @@ -3914,13 +4493,6 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b outptr1[0] = sum1; outptr0++; outptr1++; - if (pC0) - { - pC0++; - pC1++; - } - if (pC && broadcast_type_C == 4) - pC++; } outptr += out_hstep * 2; } @@ -3975,6 +4547,7 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b } _f0 = __msa_fadd_w(_f0, _c0); _f1 = __msa_fadd_w(_f1, _c1); + pC0 += 8; } } @@ -3988,8 +4561,6 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b __msa_st_w((v4i32)_f0, outptr0, 0); __msa_st_w((v4i32)_f1, outptr0 + 4, 0); outptr0 += 8; - if (pC0) - pC0 += 8; } for (; jj + 3 < max_jj; jj += 4) { @@ -4006,6 +4577,7 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b if (beta != 1.f) _c0 = __msa_fmul_w(_c0, __msa_fill_w_f32(beta)); _f0 = __msa_fadd_w(_f0, _c0); + pC0 += 4; } } @@ -4014,8 +4586,6 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b __msa_st_w((v4i32)_f0, outptr0, 0); outptr0 += 4; - if (pC0) - pC0 += 4; } for (; jj + 1 < max_jj; jj += 2) { @@ -4038,6 +4608,7 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b if (beta != 1.f) _c = __msa_fmul_w(_c, __msa_fill_w_f32(beta)); _f = __msa_fadd_w(_f, _c); + pC0 += 2; } } @@ -4046,8 +4617,6 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b __msa_storel_d((v4i32)_f, outptr0); outptr0 += 2; - if (pC0) - pC0 += 2; } #endif // __mips_msa for (; jj + 1 < max_jj; jj += 2) @@ -4078,6 +4647,7 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b } sum0 += c0; sum1 += c1; + pC0 += 2; } if (broadcast_type_C == 4) { @@ -4090,6 +4660,7 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b } sum0 += c0; sum1 += c1; + pC0 += 2; } } if (alpha != 1.f) @@ -4100,8 +4671,6 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b outptr0[0] = sum0; outptr0[1] = sum1; outptr0 += 2; - if (pC0) - pC0 += 2; } for (; jj < max_jj; jj++) { @@ -4113,12 +4682,11 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b if (broadcast_type_C == 3 || broadcast_type_C == 4) c = pC0[0]; if ((broadcast_type_C == 3 || broadcast_type_C == 4) && beta != 1.f) c *= beta; sum0 += c; + if (broadcast_type_C == 3 || broadcast_type_C == 4) pC0++; } if (alpha != 1.f) sum0 *= alpha; outptr0[0] = sum0; outptr0++; - if (pC0) - pC0++; } outptr += out_hstep; } @@ -4165,7 +4733,7 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, _c4567 = (v4f32)__msa_ld_w(pC + 4, 0); if (beta != 1.f) { - const v4f32 _beta = __msa_fill_w_f32(beta); + v4f32 _beta = __msa_fill_w_f32(beta); _c0123 = __msa_fmul_w(_c0123, _beta); _c4567 = __msa_fmul_w(_c4567, _beta); } @@ -4235,7 +4803,7 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, }; if (beta != 1.f) { - const v4f32 _beta = __msa_fill_w_f32(beta); + v4f32 _beta = __msa_fill_w_f32(beta); _cl0 = __msa_fmul_w(_cl0, _beta); _ch0 = __msa_fmul_w(_ch0, _beta); _cl1 = __msa_fmul_w(_cl1, _beta); @@ -4253,6 +4821,14 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, _f6 = __msa_fadd_w(_f6, _ch2); _f3 = __msa_fadd_w(_f3, _cl3); _f7 = __msa_fadd_w(_f7, _ch3); + pC0 += 4; + pC1 += 4; + pC2 += 4; + pC3 += 4; + pC4 += 4; + pC5 += 4; + pC6 += 4; + pC7 += 4; } if (broadcast_type_C == 4) { @@ -4267,12 +4843,13 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, _f6 = __msa_fadd_w(_f6, (v4f32)__msa_splati_w((v4i32)_c, 2)); _f3 = __msa_fadd_w(_f3, (v4f32)__msa_splati_w((v4i32)_c, 3)); _f7 = __msa_fadd_w(_f7, (v4f32)__msa_splati_w((v4i32)_c, 3)); + pC += 4; } } if (alpha != 1.f) { - const v4f32 _alpha = __msa_fill_w_f32(alpha); + v4f32 _alpha = __msa_fill_w_f32(alpha); _f0 = __msa_fmul_w(_f0, _alpha); _f1 = __msa_fmul_w(_f1, _alpha); _f2 = __msa_fmul_w(_f2, _alpha); @@ -4291,19 +4868,6 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, __msa_st_w((v4i32)_f3, outptr + out_hstep * 3, 0); __msa_st_w((v4i32)_f7, outptr + out_hstep * 3 + 4, 0); outptr += out_hstep * 4; - if (pC0) - { - pC0 += 4; - pC1 += 4; - pC2 += 4; - pC3 += 4; - pC4 += 4; - pC5 += 4; - pC6 += 4; - pC7 += 4; - } - if (pC && broadcast_type_C == 4) - pC += 4; } for (; jj + 1 < max_jj; jj += 2) { @@ -4338,7 +4902,7 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, }; if (beta != 1.f) { - const v4f32 _beta = __msa_fill_w_f32(beta); + v4f32 _beta = __msa_fill_w_f32(beta); _c0 = __msa_fmul_w(_c0, _beta); _c4 = __msa_fmul_w(_c4, _beta); _c1 = __msa_fmul_w(_c1, _beta); @@ -4348,6 +4912,14 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, _f4 = __msa_fadd_w(_f4, _c4); _f1 = __msa_fadd_w(_f1, _c1); _f5 = __msa_fadd_w(_f5, _c5); + pC0 += 2; + pC1 += 2; + pC2 += 2; + pC3 += 2; + pC4 += 2; + pC5 += 2; + pC6 += 2; + pC7 += 2; } if (broadcast_type_C == 4) { @@ -4362,12 +4934,13 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, _f4 = __msa_fadd_w(_f4, __msa_fill_w_f32(c0)); _f1 = __msa_fadd_w(_f1, __msa_fill_w_f32(c1)); _f5 = __msa_fadd_w(_f5, __msa_fill_w_f32(c1)); + pC += 2; } } if (alpha != 1.f) { - const v4f32 _alpha = __msa_fill_w_f32(alpha); + v4f32 _alpha = __msa_fill_w_f32(alpha); _f0 = __msa_fmul_w(_f0, _alpha); _f4 = __msa_fmul_w(_f4, _alpha); _f1 = __msa_fmul_w(_f1, _alpha); @@ -4378,19 +4951,6 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, __msa_st_w((v4i32)_f1, outptr + out_hstep, 0); __msa_st_w((v4i32)_f5, outptr + out_hstep + 4, 0); outptr += out_hstep * 2; - if (pC0) - { - pC0 += 2; - pC1 += 2; - pC2 += 2; - pC3 += 2; - pC4 += 2; - pC5 += 2; - pC6 += 2; - pC7 += 2; - } - if (pC && broadcast_type_C == 4) - pC += 2; } for (; jj < max_jj; jj++) { @@ -4415,12 +4975,20 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, }; if (beta != 1.f) { - const v4f32 _beta = __msa_fill_w_f32(beta); + v4f32 _beta = __msa_fill_w_f32(beta); _c0 = __msa_fmul_w(_c0, _beta); _c4 = __msa_fmul_w(_c4, _beta); } _f0 = __msa_fadd_w(_f0, _c0); _f4 = __msa_fadd_w(_f4, _c4); + pC0++; + pC1++; + pC2++; + pC3++; + pC4++; + pC5++; + pC6++; + pC7++; } if (broadcast_type_C == 4) { @@ -4429,12 +4997,13 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, c *= beta; _f0 = __msa_fadd_w(_f0, __msa_fill_w_f32(c)); _f4 = __msa_fadd_w(_f4, __msa_fill_w_f32(c)); + pC++; } } if (alpha != 1.f) { - const v4f32 _alpha = __msa_fill_w_f32(alpha); + v4f32 _alpha = __msa_fill_w_f32(alpha); _f0 = __msa_fmul_w(_f0, _alpha); _f4 = __msa_fmul_w(_f4, _alpha); } @@ -4540,6 +5109,10 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, _f5 = __msa_fadd_w(_f5, _c5); _f6 = __msa_fadd_w(_f6, _c6); _f7 = __msa_fadd_w(_f7, _c7); + pC0 += 8; + pC1 += 8; + pC2 += 8; + pC3 += 8; } if (broadcast_type_C == 4) { @@ -4559,6 +5132,7 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, _f5 = __msa_fadd_w(_f5, (v4f32)__msa_splati_w((v4i32)_c4, 1)); _f6 = __msa_fadd_w(_f6, (v4f32)__msa_splati_w((v4i32)_c4, 2)); _f7 = __msa_fadd_w(_f7, (v4f32)__msa_splati_w((v4i32)_c4, 3)); + pC += 8; } } @@ -4584,15 +5158,6 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, __msa_st_w((v4i32)_f6, outptr + out_hstep * 6, 0); __msa_st_w((v4i32)_f7, outptr + out_hstep * 7, 0); outptr += out_hstep * 8; - if (pC0) - { - pC0 += 8; - pC1 += 8; - pC2 += 8; - pC3 += 8; - } - if (pC && broadcast_type_C == 4) - pC += 8; } for (; jj + 3 < max_jj; jj += 4) { @@ -4630,6 +5195,10 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, _f1 = __msa_fadd_w(_f1, _c1); _f2 = __msa_fadd_w(_f2, _c2); _f3 = __msa_fadd_w(_f3, _c3); + pC0 += 4; + pC1 += 4; + pC2 += 4; + pC3 += 4; } if (broadcast_type_C == 4) { @@ -4640,6 +5209,7 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, _f1 = __msa_fadd_w(_f1, (v4f32)__msa_splati_w((v4i32)_c0, 1)); _f2 = __msa_fadd_w(_f2, (v4f32)__msa_splati_w((v4i32)_c0, 2)); _f3 = __msa_fadd_w(_f3, (v4f32)__msa_splati_w((v4i32)_c0, 3)); + pC += 4; } } @@ -4657,15 +5227,6 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, __msa_st_w((v4i32)_f2, outptr + out_hstep * 2, 0); __msa_st_w((v4i32)_f3, outptr + out_hstep * 3, 0); outptr += out_hstep * 4; - if (pC0) - { - pC0 += 4; - pC1 += 4; - pC2 += 4; - pC3 += 4; - } - if (pC && broadcast_type_C == 4) - pC += 4; } for (; jj + 1 < max_jj; jj += 2) { @@ -4696,6 +5257,10 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, } _f0 = __msa_fadd_w(_f0, _c0); _f1 = __msa_fadd_w(_f1, _c1); + pC0 += 2; + pC1 += 2; + pC2 += 2; + pC3 += 2; } if (broadcast_type_C == 4) { @@ -4708,6 +5273,7 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, } _f0 = __msa_fadd_w(_f0, __msa_fill_w_f32(c0)); _f1 = __msa_fadd_w(_f1, __msa_fill_w_f32(c1)); + pC += 2; } } @@ -4721,15 +5287,6 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, __msa_st_w((v4i32)_f0, outptr, 0); __msa_st_w((v4i32)_f1, outptr + out_hstep, 0); outptr += out_hstep * 2; - if (pC0) - { - pC0 += 2; - pC1 += 2; - pC2 += 2; - pC3 += 2; - } - if (pC && broadcast_type_C == 4) - pC += 2; } for (; jj < max_jj; jj++) { @@ -4749,6 +5306,10 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, if (beta != 1.f) _c0 = __msa_fmul_w(_c0, __msa_fill_w_f32(beta)); _f0 = __msa_fadd_w(_f0, _c0); + pC0++; + pC1++; + pC2++; + pC3++; } if (broadcast_type_C == 4) { @@ -4756,21 +5317,13 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, if (beta != 1.f) c *= beta; _f0 = __msa_fadd_w(_f0, __msa_fill_w_f32(c)); + pC++; } } if (alpha != 1.f) _f0 = __msa_fmul_w(_f0, __msa_fill_w_f32(alpha)); __msa_st_w((v4i32)_f0, outptr, 0); outptr += out_hstep; - if (pC0) - { - pC0++; - pC1++; - pC2++; - pC3++; - } - if (pC && broadcast_type_C == 4) - pC++; } outptr0 += 4; } @@ -4852,6 +5405,8 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, _f1 = __msa_fadd_w(_f1, _c1); _f2 = __msa_fadd_w(_f2, _c2); _f3 = __msa_fadd_w(_f3, _c3); + pC0 += 8; + pC1 += 8; } if (broadcast_type_C == 4) { @@ -4886,6 +5441,7 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, _f3 = __msa_fadd_w(_f3, (v4f32) { c06, c06, c07, c07 }); + pC += 8; } } @@ -4907,13 +5463,6 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, __msa_storel_d((v4i32)_f3, outptr + out_hstep * 6); __msa_storel_d((v4i32)__msa_sldi_b((v16i8)_f3, (v16i8)_f3, 8), outptr + out_hstep * 7); outptr += out_hstep * 8; - if (pC0) - { - pC0 += 8; - pC1 += 8; - } - if (pC && broadcast_type_C == 4) - pC += 8; } for (; jj + 3 < max_jj; jj += 4) { @@ -4947,6 +5496,8 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, } _f0 = __msa_fadd_w(_f0, _c0); _f1 = __msa_fadd_w(_f1, _c1); + pC0 += 4; + pC1 += 4; } if (broadcast_type_C == 4) { @@ -4967,6 +5518,7 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, _f1 = __msa_fadd_w(_f1, (v4f32) { c02, c02, c03, c03 }); + pC += 4; } } @@ -4982,13 +5534,6 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, __msa_storel_d((v4i32)_f1, outptr + out_hstep * 2); __msa_storel_d((v4i32)__msa_sldi_b((v16i8)_f1, (v16i8)_f1, 8), outptr + out_hstep * 3); outptr += out_hstep * 4; - if (pC0) - { - pC0 += 4; - pC1 += 4; - } - if (pC && broadcast_type_C == 4) - pC += 4; } for (; jj + 1 < max_jj; jj += 2) { @@ -5009,6 +5554,8 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, if (beta != 1.f) _c = __msa_fmul_w(_c, __msa_fill_w_f32(beta)); _f = __msa_fadd_w(_f, _c); + pC0 += 2; + pC1 += 2; } if (broadcast_type_C == 4) { @@ -5022,6 +5569,7 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, _f = __msa_fadd_w(_f, (v4f32) { cc0, cc0, cc1, cc1 }); + pC += 2; } } @@ -5031,13 +5579,6 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, __msa_storel_d((v4i32)_f, outptr); __msa_storel_d((v4i32)__msa_sldi_b((v16i8)_f, (v16i8)_f, 8), outptr + out_hstep); outptr += out_hstep * 2; - if (pC0) - { - pC0 += 2; - pC1 += 2; - } - if (pC && broadcast_type_C == 4) - pC += 2; } #endif // __mips_msa for (; jj + 1 < max_jj; jj += 2) @@ -5080,6 +5621,8 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, sum01 += c01; sum10 += c10; sum11 += c11; + pC0 += 2; + pC1 += 2; } if (broadcast_type_C == 4) { @@ -5094,6 +5637,7 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, sum01 += c0; sum10 += c1; sum11 += c1; + pC += 2; } } if (alpha != 1.f) @@ -5108,13 +5652,6 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, outptr[out_hstep] = sum10; outptr[out_hstep + 1] = sum11; outptr += out_hstep * 2; - if (pC0) - { - pC0 += 2; - pC1 += 2; - } - if (pC && broadcast_type_C == 4) - pC += 2; } for (; jj < max_jj; jj++) { @@ -5144,6 +5681,8 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, } sum0 += c0; sum1 += c1; + pC0++; + pC1++; } if (broadcast_type_C == 4) { @@ -5151,6 +5690,7 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, if (beta != 1.f) c *= beta; sum0 += c; sum1 += c; + pC++; } } if (alpha != 1.f) @@ -5161,13 +5701,6 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, outptr[0] = sum0; outptr[1] = sum1; outptr += out_hstep; - if (pC0) - { - pC0++; - pC1++; - } - if (pC && broadcast_type_C == 4) - pC++; } outptr0 += 2; } @@ -5220,6 +5753,7 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, } _f0 = __msa_fadd_w(_f0, _c0); _f1 = __msa_fadd_w(_f1, _c1); + pC0 += 8; } } if (alpha != 1.f) @@ -5245,8 +5779,6 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, *(int*)(outptr + out_hstep * 7) = __msa_copy_s_w((v4i32)_f1, 3); } outptr += out_hstep * 8; - if (pC0) - pC0 += 8; } for (; jj + 3 < max_jj; jj += 4) { @@ -5262,6 +5794,7 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, if (beta != 1.f) _c = __msa_fmul_w(_c, __msa_fill_w_f32(beta)); _f0 = __msa_fadd_w(_f0, _c); + pC0 += 4; } } if (alpha != 1.f) @@ -5278,8 +5811,6 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, *(int*)(outptr + out_hstep * 3) = __msa_copy_s_w((v4i32)_f0, 3); } outptr += out_hstep * 4; - if (pC0) - pC0 += 4; } for (; jj + 1 < max_jj; jj += 2) { @@ -5301,6 +5832,7 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, if (beta != 1.f) _c = __msa_fmul_w(_c, __msa_fill_w_f32(beta)); _f0 = __msa_fadd_w(_f0, _c); + pC0 += 2; } } if (alpha != 1.f) @@ -5315,8 +5847,6 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, *(int*)(outptr + out_hstep) = __msa_copy_s_w((v4i32)_f0, 1); } outptr += out_hstep * 2; - if (pC0) - pC0 += 2; } #endif // __mips_msa for (; jj + 1 < max_jj; jj += 2) @@ -5342,6 +5872,7 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, } sum0 += c0; sum1 += c1; + pC0 += 2; } } if (alpha != 1.f) @@ -5352,8 +5883,6 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, outptr[0] = sum0; outptr[out_hstep] = sum1; outptr += out_hstep * 2; - if (pC0) - pC0 += 2; } for (; jj < max_jj; jj++) { @@ -5366,23 +5895,23 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, { c = pC0[0]; if (beta != 1.f) c *= beta; + pC0++; } sum0 += c; } if (alpha != 1.f) sum0 *= alpha; outptr[0] = sum0; outptr += out_hstep; - if (pC0) - pC0++; } outptr0++; } } -static void get_optimal_tile_mnk_wq_int8(int M, int N, int K, int constant_TILE_M, int constant_TILE_N, int constant_TILE_K, int& TILE_M, int& TILE_N, int& TILE_K, int nT) +static void get_optimal_tile_mnk_wq_int8(int M, int N, int K, int block_size, int constant_TILE_M, int constant_TILE_N, int constant_TILE_K, int& TILE_M, int& TILE_N, int& TILE_K, int nT) { // resolve optimal tile size from cache size const size_t l2_cache_size = get_cpu_level2_cache_size(); + const int l2_cache_size_int8 = (int)(l2_cache_size / sizeof(signed char)); const int tile_size = std::max(1, (int)((float)l2_cache_size / 2 / sizeof(signed char) / std::max(1, K))); @@ -5390,13 +5919,33 @@ static void get_optimal_tile_mnk_wq_int8(int M, int N, int K, int constant_TILE_ const int tile_m_align = 8; const int tile_n_align = 8; #else - const int tile_m_align = 4; + const int tile_m_align = 2; const int tile_n_align = 2; #endif // one driver M tile follows the natural producer slab TILE_M = tile_m_align; TILE_N = std::max(tile_n_align, tile_size / tile_n_align * tile_n_align); - TILE_K = K; + +#if __mips_msa + int tile_k = (l2_cache_size_int8 - 16) / 8; +#else + int tile_k = (l2_cache_size_int8 - 2) / 3; +#endif + TILE_K = std::max(block_size, tile_k / block_size * block_size); + + if (K > 0) + { + if (TILE_K >= K) + { + TILE_K = K; + } + else + { + const int nn_K = (K + TILE_K - 1) / TILE_K; + tile_k = (K + nn_K - 1) / nn_K; + TILE_K = std::max(block_size, tile_k / block_size * block_size); + } + } if (N > 0) { @@ -5408,8 +5957,14 @@ static void get_optimal_tile_mnk_wq_int8(int M, int N, int K, int constant_TILE_ if (constant_TILE_N > 0) TILE_N = (constant_TILE_N + tile_n_align - 1) / tile_n_align * tile_n_align; + if (constant_TILE_K > 0) + { + TILE_K = std::max(block_size, constant_TILE_K / block_size * block_size); + if (K > 0) + TILE_K = std::min(TILE_K, K); + } + (void)M; (void)constant_TILE_M; - (void)constant_TILE_K; (void)nT; } diff --git a/src/layer/multiheadattention.cpp b/src/layer/multiheadattention.cpp index ae32564d4717..01311277d36e 100644 --- a/src/layer/multiheadattention.cpp +++ b/src/layer/multiheadattention.cpp @@ -597,8 +597,7 @@ static void mha_weight_block_quantize_activation_row_int8(const Mat& A, int tran continue; } - volatile double scale_fp64 = 127.0 / (double)absmax; - const float scale = (float)scale_fp64; + const float scale = 127.f / absmax; descale_ptr[g] = absmax / 127.f; for (int kk = 0; kk < max_kk; kk++) @@ -606,11 +605,7 @@ static void mha_weight_block_quantize_activation_row_int8(const Mat& A, int tran const int k = k0 + kk; float v = transA ? ((const float*)A)[k * A_hstep + i] : ptrA[k]; if (input_scale_ptr) - { v *= input_scale_ptr[k]; - volatile float v_ordered = v; - v = v_ordered; - } outptr[k] = mha_weight_block_quantize_float2int8(v * scale); } } diff --git a/src/layer/riscv/gemm_riscv.cpp b/src/layer/riscv/gemm_riscv.cpp index ea5bf48477ca..06cf0216f9f8 100644 --- a/src/layer/riscv/gemm_riscv.cpp +++ b/src/layer/riscv/gemm_riscv.cpp @@ -1879,13 +1879,13 @@ static int gemm_BT_riscv_wq_int8(const Mat& A, const Mat& packed_B, const Mat& p const int M = transA ? A.w : (A.dims == 3 ? A.c : A.h) * A.elempack; const int block_count = (K + block_size - 1) / block_size; int TILE_M, TILE_N, TILE_K; - get_optimal_tile_mnk_wq_int8(M, N, K, constant_TILE_M, constant_TILE_N, constant_TILE_K, TILE_M, TILE_N, TILE_K, nT); + get_optimal_tile_mnk_wq_int8(M, N, K, block_size, constant_TILE_M, constant_TILE_N, constant_TILE_K, TILE_M, TILE_N, TILE_K, nT); - (void)TILE_K; const int mr = std::min(M, TILE_M); const int nr = std::min(N, TILE_N); const int nn_M = (M + TILE_M - 1) / TILE_M; const int nn_N = (N + TILE_N - 1) / TILE_N; + const int nn_K = (K + TILE_K - 1) / TILE_K; const float* input_scale_ptr = input_scales; Mat BT = packed_B.reshape(K, N); Mat BT_descales = packed_B_descales.reshape(block_count, N); @@ -1901,19 +1901,27 @@ static int gemm_BT_riscv_wq_int8(const Mat& A, const Mat& packed_B, const Mat& p if (AT.empty() || AT_descales.empty()) return -100; + const int nn_MK = nn_M * nn_K; #pragma omp parallel for num_threads(nT) - for (int ppi = 0; ppi < nn_M; ppi++) + for (int ppik = 0; ppik < nn_MK; ppik++) { + const int ppi = ppik / nn_K; + const int ppk = ppik % nn_K; + const int i = ppi * TILE_M; + const int k = ppk * TILE_K; const int max_ii = std::min(M - i, TILE_M); + const int max_kk = std::min(K - k, TILE_K); - Mat AT_tile = AT.channel(i / TILE_M).row_range(0, max_ii); - Mat AT_descales_tile = AT_descales.channel(i / TILE_M).row_range(0, max_ii); + Mat AT_channel = AT.channel(i / TILE_M).reshape(K * mr); + Mat AT_descales_channel = AT_descales.channel(i / TILE_M).reshape(block_count * mr); + Mat AT_tile = AT_channel.range(k * mr, max_kk * mr); + Mat AT_descales_tile = AT_descales_channel.range(k / block_size * mr, (max_kk + block_size - 1) / block_size * mr); if (transA) - transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, block_size, input_scale_ptr); + transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, k, max_kk, block_size, input_scale_ptr); else - quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, block_size, input_scale_ptr); + quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, k, max_kk, block_size, input_scale_ptr); } const int nn_MN = nn_M * nn_N; @@ -1928,13 +1936,20 @@ static int gemm_BT_riscv_wq_int8(const Mat& A, const Mat& packed_B, const Mat& p const int max_ii = std::min(M - i, TILE_M); const int max_jj = std::min(N - j, TILE_N); - Mat AT_tile = AT.channel(i / TILE_M).row_range(0, max_ii); - Mat AT_descales_tile = AT_descales.channel(i / TILE_M).row_range(0, max_ii); Mat BT_tile = BT.row_range(j, max_jj); Mat BT_descales_tile = BT_descales.row_range(j, max_jj); Mat topT_tile = topT.channel(get_omp_thread_num()); + Mat AT_channel = AT.channel(i / TILE_M).reshape(K * mr); + Mat AT_descales_channel = AT_descales.channel(i / TILE_M).reshape(block_count * mr); + + for (int k = 0; k < K; k += TILE_K) + { + const int max_kk = std::min(K - k, TILE_K); + Mat AT_tile = AT_channel.range(k * mr, max_kk * mr); + Mat AT_descales_tile = AT_descales_channel.range(k / block_size * mr, (max_kk + block_size - 1) / block_size * mr); - gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, block_size); + gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, k, max_kk, K, block_size); + } if (output_transpose) transpose_unpack_output_tile_wq_int8(topT_tile, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, N, alpha, beta); @@ -1955,21 +1970,32 @@ static int gemm_BT_riscv_wq_int8(const Mat& A, const Mat& packed_B, const Mat& p const int i = ppi * TILE_M; const int max_ii = std::min(M - i, TILE_M); - Mat AT_tile = ATX.channel(get_omp_thread_num()).row_range(0, max_ii); - Mat AT_descales_tile = ATX_descales.channel(get_omp_thread_num()).row_range(0, max_ii); Mat topT_tile = topT.channel(get_omp_thread_num()); - - if (transA) - transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, block_size, input_scale_ptr); - else - quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, block_size, input_scale_ptr); + Mat ATX_channel = ATX.channel(get_omp_thread_num()).reshape(K * mr); + Mat ATX_descales_channel = ATX_descales.channel(get_omp_thread_num()).reshape(block_count * mr); for (int j = 0; j < N; j += TILE_N) { const int max_jj = std::min(N - j, TILE_N); Mat BT_tile = BT.row_range(j, max_jj); Mat BT_descales_tile = BT_descales.row_range(j, max_jj); - gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, block_size); + + for (int k = 0; k < K; k += TILE_K) + { + const int max_kk = std::min(K - k, TILE_K); + Mat AT_tile = ATX_channel.range(k * mr, max_kk * mr); + Mat AT_descales_tile = ATX_descales_channel.range(k / block_size * mr, (max_kk + block_size - 1) / block_size * mr); + + if (j == 0) + { + if (transA) + transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, k, max_kk, block_size, input_scale_ptr); + else + quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, k, max_kk, block_size, input_scale_ptr); + } + + gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, k, max_kk, K, block_size); + } if (output_transpose) transpose_unpack_output_tile_wq_int8(topT_tile, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, N, alpha, beta); diff --git a/src/layer/riscv/gemm_wq_int8.h b/src/layer/riscv/gemm_wq_int8.h index 9edd10147703..ce05ea98103b 100644 --- a/src/layer/riscv/gemm_wq_int8.h +++ b/src/layer/riscv/gemm_wq_int8.h @@ -34,6 +34,19 @@ static int pack_B_wq_int8(const Mat& B, const Mat& B_scales, Mat& packed_B, Mat& int jj = 0; for (; jj + 3 < max_jj; jj += 4) { + const signed char* p0 = B.row(j + jj); + const signed char* p1 = B.row(j + jj + 1); + const signed char* p2 = B.row(j + jj + 2); + const signed char* p3 = B.row(j + jj + 3); + const float* ps0 = B_scales.row(j + jj); + const float* ps1 = B_scales.row(j + jj + 1); + const float* ps2 = B_scales.row(j + jj + 2); + const float* ps3 = B_scales.row(j + jj + 3); +#if __riscv_vector + const size_t vl = __riscv_vsetvl_e8mf4(4); + const ptrdiff_t B_stride = (ptrdiff_t)B.w; +#endif + for (int g = 0; g < block_count; g++) { const int k0 = g * block_size; @@ -42,18 +55,12 @@ static int pack_B_wq_int8(const Mat& B, const Mat& B_scales, Mat& packed_B, Mat& int kk = 0; for (; kk + 3 < max_kk; kk += 4) { - const signed char* p0 = B.row(j + jj) + k0 + kk; #if __riscv_vector - const size_t vl = __riscv_vsetvl_e8mf4(4); - const ptrdiff_t B_stride = (ptrdiff_t)B.w; __riscv_vse8_v_i8mf4(pp, __riscv_vlse8_v_i8mf4(p0, B_stride, vl), vl); __riscv_vse8_v_i8mf4(pp + 4, __riscv_vlse8_v_i8mf4(p0 + 1, B_stride, vl), vl); __riscv_vse8_v_i8mf4(pp + 8, __riscv_vlse8_v_i8mf4(p0 + 2, B_stride, vl), vl); __riscv_vse8_v_i8mf4(pp + 12, __riscv_vlse8_v_i8mf4(p0 + 3, B_stride, vl), vl); #else - const signed char* p1 = B.row(j + jj + 1) + k0 + kk; - const signed char* p2 = B.row(j + jj + 2) + k0 + kk; - const signed char* p3 = B.row(j + jj + 3) + k0 + kk; pp[0] = p0[0]; pp[1] = p0[1]; pp[2] = p0[2]; @@ -71,20 +78,18 @@ static int pack_B_wq_int8(const Mat& B, const Mat& B_scales, Mat& packed_B, Mat& pp[14] = p3[2]; pp[15] = p3[3]; #endif + p0 += 4; + p1 += 4; + p2 += 4; + p3 += 4; pp += 16; } for (; kk + 1 < max_kk; kk += 2) { - const signed char* p0 = B.row(j + jj) + k0 + kk; #if __riscv_vector - const size_t vl = __riscv_vsetvl_e8mf4(4); - const ptrdiff_t B_stride = (ptrdiff_t)B.w; __riscv_vse8_v_i8mf4(pp, __riscv_vlse8_v_i8mf4(p0, B_stride, vl), vl); __riscv_vse8_v_i8mf4(pp + 4, __riscv_vlse8_v_i8mf4(p0 + 1, B_stride, vl), vl); #else - const signed char* p1 = B.row(j + jj + 1) + k0 + kk; - const signed char* p2 = B.row(j + jj + 2) + k0 + kk; - const signed char* p3 = B.row(j + jj + 3) + k0 + kk; pp[0] = p0[0]; pp[1] = p0[1]; pp[2] = p1[0]; @@ -94,26 +99,39 @@ static int pack_B_wq_int8(const Mat& B, const Mat& B_scales, Mat& packed_B, Mat& pp[6] = p3[0]; pp[7] = p3[1]; #endif + p0 += 2; + p1 += 2; + p2 += 2; + p3 += 2; pp += 8; } for (; kk < max_kk; kk++) { - pp[0] = B.row(j + jj)[k0 + kk]; - pp[1] = B.row(j + jj + 1)[k0 + kk]; - pp[2] = B.row(j + jj + 2)[k0 + kk]; - pp[3] = B.row(j + jj + 3)[k0 + kk]; + pp[0] = *p0++; + pp[1] = *p1++; + pp[2] = *p2++; + pp[3] = *p3++; pp += 4; } - pd[0] = 1.f / B_scales.row(j + jj)[g]; - pd[1] = 1.f / B_scales.row(j + jj + 1)[g]; - pd[2] = 1.f / B_scales.row(j + jj + 2)[g]; - pd[3] = 1.f / B_scales.row(j + jj + 3)[g]; + pd[0] = 1.f / *ps0++; + pd[1] = 1.f / *ps1++; + pd[2] = 1.f / *ps2++; + pd[3] = 1.f / *ps3++; pd += 4; } } for (; jj + 1 < max_jj; jj += 2) { + const signed char* p0 = B.row(j + jj); + const signed char* p1 = B.row(j + jj + 1); + const float* ps0 = B_scales.row(j + jj); + const float* ps1 = B_scales.row(j + jj + 1); +#if __riscv_vector + const size_t vl = __riscv_vsetvl_e8mf4(2); + const ptrdiff_t B_stride = (ptrdiff_t)B.w; +#endif + for (int g = 0; g < block_count; g++) { const int k0 = g * block_size; @@ -122,16 +140,12 @@ static int pack_B_wq_int8(const Mat& B, const Mat& B_scales, Mat& packed_B, Mat& int kk = 0; for (; kk + 3 < max_kk; kk += 4) { - const signed char* p0 = B.row(j + jj) + k0 + kk; #if __riscv_vector - const size_t vl = __riscv_vsetvl_e8mf4(2); - const ptrdiff_t B_stride = (ptrdiff_t)B.w; __riscv_vse8_v_i8mf4(pp, __riscv_vlse8_v_i8mf4(p0, B_stride, vl), vl); __riscv_vse8_v_i8mf4(pp + 2, __riscv_vlse8_v_i8mf4(p0 + 1, B_stride, vl), vl); __riscv_vse8_v_i8mf4(pp + 4, __riscv_vlse8_v_i8mf4(p0 + 2, B_stride, vl), vl); __riscv_vse8_v_i8mf4(pp + 6, __riscv_vlse8_v_i8mf4(p0 + 3, B_stride, vl), vl); #else - const signed char* p1 = B.row(j + jj + 1) + k0 + kk; pp[0] = p0[0]; pp[1] = p0[1]; pp[2] = p0[2]; @@ -141,47 +155,48 @@ static int pack_B_wq_int8(const Mat& B, const Mat& B_scales, Mat& packed_B, Mat& pp[6] = p1[2]; pp[7] = p1[3]; #endif + p0 += 4; + p1 += 4; pp += 8; } for (; kk + 1 < max_kk; kk += 2) { - const signed char* p0 = B.row(j + jj) + k0 + kk; #if __riscv_vector - const size_t vl = __riscv_vsetvl_e8mf4(2); - const ptrdiff_t B_stride = (ptrdiff_t)B.w; __riscv_vse8_v_i8mf4(pp, __riscv_vlse8_v_i8mf4(p0, B_stride, vl), vl); __riscv_vse8_v_i8mf4(pp + 2, __riscv_vlse8_v_i8mf4(p0 + 1, B_stride, vl), vl); #else - const signed char* p1 = B.row(j + jj + 1) + k0 + kk; pp[0] = p0[0]; pp[1] = p0[1]; pp[2] = p1[0]; pp[3] = p1[1]; #endif + p0 += 2; + p1 += 2; pp += 4; } for (; kk < max_kk; kk++) { - pp[0] = B.row(j + jj)[k0 + kk]; - pp[1] = B.row(j + jj + 1)[k0 + kk]; + pp[0] = *p0++; + pp[1] = *p1++; pp += 2; } - pd[0] = 1.f / B_scales.row(j + jj)[g]; - pd[1] = 1.f / B_scales.row(j + jj + 1)[g]; + pd[0] = 1.f / *ps0++; + pd[1] = 1.f / *ps1++; pd += 2; } } for (; jj < max_jj; jj++) { + const signed char* p0 = B.row(j + jj); + const float* ps0 = B_scales.row(j + jj); + for (int g = 0; g < block_count; g++) { - const int k0 = g * block_size; - const int max_kk = std::min(K - k0, block_size); - const signed char* p0 = B.row(j + jj) + k0; + const int max_kk = std::min(K - g * block_size, block_size); for (int kk = 0; kk < max_kk; kk++) - *pp++ = p0[kk]; - *pd++ = 1.f / B_scales.row(j + jj)[g]; + *pp++ = *p0++; + *pd++ = 1.f / *ps0++; } } } @@ -190,13 +205,15 @@ static int pack_B_wq_int8(const Mat& B, const Mat& B_scales, Mat& packed_B, Mat& } // K-major, row-interleaved MR-packn/MR2/MR1 -static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int block_size, const float* input_scale_ptr) +static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int k, int max_kk, int block_size, const float* input_scale_ptr) { signed char* pp = AT_tile; float* pd = AT_descales_tile; - const int K = AT_tile.w; - const int block_count = AT_descales_tile.w; + const int K = max_kk; + const int block_count = (max_kk + block_size - 1) / block_size; const size_t A_hstep = A.dims == 3 ? A.cstep : (size_t)A.w; + const float* A_data = (const float*)A + k; + input_scale_ptr = input_scale_ptr ? input_scale_ptr + k : 0; int ii = 0; #if __riscv_vector @@ -205,19 +222,20 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales const ptrdiff_t A_stride = (ptrdiff_t)A_hstep * sizeof(float); for (; ii + (packn - 1) < max_ii; ii += packn) { - const float* p0 = (const float*)A + (size_t)(i + ii) * A_hstep; + const float* p0 = A_data + (size_t)(i + ii) * A_hstep; + const float* p0g = p0; + const float* psg = input_scale_ptr; for (int g = 0; g < block_count; g++) { - const int k0 = g * block_size; - const int max_kk = std::min(K - k0, block_size); + const int max_kk = std::min(K - g * block_size, block_size); vfloat32m1_t _absmax = __riscv_vfmv_v_f_f32m1(0.f, vl); for (int kk = 0; kk < max_kk; kk++) { - vfloat32m1_t _v = __riscv_vlse32_v_f32m1(p0 + k0 + kk, A_stride, vl); - if (input_scale_ptr) - _v = __riscv_vfmul_vf_f32m1(_v, input_scale_ptr[k0 + kk], vl); + vfloat32m1_t _v = __riscv_vlse32_v_f32m1(p0g + kk, A_stride, vl); + if (psg) + _v = __riscv_vfmul_vf_f32m1(_v, psg[kk], vl); _absmax = __riscv_vfmax_vv_f32m1(_absmax, __riscv_vfabs_v_f32m1(_v, vl), vl); } @@ -228,16 +246,19 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales for (int kk = 0; kk < max_kk; kk++) { - vfloat32m1_t _v = __riscv_vlse32_v_f32m1(p0 + k0 + kk, A_stride, vl); - if (input_scale_ptr) - _v = __riscv_vfmul_vf_f32m1(_v, input_scale_ptr[k0 + kk], vl); + vfloat32m1_t _v = __riscv_vlse32_v_f32m1(p0g + kk, A_stride, vl); + if (psg) + _v = __riscv_vfmul_vf_f32m1(_v, psg[kk], vl); vint32m1_t _v32 = __riscv_vfcvt_x_f_v_i32m1_rm(__riscv_vfmul_vv_f32m1(_v, _scale, vl), __RISCV_FRM_RMM, vl); _v32 = __riscv_vmax_vx_i32m1(_v32, -127, vl); _v32 = __riscv_vmin_vx_i32m1(_v32, 127, vl); - const vint16mf2_t _v16 = __riscv_vnclip_wx_i16mf2(_v32, 0, __RISCV_VXRM_RNU, vl); + vint16mf2_t _v16 = __riscv_vnclip_wx_i16mf2(_v32, 0, __RISCV_VXRM_RNU, vl); __riscv_vse8_v_i8mf4(pp, __riscv_vnclip_wx_i8mf4(_v16, 0, __RISCV_VXRM_RNU, vl), vl); pp += packn; } + p0g += max_kk; + if (psg) + psg += max_kk; } } #endif @@ -245,13 +266,15 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales { const int i0 = i + ii; const int i1 = i + ii + 1; - const float* p0 = (const float*)A + (size_t)i0 * A_hstep; - const float* p1 = (const float*)A + (size_t)i1 * A_hstep; + const float* p0 = A_data + (size_t)i0 * A_hstep; + const float* p1 = A_data + (size_t)i1 * A_hstep; + const float* p0g = p0; + const float* p1g = p1; + const float* psg = input_scale_ptr; for (int g = 0; g < block_count; g++) { - const int k0 = g * block_size; - const int max_kk = std::min(K - k0, block_size); + const int max_kk = std::min(K - g * block_size, block_size); float absmax0 = 0.f; float absmax1 = 0.f; @@ -260,11 +283,11 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales while (kk < max_kk) { const size_t vl = __riscv_vsetvl_e32m8(max_kk - kk); - vfloat32m8_t _v0 = __riscv_vle32_v_f32m8(p0 + k0 + kk, vl); - vfloat32m8_t _v1 = __riscv_vle32_v_f32m8(p1 + k0 + kk, vl); - if (input_scale_ptr) + vfloat32m8_t _v0 = __riscv_vle32_v_f32m8(p0g + kk, vl); + vfloat32m8_t _v1 = __riscv_vle32_v_f32m8(p1g + kk, vl); + if (psg) { - const vfloat32m8_t _s = __riscv_vle32_v_f32m8(input_scale_ptr + k0 + kk, vl); + vfloat32m8_t _s = __riscv_vle32_v_f32m8(psg + kk, vl); _v0 = __riscv_vfmul_vv_f32m8(_v0, _s, vl); _v1 = __riscv_vfmul_vv_f32m8(_v1, _s, vl); } @@ -277,23 +300,20 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales #else for (int kk = 0; kk < max_kk; kk++) { - const int k = k0 + kk; - float v0 = p0[k]; - float v1 = p1[k]; - if (input_scale_ptr) + float v0 = p0g[kk]; + float v1 = p1g[kk]; + if (psg) { - v0 *= input_scale_ptr[k]; - v1 *= input_scale_ptr[k]; + v0 *= psg[kk]; + v1 *= psg[kk]; } absmax0 = std::max(absmax0, fabsf(v0)); absmax1 = std::max(absmax1, fabsf(v1)); } #endif - volatile double scale0_fp64 = absmax0 == 0.f ? 0.0 : 127.0 / (double)absmax0; - volatile double scale1_fp64 = absmax1 == 0.f ? 0.0 : 127.0 / (double)absmax1; - const float scale0 = (float)scale0_fp64; - const float scale1 = (float)scale1_fp64; + const float scale0 = absmax0 == 0.f ? 0.f : 127.f / absmax0; + const float scale1 = absmax1 == 0.f ? 0.f : 127.f / absmax1; pd[0] = absmax0 / 127.f; pd[1] = absmax1 / 127.f; pd += 2; @@ -303,16 +323,16 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales while (kk < max_kk) { const size_t vl = __riscv_vsetvl_e32m8(max_kk - kk); - vfloat32m8_t _v0 = __riscv_vle32_v_f32m8(p0 + k0 + kk, vl); - vfloat32m8_t _v1 = __riscv_vle32_v_f32m8(p1 + k0 + kk, vl); - if (input_scale_ptr) + vfloat32m8_t _v0 = __riscv_vle32_v_f32m8(p0g + kk, vl); + vfloat32m8_t _v1 = __riscv_vle32_v_f32m8(p1g + kk, vl); + if (psg) { - const vfloat32m8_t _s = __riscv_vle32_v_f32m8(input_scale_ptr + k0 + kk, vl); + vfloat32m8_t _s = __riscv_vle32_v_f32m8(psg + kk, vl); _v0 = __riscv_vfmul_vv_f32m8(_v0, _s, vl); _v1 = __riscv_vfmul_vv_f32m8(_v1, _s, vl); } - const vint8m2_t _q0 = float2int8(__riscv_vfmul_vf_f32m8(_v0, scale0, vl), vl); - const vint8m2_t _q1 = float2int8(__riscv_vfmul_vf_f32m8(_v1, scale1, vl), vl); + vint8m2_t _q0 = float2int8(__riscv_vfmul_vf_f32m8(_v0, scale0, vl), vl); + vint8m2_t _q1 = float2int8(__riscv_vfmul_vf_f32m8(_v1, scale1, vl), vl); __riscv_vsse8_v_i8m2(pp, 2, _q0, vl); __riscv_vsse8_v_i8m2(pp + 1, 2, _q1, vl); pp += vl * 2; @@ -321,32 +341,34 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales #else for (int kk = 0; kk < max_kk; kk++) { - const int k = k0 + kk; - float v0 = p0[k]; - float v1 = p1[k]; - if (input_scale_ptr) + float v0 = p0g[kk]; + float v1 = p1g[kk]; + if (psg) { - v0 *= input_scale_ptr[k]; - v1 *= input_scale_ptr[k]; - asm volatile("" - : "+f"(v0), "+f"(v1)); + v0 *= psg[kk]; + v1 *= psg[kk]; } pp[0] = float2int8(v0 * scale0); pp[1] = float2int8(v1 * scale1); pp += 2; } #endif + p0g += max_kk; + p1g += max_kk; + if (psg) + psg += max_kk; } } for (; ii < max_ii; ii++) { const int i0 = i + ii; - const float* p0 = (const float*)A + (size_t)i0 * A_hstep; + const float* p0 = A_data + (size_t)i0 * A_hstep; + const float* p0g = p0; + const float* psg = input_scale_ptr; for (int g = 0; g < block_count; g++) { - const int k0 = g * block_size; - const int max_kk = std::min(K - k0, block_size); + const int max_kk = std::min(K - g * block_size, block_size); float absmax = 0.f; #if __riscv_vector @@ -354,9 +376,9 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales while (kk < max_kk) { const size_t vl = __riscv_vsetvl_e32m8(max_kk - kk); - vfloat32m8_t _v = __riscv_vle32_v_f32m8(p0 + k0 + kk, vl); - if (input_scale_ptr) - _v = __riscv_vfmul_vv_f32m8(_v, __riscv_vle32_v_f32m8(input_scale_ptr + k0 + kk, vl), vl); + vfloat32m8_t _v = __riscv_vle32_v_f32m8(p0g + kk, vl); + if (psg) + _v = __riscv_vfmul_vv_f32m8(_v, __riscv_vle32_v_f32m8(psg + kk, vl), vl); _v = __riscv_vfabs_v_f32m8(_v, vl); absmax = __riscv_vfmv_f_s_f32m1_f32(__riscv_vfredmax_vs_f32m8_f32m1(_v, __riscv_vfmv_s_f_f32m1(absmax, 1), vl)); kk += vl; @@ -364,16 +386,14 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales #else for (int kk = 0; kk < max_kk; kk++) { - const int k = k0 + kk; - float v = p0[k]; - if (input_scale_ptr) - v *= input_scale_ptr[k]; + float v = p0g[kk]; + if (psg) + v *= psg[kk]; absmax = std::max(absmax, fabsf(v)); } #endif - volatile double scale_fp64 = absmax == 0.f ? 0.0 : 127.0 / (double)absmax; - const float scale = (float)scale_fp64; + const float scale = absmax == 0.f ? 0.f : 127.f / absmax; *pd++ = absmax / 127.f; #if __riscv_vector @@ -381,9 +401,9 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales while (kk < max_kk) { const size_t vl = __riscv_vsetvl_e32m8(max_kk - kk); - vfloat32m8_t _v = __riscv_vle32_v_f32m8(p0 + k0 + kk, vl); - if (input_scale_ptr) - _v = __riscv_vfmul_vv_f32m8(_v, __riscv_vle32_v_f32m8(input_scale_ptr + k0 + kk, vl), vl); + vfloat32m8_t _v = __riscv_vle32_v_f32m8(p0g + kk, vl); + if (psg) + _v = __riscv_vfmul_vv_f32m8(_v, __riscv_vle32_v_f32m8(psg + kk, vl), vl); __riscv_vse8_v_i8m2(pp, float2int8(__riscv_vfmul_vf_f32m8(_v, scale, vl), vl), vl); pp += vl; kk += vl; @@ -391,29 +411,31 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales #else for (int kk = 0; kk < max_kk; kk++) { - const int k = k0 + kk; - float v = p0[k]; - if (input_scale_ptr) + float v = p0g[kk]; + if (psg) { - v *= input_scale_ptr[k]; - asm volatile("" - : "+f"(v)); + v *= psg[kk]; } *pp++ = float2int8(v * scale); } #endif + p0g += max_kk; + if (psg) + psg += max_kk; } } } // K-major, row-interleaved MR-packn/MR2/MR1 -static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int block_size, const float* input_scale_ptr) +static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int k, int max_kk, int block_size, const float* input_scale_ptr) { signed char* pp = AT_tile; float* pd = AT_descales_tile; - const int K = AT_tile.w; - const int block_count = AT_descales_tile.w; + const int K = max_kk; + const int block_count = (max_kk + block_size - 1) / block_size; const size_t A_hstep = A.dims == 3 ? A.cstep : (size_t)A.w; + const float* A_data = (const float*)A + (size_t)k * A_hstep; + input_scale_ptr = input_scale_ptr ? input_scale_ptr + k : 0; int ii = 0; #if __riscv_vector @@ -422,61 +444,71 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A for (; ii + (packn - 1) < max_ii; ii += packn) { const int i0 = i + ii; + const float* p0g = A_data + i0; + const float* psg = input_scale_ptr; for (int g = 0; g < block_count; g++) { - const int k0 = g * block_size; - const int max_kk = std::min(K - k0, block_size); + const int max_kk = std::min(K - g * block_size, block_size); vfloat32m1_t _absmax = __riscv_vfmv_v_f_f32m1(0.f, vl); + const float* pAk = p0g; for (int kk = 0; kk < max_kk; kk++) { - vfloat32m1_t _v = __riscv_vle32_v_f32m1((const float*)A + (size_t)(k0 + kk) * A_hstep + i0, vl); - if (input_scale_ptr) - _v = __riscv_vfmul_vf_f32m1(_v, input_scale_ptr[k0 + kk], vl); + vfloat32m1_t _v = __riscv_vle32_v_f32m1(pAk, vl); + if (psg) + _v = __riscv_vfmul_vf_f32m1(_v, psg[kk], vl); _absmax = __riscv_vfmax_vv_f32m1(_absmax, __riscv_vfabs_v_f32m1(_v, vl), vl); + pAk += A_hstep; } vfloat32m1_t _scale = __riscv_vfrdiv_vf_f32m1(_absmax, 127.f, vl); _scale = __riscv_vfmerge_vfm_f32m1(_scale, 0.f, __riscv_vmfeq_vf_f32m1_b32(_absmax, 0.f, vl), vl); __riscv_vse32_v_f32m1(pd, __riscv_vfmul_vf_f32m1(_absmax, 1.f / 127.f, vl), vl); pd += packn; + pAk = p0g; for (int kk = 0; kk < max_kk; kk++) { - vfloat32m1_t _v = __riscv_vle32_v_f32m1((const float*)A + (size_t)(k0 + kk) * A_hstep + i0, vl); - if (input_scale_ptr) - _v = __riscv_vfmul_vf_f32m1(_v, input_scale_ptr[k0 + kk], vl); + vfloat32m1_t _v = __riscv_vle32_v_f32m1(pAk, vl); + if (psg) + _v = __riscv_vfmul_vf_f32m1(_v, psg[kk], vl); vint32m1_t _v32 = __riscv_vfcvt_x_f_v_i32m1_rm(__riscv_vfmul_vv_f32m1(_v, _scale, vl), __RISCV_FRM_RMM, vl); _v32 = __riscv_vmax_vx_i32m1(_v32, -127, vl); _v32 = __riscv_vmin_vx_i32m1(_v32, 127, vl); - const vint16mf2_t _v16 = __riscv_vnclip_wx_i16mf2(_v32, 0, __RISCV_VXRM_RNU, vl); + vint16mf2_t _v16 = __riscv_vnclip_wx_i16mf2(_v32, 0, __RISCV_VXRM_RNU, vl); __riscv_vse8_v_i8mf4(pp, __riscv_vnclip_wx_i8mf4(_v16, 0, __RISCV_VXRM_RNU, vl), vl); pp += packn; + pAk += A_hstep; } + p0g += (size_t)max_kk * A_hstep; + if (psg) + psg += max_kk; } } #endif for (; ii + 1 < max_ii; ii += 2) { const int i0 = i + ii; + const float* p0g = A_data + i0; + const float* psg = input_scale_ptr; for (int g = 0; g < block_count; g++) { - const int k0 = g * block_size; - const int max_kk = std::min(K - k0, block_size); + const int max_kk = std::min(K - g * block_size, block_size); float absmax0 = 0.f; float absmax1 = 0.f; #if __riscv_vector int kk = 0; + const float* pAk = p0g; while (kk < max_kk) { const size_t vl = __riscv_vsetvl_e32m8(max_kk - kk); - vfloat32m8_t _v0 = __riscv_vlse32_v_f32m8((const float*)A + (size_t)(k0 + kk) * A_hstep + i0, (ptrdiff_t)A_hstep * sizeof(float), vl); - vfloat32m8_t _v1 = __riscv_vlse32_v_f32m8((const float*)A + (size_t)(k0 + kk) * A_hstep + i0 + 1, (ptrdiff_t)A_hstep * sizeof(float), vl); - if (input_scale_ptr) + vfloat32m8_t _v0 = __riscv_vlse32_v_f32m8(pAk, (ptrdiff_t)A_hstep * sizeof(float), vl); + vfloat32m8_t _v1 = __riscv_vlse32_v_f32m8(pAk + 1, (ptrdiff_t)A_hstep * sizeof(float), vl); + if (psg) { - const vfloat32m8_t _s = __riscv_vle32_v_f32m8(input_scale_ptr + k0 + kk, vl); + vfloat32m8_t _s = __riscv_vle32_v_f32m8(psg + kk, vl); _v0 = __riscv_vfmul_vv_f32m8(_v0, _s, vl); _v1 = __riscv_vfmul_vv_f32m8(_v1, _s, vl); } @@ -484,167 +516,351 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A _v1 = __riscv_vfabs_v_f32m8(_v1, vl); absmax0 = __riscv_vfmv_f_s_f32m1_f32(__riscv_vfredmax_vs_f32m8_f32m1(_v0, __riscv_vfmv_s_f_f32m1(absmax0, 1), vl)); absmax1 = __riscv_vfmv_f_s_f32m1_f32(__riscv_vfredmax_vs_f32m8_f32m1(_v1, __riscv_vfmv_s_f_f32m1(absmax1, 1), vl)); + pAk += vl * A_hstep; kk += vl; } #else + const float* pAk = p0g; for (int kk = 0; kk < max_kk; kk++) { - const int k = k0 + kk; - const float* p0 = (const float*)A + (size_t)k * A_hstep + i0; - float v0 = p0[0]; - float v1 = p0[1]; - if (input_scale_ptr) + float v0 = pAk[0]; + float v1 = pAk[1]; + if (psg) { - const float s = input_scale_ptr[k]; + const float s = psg[kk]; v0 *= s; v1 *= s; } absmax0 = std::max(absmax0, fabsf(v0)); absmax1 = std::max(absmax1, fabsf(v1)); + pAk += A_hstep; } #endif - volatile double scale0_fp64 = absmax0 == 0.f ? 0.0 : 127.0 / (double)absmax0; - volatile double scale1_fp64 = absmax1 == 0.f ? 0.0 : 127.0 / (double)absmax1; - const float scale0 = (float)scale0_fp64; - const float scale1 = (float)scale1_fp64; + const float scale0 = absmax0 == 0.f ? 0.f : 127.f / absmax0; + const float scale1 = absmax1 == 0.f ? 0.f : 127.f / absmax1; pd[0] = absmax0 / 127.f; pd[1] = absmax1 / 127.f; pd += 2; #if __riscv_vector kk = 0; + pAk = p0g; while (kk < max_kk) { const size_t vl = __riscv_vsetvl_e32m8(max_kk - kk); - vfloat32m8_t _v0 = __riscv_vlse32_v_f32m8((const float*)A + (size_t)(k0 + kk) * A_hstep + i0, (ptrdiff_t)A_hstep * sizeof(float), vl); - vfloat32m8_t _v1 = __riscv_vlse32_v_f32m8((const float*)A + (size_t)(k0 + kk) * A_hstep + i0 + 1, (ptrdiff_t)A_hstep * sizeof(float), vl); - if (input_scale_ptr) + vfloat32m8_t _v0 = __riscv_vlse32_v_f32m8(pAk, (ptrdiff_t)A_hstep * sizeof(float), vl); + vfloat32m8_t _v1 = __riscv_vlse32_v_f32m8(pAk + 1, (ptrdiff_t)A_hstep * sizeof(float), vl); + if (psg) { - const vfloat32m8_t _s = __riscv_vle32_v_f32m8(input_scale_ptr + k0 + kk, vl); + vfloat32m8_t _s = __riscv_vle32_v_f32m8(psg + kk, vl); _v0 = __riscv_vfmul_vv_f32m8(_v0, _s, vl); _v1 = __riscv_vfmul_vv_f32m8(_v1, _s, vl); } - const vint8m2_t _q0 = float2int8(__riscv_vfmul_vf_f32m8(_v0, scale0, vl), vl); - const vint8m2_t _q1 = float2int8(__riscv_vfmul_vf_f32m8(_v1, scale1, vl), vl); + vint8m2_t _q0 = float2int8(__riscv_vfmul_vf_f32m8(_v0, scale0, vl), vl); + vint8m2_t _q1 = float2int8(__riscv_vfmul_vf_f32m8(_v1, scale1, vl), vl); __riscv_vsse8_v_i8m2(pp, 2, _q0, vl); __riscv_vsse8_v_i8m2(pp + 1, 2, _q1, vl); pp += vl * 2; + pAk += vl * A_hstep; kk += vl; } #else + pAk = p0g; for (int kk = 0; kk < max_kk; kk++) { - const int k = k0 + kk; - const float* p0 = (const float*)A + (size_t)k * A_hstep + i0; - float v0 = p0[0]; - float v1 = p0[1]; - if (input_scale_ptr) + float v0 = pAk[0]; + float v1 = pAk[1]; + if (psg) { - const float s = input_scale_ptr[k]; + const float s = psg[kk]; v0 *= s; v1 *= s; - asm volatile("" - : "+f"(v0), "+f"(v1)); } pp[0] = float2int8(v0 * scale0); pp[1] = float2int8(v1 * scale1); pp += 2; + pAk += A_hstep; } #endif + p0g += (size_t)max_kk * A_hstep; + if (psg) + psg += max_kk; } } for (; ii < max_ii; ii++) { const int i0 = i + ii; + const float* p0g = A_data + i0; + const float* psg = input_scale_ptr; for (int g = 0; g < block_count; g++) { - const int k0 = g * block_size; - const int max_kk = std::min(K - k0, block_size); + const int max_kk = std::min(K - g * block_size, block_size); float absmax = 0.f; #if __riscv_vector int kk = 0; + const float* pAk = p0g; while (kk < max_kk) { const size_t vl = __riscv_vsetvl_e32m8(max_kk - kk); - vfloat32m8_t _v = __riscv_vlse32_v_f32m8((const float*)A + (size_t)(k0 + kk) * A_hstep + i0, (ptrdiff_t)A_hstep * sizeof(float), vl); - if (input_scale_ptr) - _v = __riscv_vfmul_vv_f32m8(_v, __riscv_vle32_v_f32m8(input_scale_ptr + k0 + kk, vl), vl); + vfloat32m8_t _v = __riscv_vlse32_v_f32m8(pAk, (ptrdiff_t)A_hstep * sizeof(float), vl); + if (psg) + _v = __riscv_vfmul_vv_f32m8(_v, __riscv_vle32_v_f32m8(psg + kk, vl), vl); _v = __riscv_vfabs_v_f32m8(_v, vl); absmax = __riscv_vfmv_f_s_f32m1_f32(__riscv_vfredmax_vs_f32m8_f32m1(_v, __riscv_vfmv_s_f_f32m1(absmax, 1), vl)); + pAk += vl * A_hstep; kk += vl; } #else + const float* pAk = p0g; for (int kk = 0; kk < max_kk; kk++) { - const int k = k0 + kk; - float v = ((const float*)A)[(size_t)k * A_hstep + i0]; - if (input_scale_ptr) - v *= input_scale_ptr[k]; + float v = *pAk; + if (psg) + v *= psg[kk]; absmax = std::max(absmax, fabsf(v)); + pAk += A_hstep; } #endif - volatile double scale_fp64 = absmax == 0.f ? 0.0 : 127.0 / (double)absmax; - const float scale = (float)scale_fp64; + const float scale = absmax == 0.f ? 0.f : 127.f / absmax; *pd++ = absmax / 127.f; #if __riscv_vector kk = 0; + pAk = p0g; while (kk < max_kk) { const size_t vl = __riscv_vsetvl_e32m8(max_kk - kk); - vfloat32m8_t _v = __riscv_vlse32_v_f32m8((const float*)A + (size_t)(k0 + kk) * A_hstep + i0, (ptrdiff_t)A_hstep * sizeof(float), vl); - if (input_scale_ptr) - _v = __riscv_vfmul_vv_f32m8(_v, __riscv_vle32_v_f32m8(input_scale_ptr + k0 + kk, vl), vl); + vfloat32m8_t _v = __riscv_vlse32_v_f32m8(pAk, (ptrdiff_t)A_hstep * sizeof(float), vl); + if (psg) + _v = __riscv_vfmul_vv_f32m8(_v, __riscv_vle32_v_f32m8(psg + kk, vl), vl); __riscv_vse8_v_i8m2(pp, float2int8(__riscv_vfmul_vf_f32m8(_v, scale, vl), vl), vl); pp += vl; + pAk += vl * A_hstep; kk += vl; } #else + pAk = p0g; for (int kk = 0; kk < max_kk; kk++) { - const int k = k0 + kk; - float v = ((const float*)A)[(size_t)k * A_hstep + i0]; - if (input_scale_ptr) + float v = *pAk; + if (psg) { - v *= input_scale_ptr[k]; - asm volatile("" - : "+f"(v)); + v *= psg[kk]; } *pp++ = float2int8(v * scale); + pAk += A_hstep; } #endif + p0g += (size_t)max_kk * A_hstep; + if (psg) + psg += max_kk; } } } -static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int block_size) +static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int k0, int max_kk0, int B_hstep, int block_size) { const signed char* pAT = AT_tile; const float* pAT_descales = AT_descales_tile; const signed char* pBT = BT_tile; const float* pBT_descales = BT_descales_tile; float* outptr = topT_tile; - const int K = AT_tile.w; - const int block_count = AT_descales_tile.w; + const int K = max_kk0; + const int block_count = (max_kk0 + block_size - 1) / block_size; + const int num_blocks = (B_hstep + block_size - 1) / block_size; + const int block_start = k0 / block_size; int ii = 0; #if __riscv_vector const int packn = csrr_vlenb() / 4; const size_t vl4 = __riscv_vsetvl_e8mf4(4); - for (; ii < max_ii;) + for (; ii + (packn - 1) < max_ii; ii += packn) { const signed char* pB = pBT; const float* pB_descales = pBT_descales; - const int mr = ii + (packn - 1) < max_ii ? packn : ii + 1 < max_ii ? 2 : 1; + const int mr = packn; const size_t vl = __riscv_vsetvl_e32m1(mr); int jj = 0; + for (; jj + 7 < max_jj; jj += 8) + { + vfloat32m1_t _fsum0 = __riscv_vfmv_v_f_f32m1(0.f, vl); + vfloat32m1_t _fsum1 = __riscv_vfmv_v_f_f32m1(0.f, vl); + vfloat32m1_t _fsum2 = __riscv_vfmv_v_f_f32m1(0.f, vl); + vfloat32m1_t _fsum3 = __riscv_vfmv_v_f_f32m1(0.f, vl); + vfloat32m1_t _fsum4 = __riscv_vfmv_v_f_f32m1(0.f, vl); + vfloat32m1_t _fsum5 = __riscv_vfmv_v_f_f32m1(0.f, vl); + vfloat32m1_t _fsum6 = __riscv_vfmv_v_f_f32m1(0.f, vl); + vfloat32m1_t _fsum7 = __riscv_vfmv_v_f_f32m1(0.f, vl); + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + const signed char* pB0 = pB + (size_t)4 * k0; + const signed char* pB1 = pB + (size_t)4 * B_hstep + (size_t)4 * k0; + const float* pB_descales0 = pB_descales + (size_t)4 * block_start; + const float* pB_descales1 = pB_descales + (size_t)4 * num_blocks + (size_t)4 * block_start; + for (int k = 0; k < K; k += block_size) + { + const int max_kk = std::min(K - k, block_size); + vint32m1_t _sum0 = __riscv_vmv_v_x_i32m1(0, vl); + vint32m1_t _sum1 = __riscv_vmv_v_x_i32m1(0, vl); + vint32m1_t _sum2 = __riscv_vmv_v_x_i32m1(0, vl); + vint32m1_t _sum3 = __riscv_vmv_v_x_i32m1(0, vl); + vint32m1_t _sum4 = __riscv_vmv_v_x_i32m1(0, vl); + vint32m1_t _sum5 = __riscv_vmv_v_x_i32m1(0, vl); + vint32m1_t _sum6 = __riscv_vmv_v_x_i32m1(0, vl); + vint32m1_t _sum7 = __riscv_vmv_v_x_i32m1(0, vl); + + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + vint16mf2_t _a = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pA, vl), vl); + uint32_t b0 = *(const uint32_t*)pB0; + uint32_t b1 = *(const uint32_t*)pB1; + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, (signed char)b0, _a, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, (signed char)(b0 >> 8), _a, vl); + _sum2 = __riscv_vwmacc_vx_i32m1(_sum2, (signed char)(b0 >> 16), _a, vl); + _sum3 = __riscv_vwmacc_vx_i32m1(_sum3, (signed char)(b0 >> 24), _a, vl); + _sum4 = __riscv_vwmacc_vx_i32m1(_sum4, (signed char)b1, _a, vl); + _sum5 = __riscv_vwmacc_vx_i32m1(_sum5, (signed char)(b1 >> 8), _a, vl); + _sum6 = __riscv_vwmacc_vx_i32m1(_sum6, (signed char)(b1 >> 16), _a, vl); + _sum7 = __riscv_vwmacc_vx_i32m1(_sum7, (signed char)(b1 >> 24), _a, vl); + _a = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pA + mr, vl), vl); + b0 = *(const uint32_t*)(pB0 + 4); + b1 = *(const uint32_t*)(pB1 + 4); + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, (signed char)b0, _a, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, (signed char)(b0 >> 8), _a, vl); + _sum2 = __riscv_vwmacc_vx_i32m1(_sum2, (signed char)(b0 >> 16), _a, vl); + _sum3 = __riscv_vwmacc_vx_i32m1(_sum3, (signed char)(b0 >> 24), _a, vl); + _sum4 = __riscv_vwmacc_vx_i32m1(_sum4, (signed char)b1, _a, vl); + _sum5 = __riscv_vwmacc_vx_i32m1(_sum5, (signed char)(b1 >> 8), _a, vl); + _sum6 = __riscv_vwmacc_vx_i32m1(_sum6, (signed char)(b1 >> 16), _a, vl); + _sum7 = __riscv_vwmacc_vx_i32m1(_sum7, (signed char)(b1 >> 24), _a, vl); + _a = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pA + mr * 2, vl), vl); + b0 = *(const uint32_t*)(pB0 + 8); + b1 = *(const uint32_t*)(pB1 + 8); + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, (signed char)b0, _a, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, (signed char)(b0 >> 8), _a, vl); + _sum2 = __riscv_vwmacc_vx_i32m1(_sum2, (signed char)(b0 >> 16), _a, vl); + _sum3 = __riscv_vwmacc_vx_i32m1(_sum3, (signed char)(b0 >> 24), _a, vl); + _sum4 = __riscv_vwmacc_vx_i32m1(_sum4, (signed char)b1, _a, vl); + _sum5 = __riscv_vwmacc_vx_i32m1(_sum5, (signed char)(b1 >> 8), _a, vl); + _sum6 = __riscv_vwmacc_vx_i32m1(_sum6, (signed char)(b1 >> 16), _a, vl); + _sum7 = __riscv_vwmacc_vx_i32m1(_sum7, (signed char)(b1 >> 24), _a, vl); + _a = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pA + mr * 3, vl), vl); + b0 = *(const uint32_t*)(pB0 + 12); + b1 = *(const uint32_t*)(pB1 + 12); + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, (signed char)b0, _a, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, (signed char)(b0 >> 8), _a, vl); + _sum2 = __riscv_vwmacc_vx_i32m1(_sum2, (signed char)(b0 >> 16), _a, vl); + _sum3 = __riscv_vwmacc_vx_i32m1(_sum3, (signed char)(b0 >> 24), _a, vl); + _sum4 = __riscv_vwmacc_vx_i32m1(_sum4, (signed char)b1, _a, vl); + _sum5 = __riscv_vwmacc_vx_i32m1(_sum5, (signed char)(b1 >> 8), _a, vl); + _sum6 = __riscv_vwmacc_vx_i32m1(_sum6, (signed char)(b1 >> 16), _a, vl); + _sum7 = __riscv_vwmacc_vx_i32m1(_sum7, (signed char)(b1 >> 24), _a, vl); + pA += mr * 4; + pB0 += 16; + pB1 += 16; + } + for (; kk + 1 < max_kk; kk += 2) + { + vint16mf2_t _a = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pA, vl), vl); + uint32_t b0 = *(const uint32_t*)pB0; + uint32_t b1 = *(const uint32_t*)pB1; + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, (signed char)b0, _a, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, (signed char)(b0 >> 8), _a, vl); + _sum2 = __riscv_vwmacc_vx_i32m1(_sum2, (signed char)(b0 >> 16), _a, vl); + _sum3 = __riscv_vwmacc_vx_i32m1(_sum3, (signed char)(b0 >> 24), _a, vl); + _sum4 = __riscv_vwmacc_vx_i32m1(_sum4, (signed char)b1, _a, vl); + _sum5 = __riscv_vwmacc_vx_i32m1(_sum5, (signed char)(b1 >> 8), _a, vl); + _sum6 = __riscv_vwmacc_vx_i32m1(_sum6, (signed char)(b1 >> 16), _a, vl); + _sum7 = __riscv_vwmacc_vx_i32m1(_sum7, (signed char)(b1 >> 24), _a, vl); + _a = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pA + mr, vl), vl); + b0 = *(const uint32_t*)(pB0 + 4); + b1 = *(const uint32_t*)(pB1 + 4); + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, (signed char)b0, _a, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, (signed char)(b0 >> 8), _a, vl); + _sum2 = __riscv_vwmacc_vx_i32m1(_sum2, (signed char)(b0 >> 16), _a, vl); + _sum3 = __riscv_vwmacc_vx_i32m1(_sum3, (signed char)(b0 >> 24), _a, vl); + _sum4 = __riscv_vwmacc_vx_i32m1(_sum4, (signed char)b1, _a, vl); + _sum5 = __riscv_vwmacc_vx_i32m1(_sum5, (signed char)(b1 >> 8), _a, vl); + _sum6 = __riscv_vwmacc_vx_i32m1(_sum6, (signed char)(b1 >> 16), _a, vl); + _sum7 = __riscv_vwmacc_vx_i32m1(_sum7, (signed char)(b1 >> 24), _a, vl); + pA += mr * 2; + pB0 += 8; + pB1 += 8; + } + for (; kk < max_kk; kk++) + { + vint16mf2_t _a = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pA, vl), vl); + const uint32_t b0 = *(const uint32_t*)pB0; + const uint32_t b1 = *(const uint32_t*)pB1; + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, (signed char)b0, _a, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, (signed char)(b0 >> 8), _a, vl); + _sum2 = __riscv_vwmacc_vx_i32m1(_sum2, (signed char)(b0 >> 16), _a, vl); + _sum3 = __riscv_vwmacc_vx_i32m1(_sum3, (signed char)(b0 >> 24), _a, vl); + _sum4 = __riscv_vwmacc_vx_i32m1(_sum4, (signed char)b1, _a, vl); + _sum5 = __riscv_vwmacc_vx_i32m1(_sum5, (signed char)(b1 >> 8), _a, vl); + _sum6 = __riscv_vwmacc_vx_i32m1(_sum6, (signed char)(b1 >> 16), _a, vl); + _sum7 = __riscv_vwmacc_vx_i32m1(_sum7, (signed char)(b1 >> 24), _a, vl); + pA += mr; + pB0 += 4; + pB1 += 4; + } + + vfloat32m1_t _ad = __riscv_vle32_v_f32m1(pA_descales, vl); + vfloat32m1_t _v = __riscv_vfmul_vv_f32m1(__riscv_vfcvt_f_x_v_f32m1(_sum0, vl), _ad, vl); + _fsum0 = __riscv_vfmacc_vf_f32m1(_fsum0, pB_descales0[0], _v, vl); + _v = __riscv_vfmul_vv_f32m1(__riscv_vfcvt_f_x_v_f32m1(_sum1, vl), _ad, vl); + _fsum1 = __riscv_vfmacc_vf_f32m1(_fsum1, pB_descales0[1], _v, vl); + _v = __riscv_vfmul_vv_f32m1(__riscv_vfcvt_f_x_v_f32m1(_sum2, vl), _ad, vl); + _fsum2 = __riscv_vfmacc_vf_f32m1(_fsum2, pB_descales0[2], _v, vl); + _v = __riscv_vfmul_vv_f32m1(__riscv_vfcvt_f_x_v_f32m1(_sum3, vl), _ad, vl); + _fsum3 = __riscv_vfmacc_vf_f32m1(_fsum3, pB_descales0[3], _v, vl); + _v = __riscv_vfmul_vv_f32m1(__riscv_vfcvt_f_x_v_f32m1(_sum4, vl), _ad, vl); + _fsum4 = __riscv_vfmacc_vf_f32m1(_fsum4, pB_descales1[0], _v, vl); + _v = __riscv_vfmul_vv_f32m1(__riscv_vfcvt_f_x_v_f32m1(_sum5, vl), _ad, vl); + _fsum5 = __riscv_vfmacc_vf_f32m1(_fsum5, pB_descales1[1], _v, vl); + _v = __riscv_vfmul_vv_f32m1(__riscv_vfcvt_f_x_v_f32m1(_sum6, vl), _ad, vl); + _fsum6 = __riscv_vfmacc_vf_f32m1(_fsum6, pB_descales1[2], _v, vl); + _v = __riscv_vfmul_vv_f32m1(__riscv_vfcvt_f_x_v_f32m1(_sum7, vl), _ad, vl); + _fsum7 = __riscv_vfmacc_vf_f32m1(_fsum7, pB_descales1[3], _v, vl); + pA_descales += mr; + pB_descales0 += 4; + pB_descales1 += 4; + } + + if (k0 != 0) + { + _fsum0 = __riscv_vfadd_vv_f32m1(_fsum0, __riscv_vle32_v_f32m1(outptr, vl), vl); + _fsum1 = __riscv_vfadd_vv_f32m1(_fsum1, __riscv_vle32_v_f32m1(outptr + mr, vl), vl); + _fsum2 = __riscv_vfadd_vv_f32m1(_fsum2, __riscv_vle32_v_f32m1(outptr + mr * 2, vl), vl); + _fsum3 = __riscv_vfadd_vv_f32m1(_fsum3, __riscv_vle32_v_f32m1(outptr + mr * 3, vl), vl); + _fsum4 = __riscv_vfadd_vv_f32m1(_fsum4, __riscv_vle32_v_f32m1(outptr + mr * 4, vl), vl); + _fsum5 = __riscv_vfadd_vv_f32m1(_fsum5, __riscv_vle32_v_f32m1(outptr + mr * 5, vl), vl); + _fsum6 = __riscv_vfadd_vv_f32m1(_fsum6, __riscv_vle32_v_f32m1(outptr + mr * 6, vl), vl); + _fsum7 = __riscv_vfadd_vv_f32m1(_fsum7, __riscv_vle32_v_f32m1(outptr + mr * 7, vl), vl); + } + __riscv_vse32_v_f32m1(outptr, _fsum0, vl); + __riscv_vse32_v_f32m1(outptr + mr, _fsum1, vl); + __riscv_vse32_v_f32m1(outptr + mr * 2, _fsum2, vl); + __riscv_vse32_v_f32m1(outptr + mr * 3, _fsum3, vl); + __riscv_vse32_v_f32m1(outptr + mr * 4, _fsum4, vl); + __riscv_vse32_v_f32m1(outptr + mr * 5, _fsum5, vl); + __riscv_vse32_v_f32m1(outptr + mr * 6, _fsum6, vl); + __riscv_vse32_v_f32m1(outptr + mr * 7, _fsum7, vl); + outptr += mr * 8; + pB = pB1 + (size_t)4 * (B_hstep - k0 - max_kk0); + pB_descales = pB_descales1 + (size_t)4 * (num_blocks - block_start - block_count); + } for (; jj + 3 < max_jj; jj += 4) { + pB += (size_t)4 * k0; + pB_descales += (size_t)4 * block_start; vfloat32m1_t _fsum0 = __riscv_vfmv_v_f_f32m1(0.f, vl); vfloat32m1_t _fsum1 = __riscv_vfmv_v_f_f32m1(0.f, vl); vfloat32m1_t _fsum2 = __riscv_vfmv_v_f_f32m1(0.f, vl); @@ -709,7 +925,7 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } for (; kk < max_kk; kk++) { - const vint16mf2_t _a = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pA, vl), vl); + vint16mf2_t _a = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pA, vl), vl); const uint32_t b = *(const uint32_t*)pB; _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, (signed char)b, _a, vl); _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, (signed char)(b >> 8), _a, vl); @@ -719,7 +935,7 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB += 4; } - const vfloat32m1_t _ad = __riscv_vle32_v_f32m1(pA_descales, vl); + vfloat32m1_t _ad = __riscv_vle32_v_f32m1(pA_descales, vl); vfloat32m1_t _v = __riscv_vfmul_vv_f32m1(__riscv_vfcvt_f_x_v_f32m1(_sum0, vl), _ad, vl); _fsum0 = __riscv_vfmacc_vf_f32m1(_fsum0, pB_descales[0], _v, vl); _v = __riscv_vfmul_vv_f32m1(__riscv_vfcvt_f_x_v_f32m1(_sum1, vl), _ad, vl); @@ -732,14 +948,25 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB_descales += 4; } + if (k0 != 0) + { + _fsum0 = __riscv_vfadd_vv_f32m1(_fsum0, __riscv_vle32_v_f32m1(outptr, vl), vl); + _fsum1 = __riscv_vfadd_vv_f32m1(_fsum1, __riscv_vle32_v_f32m1(outptr + mr, vl), vl); + _fsum2 = __riscv_vfadd_vv_f32m1(_fsum2, __riscv_vle32_v_f32m1(outptr + mr * 2, vl), vl); + _fsum3 = __riscv_vfadd_vv_f32m1(_fsum3, __riscv_vle32_v_f32m1(outptr + mr * 3, vl), vl); + } __riscv_vse32_v_f32m1(outptr, _fsum0, vl); __riscv_vse32_v_f32m1(outptr + mr, _fsum1, vl); __riscv_vse32_v_f32m1(outptr + mr * 2, _fsum2, vl); __riscv_vse32_v_f32m1(outptr + mr * 3, _fsum3, vl); outptr += mr * 4; + pB += (size_t)4 * (B_hstep - k0 - max_kk0); + pB_descales += (size_t)4 * (num_blocks - block_start - block_count); } for (; jj + 1 < max_jj; jj += 2) { + pB += (size_t)2 * k0; + pB_descales += (size_t)2 * block_start; vfloat32m1_t _fsum0 = __riscv_vfmv_v_f_f32m1(0.f, vl); vfloat32m1_t _fsum1 = __riscv_vfmv_v_f_f32m1(0.f, vl); @@ -788,7 +1015,7 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } for (; kk < max_kk; kk++) { - const vint16mf2_t _a = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pA, vl), vl); + vint16mf2_t _a = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pA, vl), vl); const uint16_t b = *(const uint16_t*)pB; _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, (signed char)b, _a, vl); _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, (signed char)(b >> 8), _a, vl); @@ -796,7 +1023,7 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB += 2; } - const vfloat32m1_t _ad = __riscv_vle32_v_f32m1(pA_descales, vl); + vfloat32m1_t _ad = __riscv_vle32_v_f32m1(pA_descales, vl); vfloat32m1_t _v = __riscv_vfmul_vv_f32m1(__riscv_vfcvt_f_x_v_f32m1(_sum0, vl), _ad, vl); _fsum0 = __riscv_vfmacc_vf_f32m1(_fsum0, pB_descales[0], _v, vl); _v = __riscv_vfmul_vv_f32m1(__riscv_vfcvt_f_x_v_f32m1(_sum1, vl), _ad, vl); @@ -805,12 +1032,21 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB_descales += 2; } + if (k0 != 0) + { + _fsum0 = __riscv_vfadd_vv_f32m1(_fsum0, __riscv_vle32_v_f32m1(outptr, vl), vl); + _fsum1 = __riscv_vfadd_vv_f32m1(_fsum1, __riscv_vle32_v_f32m1(outptr + mr, vl), vl); + } __riscv_vse32_v_f32m1(outptr, _fsum0, vl); __riscv_vse32_v_f32m1(outptr + mr, _fsum1, vl); outptr += mr * 2; + pB += (size_t)2 * (B_hstep - k0 - max_kk0); + pB_descales += (size_t)2 * (num_blocks - block_start - block_count); } for (; jj < max_jj; jj++) { + pB += k0; + pB_descales += block_start; vfloat32m1_t _fsum = __riscv_vfmv_v_f_f32m1(0.f, vl); const signed char* pA = pAT; @@ -854,226 +1090,259 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } for (; kk < max_kk; kk++) { - const vint16mf2_t _a = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pA, vl), vl); + vint16mf2_t _a = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pA, vl), vl); _sum = __riscv_vwmacc_vx_i32m1(_sum, pB[0], _a, vl); pA += mr; pB++; } - const vfloat32m1_t _ad = __riscv_vle32_v_f32m1(pA_descales, vl); - const vfloat32m1_t _v = __riscv_vfmul_vv_f32m1(__riscv_vfcvt_f_x_v_f32m1(_sum, vl), _ad, vl); + vfloat32m1_t _ad = __riscv_vle32_v_f32m1(pA_descales, vl); + vfloat32m1_t _v = __riscv_vfmul_vv_f32m1(__riscv_vfcvt_f_x_v_f32m1(_sum, vl), _ad, vl); _fsum = __riscv_vfmacc_vf_f32m1(_fsum, pB_descales[0], _v, vl); pA_descales += mr; pB_descales++; } + if (k0 != 0) + _fsum = __riscv_vfadd_vv_f32m1(_fsum, __riscv_vle32_v_f32m1(outptr, vl), vl); __riscv_vse32_v_f32m1(outptr, _fsum, vl); outptr += mr; + pB += B_hstep - k0 - max_kk0; + pB_descales += num_blocks - block_start - block_count; } pAT += K * mr; pAT_descales += block_count * mr; - ii += mr; } -#else for (; ii + 1 < max_ii; ii += 2) { const signed char* pB = pBT; const float* pB_descales = pBT_descales; + int jj = 0; for (; jj + 3 < max_jj; jj += 4) { - float sum00 = 0.f; - float sum01 = 0.f; - float sum02 = 0.f; - float sum03 = 0.f; - float sum10 = 0.f; - float sum11 = 0.f; - float sum12 = 0.f; - float sum13 = 0.f; + pB += (size_t)4 * k0; + pB_descales += (size_t)4 * block_start; + const size_t vl = __riscv_vsetvl_e32m1(4); + vfloat32m1_t _fsum0 = __riscv_vfmv_v_f_f32m1(0.f, vl); + vfloat32m1_t _fsum1 = __riscv_vfmv_v_f_f32m1(0.f, vl); const signed char* pA = pAT; const float* pA_descales = pAT_descales; for (int k = 0; k < K; k += block_size) { const int max_kk = std::min(K - k, block_size); - int s00 = 0, s01 = 0, s02 = 0, s03 = 0; - int s10 = 0, s11 = 0, s12 = 0, s13 = 0; + vint32m1_t _sum0 = __riscv_vmv_v_x_i32m1(0, vl); + vint32m1_t _sum1 = __riscv_vmv_v_x_i32m1(0, vl); int kk = 0; for (; kk + 3 < max_kk; kk += 4) { - const signed char* b0 = pB; - const signed char* b1 = b0 + 4; - const signed char* b2 = b1 + 4; - const signed char* b3 = b2 + 4; - s00 += pA[0] * b0[0] + pA[2] * b0[1] + pA[4] * b0[2] + pA[6] * b0[3]; - s01 += pA[0] * b1[0] + pA[2] * b1[1] + pA[4] * b1[2] + pA[6] * b1[3]; - s02 += pA[0] * b2[0] + pA[2] * b2[1] + pA[4] * b2[2] + pA[6] * b2[3]; - s03 += pA[0] * b3[0] + pA[2] * b3[1] + pA[4] * b3[2] + pA[6] * b3[3]; - s10 += pA[1] * b0[0] + pA[3] * b0[1] + pA[5] * b0[2] + pA[7] * b0[3]; - s11 += pA[1] * b1[0] + pA[3] * b1[1] + pA[5] * b1[2] + pA[7] * b1[3]; - s12 += pA[1] * b2[0] + pA[3] * b2[1] + pA[5] * b2[2] + pA[7] * b2[3]; - s13 += pA[1] * b3[0] + pA[3] * b3[1] + pA[5] * b3[2] + pA[7] * b3[3]; + vint16mf2_t _b0 = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB, vl), vl); + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, pA[0], _b0, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, pA[1], _b0, vl); + _b0 = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB + 4, vl), vl); + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, pA[2], _b0, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, pA[3], _b0, vl); + _b0 = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB + 8, vl), vl); + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, pA[4], _b0, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, pA[5], _b0, vl); + _b0 = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB + 12, vl), vl); + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, pA[6], _b0, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, pA[7], _b0, vl); pA += 8; pB += 16; } for (; kk + 1 < max_kk; kk += 2) { - const signed char* b0 = pB; - const signed char* b1 = b0 + 2; - const signed char* b2 = b1 + 2; - const signed char* b3 = b2 + 2; - s00 += pA[0] * b0[0] + pA[2] * b0[1]; - s01 += pA[0] * b1[0] + pA[2] * b1[1]; - s02 += pA[0] * b2[0] + pA[2] * b2[1]; - s03 += pA[0] * b3[0] + pA[2] * b3[1]; - s10 += pA[1] * b0[0] + pA[3] * b0[1]; - s11 += pA[1] * b1[0] + pA[3] * b1[1]; - s12 += pA[1] * b2[0] + pA[3] * b2[1]; - s13 += pA[1] * b3[0] + pA[3] * b3[1]; + vint16mf2_t _b0 = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB, vl), vl); + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, pA[0], _b0, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, pA[1], _b0, vl); + _b0 = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB + 4, vl), vl); + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, pA[2], _b0, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, pA[3], _b0, vl); pA += 4; pB += 8; } for (; kk < max_kk; kk++) { - s00 += pA[0] * pB[0]; - s01 += pA[0] * pB[1]; - s02 += pA[0] * pB[2]; - s03 += pA[0] * pB[3]; - s10 += pA[1] * pB[0]; - s11 += pA[1] * pB[1]; - s12 += pA[1] * pB[2]; - s13 += pA[1] * pB[3]; + vint16mf2_t _b = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB, vl), vl); + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, pA[0], _b, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, pA[1], _b, vl); pA += 2; pB += 4; } - const float ad0 = pA_descales[0]; - const float ad1 = pA_descales[1]; - const float* bd = pB_descales; - sum00 += s00 * ad0 * bd[0]; - sum01 += s01 * ad0 * bd[1]; - sum02 += s02 * ad0 * bd[2]; - sum03 += s03 * ad0 * bd[3]; - sum10 += s10 * ad1 * bd[0]; - sum11 += s11 * ad1 * bd[1]; - sum12 += s12 * ad1 * bd[2]; - sum13 += s13 * ad1 * bd[3]; + vfloat32m1_t _bd = __riscv_vle32_v_f32m1(pB_descales, vl); + vfloat32m1_t _v = __riscv_vfmul_vf_f32m1(__riscv_vfcvt_f_x_v_f32m1(_sum0, vl), pA_descales[0], vl); + _fsum0 = __riscv_vfmacc_vv_f32m1(_fsum0, _bd, _v, vl); + _v = __riscv_vfmul_vf_f32m1(__riscv_vfcvt_f_x_v_f32m1(_sum1, vl), pA_descales[1], vl); + _fsum1 = __riscv_vfmacc_vv_f32m1(_fsum1, _bd, _v, vl); pA_descales += 2; pB_descales += 4; } - outptr[0] = sum00; - outptr[1] = sum10; - outptr[2] = sum01; - outptr[3] = sum11; - outptr[4] = sum02; - outptr[5] = sum12; - outptr[6] = sum03; - outptr[7] = sum13; + if (k0 != 0) + { + vfloat32m1x2_t _s = __riscv_vlseg2e32_v_f32m1x2(outptr, vl); + _fsum0 = __riscv_vfadd_vv_f32m1(_fsum0, __riscv_vget_v_f32m1x2_f32m1(_s, 0), vl); + _fsum1 = __riscv_vfadd_vv_f32m1(_fsum1, __riscv_vget_v_f32m1x2_f32m1(_s, 1), vl); + } + __riscv_vsseg2e32_v_f32m1x2(outptr, __riscv_vcreate_v_f32m1x2(_fsum0, _fsum1), vl); outptr += 8; + pB += (size_t)4 * (B_hstep - k0 - max_kk0); + pB_descales += (size_t)4 * (num_blocks - block_start - block_count); } for (; jj + 1 < max_jj; jj += 2) { - float sum00 = 0.f, sum01 = 0.f, sum10 = 0.f, sum11 = 0.f; + pB += (size_t)2 * k0; + pB_descales += (size_t)2 * block_start; + const size_t vl = __riscv_vsetvl_e32m1(2); + vfloat32m1_t _fsum0 = __riscv_vfmv_v_f_f32m1(0.f, vl); + vfloat32m1_t _fsum1 = __riscv_vfmv_v_f_f32m1(0.f, vl); + const signed char* pA = pAT; const float* pA_descales = pAT_descales; for (int k = 0; k < K; k += block_size) { const int max_kk = std::min(K - k, block_size); - int s00 = 0, s01 = 0, s10 = 0, s11 = 0; + vint32m1_t _sum0 = __riscv_vmv_v_x_i32m1(0, vl); + vint32m1_t _sum1 = __riscv_vmv_v_x_i32m1(0, vl); + int kk = 0; for (; kk + 3 < max_kk; kk += 4) { - const signed char* b0 = pB; - const signed char* b1 = b0 + 4; - s00 += pA[0] * b0[0] + pA[2] * b0[1] + pA[4] * b0[2] + pA[6] * b0[3]; - s01 += pA[0] * b1[0] + pA[2] * b1[1] + pA[4] * b1[2] + pA[6] * b1[3]; - s10 += pA[1] * b0[0] + pA[3] * b0[1] + pA[5] * b0[2] + pA[7] * b0[3]; - s11 += pA[1] * b1[0] + pA[3] * b1[1] + pA[5] * b1[2] + pA[7] * b1[3]; + vint16mf2_t _b0 = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB, vl), vl); + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, pA[0], _b0, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, pA[1], _b0, vl); + _b0 = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB + 2, vl), vl); + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, pA[2], _b0, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, pA[3], _b0, vl); + _b0 = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB + 4, vl), vl); + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, pA[4], _b0, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, pA[5], _b0, vl); + _b0 = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB + 6, vl), vl); + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, pA[6], _b0, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, pA[7], _b0, vl); pA += 8; pB += 8; } for (; kk + 1 < max_kk; kk += 2) { - const signed char* b0 = pB; - const signed char* b1 = b0 + 2; - s00 += pA[0] * b0[0] + pA[2] * b0[1]; - s01 += pA[0] * b1[0] + pA[2] * b1[1]; - s10 += pA[1] * b0[0] + pA[3] * b0[1]; - s11 += pA[1] * b1[0] + pA[3] * b1[1]; + vint16mf2_t _b0 = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB, vl), vl); + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, pA[0], _b0, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, pA[1], _b0, vl); + _b0 = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB + 2, vl), vl); + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, pA[2], _b0, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, pA[3], _b0, vl); pA += 4; pB += 4; } for (; kk < max_kk; kk++) { - s00 += pA[0] * pB[0]; - s01 += pA[0] * pB[1]; - s10 += pA[1] * pB[0]; - s11 += pA[1] * pB[1]; + vint16mf2_t _b = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB, vl), vl); + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, pA[0], _b, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, pA[1], _b, vl); pA += 2; pB += 2; } - const float ad0 = pA_descales[0]; - const float ad1 = pA_descales[1]; - const float* bd = pB_descales; - sum00 += s00 * ad0 * bd[0]; - sum01 += s01 * ad0 * bd[1]; - sum10 += s10 * ad1 * bd[0]; - sum11 += s11 * ad1 * bd[1]; + + vfloat32m1_t _bd = __riscv_vle32_v_f32m1(pB_descales, vl); + vfloat32m1_t _v = __riscv_vfmul_vf_f32m1(__riscv_vfcvt_f_x_v_f32m1(_sum0, vl), pA_descales[0], vl); + _fsum0 = __riscv_vfmacc_vv_f32m1(_fsum0, _bd, _v, vl); + _v = __riscv_vfmul_vf_f32m1(__riscv_vfcvt_f_x_v_f32m1(_sum1, vl), pA_descales[1], vl); + _fsum1 = __riscv_vfmacc_vv_f32m1(_fsum1, _bd, _v, vl); pA_descales += 2; pB_descales += 2; } - outptr[0] = sum00; - outptr[1] = sum10; - outptr[2] = sum01; - outptr[3] = sum11; + + if (k0 != 0) + { + vfloat32m1x2_t _s = __riscv_vlseg2e32_v_f32m1x2(outptr, vl); + _fsum0 = __riscv_vfadd_vv_f32m1(_fsum0, __riscv_vget_v_f32m1x2_f32m1(_s, 0), vl); + _fsum1 = __riscv_vfadd_vv_f32m1(_fsum1, __riscv_vget_v_f32m1x2_f32m1(_s, 1), vl); + } + __riscv_vsseg2e32_v_f32m1x2(outptr, __riscv_vcreate_v_f32m1x2(_fsum0, _fsum1), vl); outptr += 4; + pB += (size_t)2 * (B_hstep - k0 - max_kk0); + pB_descales += (size_t)2 * (num_blocks - block_start - block_count); } for (; jj < max_jj; jj++) { - float sum0 = 0.f, sum1 = 0.f; + pB += k0; + pB_descales += block_start; + const size_t vl = __riscv_vsetvl_e32m1(1); + vfloat32m1_t _fsum0 = __riscv_vfmv_v_f_f32m1(0.f, vl); + vfloat32m1_t _fsum1 = __riscv_vfmv_v_f_f32m1(0.f, vl); + const signed char* pA = pAT; const float* pA_descales = pAT_descales; for (int k = 0; k < K; k += block_size) { const int max_kk = std::min(K - k, block_size); - int s0 = 0, s1 = 0; + vint32m1_t _sum0 = __riscv_vmv_v_x_i32m1(0, vl); + vint32m1_t _sum1 = __riscv_vmv_v_x_i32m1(0, vl); + int kk = 0; for (; kk + 3 < max_kk; kk += 4) { - const signed char* b = pB; - s0 += pA[0] * b[0] + pA[2] * b[1] + pA[4] * b[2] + pA[6] * b[3]; - s1 += pA[1] * b[0] + pA[3] * b[1] + pA[5] * b[2] + pA[7] * b[3]; + vint16mf2_t _b0 = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB, vl), vl); + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, pA[0], _b0, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, pA[1], _b0, vl); + _b0 = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB + 1, vl), vl); + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, pA[2], _b0, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, pA[3], _b0, vl); + _b0 = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB + 2, vl), vl); + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, pA[4], _b0, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, pA[5], _b0, vl); + _b0 = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB + 3, vl), vl); + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, pA[6], _b0, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, pA[7], _b0, vl); pA += 8; pB += 4; } for (; kk + 1 < max_kk; kk += 2) { - const signed char* b = pB; - s0 += pA[0] * b[0] + pA[2] * b[1]; - s1 += pA[1] * b[0] + pA[3] * b[1]; + vint16mf2_t _b0 = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB, vl), vl); + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, pA[0], _b0, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, pA[1], _b0, vl); + _b0 = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB + 1, vl), vl); + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, pA[2], _b0, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, pA[3], _b0, vl); pA += 4; pB += 2; } for (; kk < max_kk; kk++) { - s0 += pA[0] * pB[0]; - s1 += pA[1] * pB[0]; + vint16mf2_t _b = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB, vl), vl); + _sum0 = __riscv_vwmacc_vx_i32m1(_sum0, pA[0], _b, vl); + _sum1 = __riscv_vwmacc_vx_i32m1(_sum1, pA[1], _b, vl); pA += 2; pB++; } - const float bd = pB_descales[0]; - sum0 += s0 * pA_descales[0] * bd; - sum1 += s1 * pA_descales[1] * bd; + + vfloat32m1_t _bd = __riscv_vle32_v_f32m1(pB_descales, vl); + vfloat32m1_t _v = __riscv_vfmul_vf_f32m1(__riscv_vfcvt_f_x_v_f32m1(_sum0, vl), pA_descales[0], vl); + _fsum0 = __riscv_vfmacc_vv_f32m1(_fsum0, _bd, _v, vl); + _v = __riscv_vfmul_vf_f32m1(__riscv_vfcvt_f_x_v_f32m1(_sum1, vl), pA_descales[1], vl); + _fsum1 = __riscv_vfmacc_vv_f32m1(_fsum1, _bd, _v, vl); pA_descales += 2; pB_descales++; } - outptr[0] = sum0; - outptr[1] = sum1; + + if (k0 != 0) + { + vfloat32m1x2_t _s = __riscv_vlseg2e32_v_f32m1x2(outptr, vl); + _fsum0 = __riscv_vfadd_vv_f32m1(_fsum0, __riscv_vget_v_f32m1x2_f32m1(_s, 0), vl); + _fsum1 = __riscv_vfadd_vv_f32m1(_fsum1, __riscv_vget_v_f32m1x2_f32m1(_s, 1), vl); + } + __riscv_vsseg2e32_v_f32m1x2(outptr, __riscv_vcreate_v_f32m1x2(_fsum0, _fsum1), vl); outptr += 2; + pB += B_hstep - k0 - max_kk0; + pB_descales += num_blocks - block_start - block_count; } + pAT += K * 2; pAT_descales += block_count * 2; } @@ -1081,92 +1350,529 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de { const signed char* pB = pBT; const float* pB_descales = pBT_descales; + int jj = 0; for (; jj + 3 < max_jj; jj += 4) { - float sum0 = 0.f, sum1 = 0.f, sum2 = 0.f, sum3 = 0.f; + pB += (size_t)4 * k0; + pB_descales += (size_t)4 * block_start; + const size_t vl = __riscv_vsetvl_e32m1(4); + vfloat32m1_t _fsum = __riscv_vfmv_v_f_f32m1(0.f, vl); + const signed char* pA = pAT; const float* pA_descales = pAT_descales; for (int k = 0; k < K; k += block_size) { const int max_kk = std::min(K - k, block_size); - int s0 = 0, s1 = 0, s2 = 0, s3 = 0; + vint32m1_t _sum = __riscv_vmv_v_x_i32m1(0, vl); + int kk = 0; for (; kk + 3 < max_kk; kk += 4) { - const signed char* b0 = pB; - const signed char* b1 = b0 + 4; - const signed char* b2 = b1 + 4; - const signed char* b3 = b2 + 4; - s0 += pA[0] * b0[0] + pA[1] * b0[1] + pA[2] * b0[2] + pA[3] * b0[3]; - s1 += pA[0] * b1[0] + pA[1] * b1[1] + pA[2] * b1[2] + pA[3] * b1[3]; - s2 += pA[0] * b2[0] + pA[1] * b2[1] + pA[2] * b2[2] + pA[3] * b2[3]; - s3 += pA[0] * b3[0] + pA[1] * b3[1] + pA[2] * b3[2] + pA[3] * b3[3]; + vint16mf2_t _b0 = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB, vl), vl); + _sum = __riscv_vwmacc_vx_i32m1(_sum, pA[0], _b0, vl); + _b0 = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB + 4, vl), vl); + _sum = __riscv_vwmacc_vx_i32m1(_sum, pA[1], _b0, vl); + _b0 = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB + 8, vl), vl); + _sum = __riscv_vwmacc_vx_i32m1(_sum, pA[2], _b0, vl); + _b0 = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB + 12, vl), vl); + _sum = __riscv_vwmacc_vx_i32m1(_sum, pA[3], _b0, vl); pA += 4; pB += 16; } for (; kk + 1 < max_kk; kk += 2) { - const signed char* b0 = pB; - const signed char* b1 = b0 + 2; - const signed char* b2 = b1 + 2; - const signed char* b3 = b2 + 2; - s0 += pA[0] * b0[0] + pA[1] * b0[1]; - s1 += pA[0] * b1[0] + pA[1] * b1[1]; - s2 += pA[0] * b2[0] + pA[1] * b2[1]; - s3 += pA[0] * b3[0] + pA[1] * b3[1]; + vint16mf2_t _b0 = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB, vl), vl); + _sum = __riscv_vwmacc_vx_i32m1(_sum, pA[0], _b0, vl); + _b0 = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB + 4, vl), vl); + _sum = __riscv_vwmacc_vx_i32m1(_sum, pA[1], _b0, vl); pA += 2; pB += 8; } for (; kk < max_kk; kk++) { - s0 += pA[0] * pB[0]; - s1 += pA[0] * pB[1]; - s2 += pA[0] * pB[2]; - s3 += pA[0] * pB[3]; + vint16mf2_t _b = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB, vl), vl); + _sum = __riscv_vwmacc_vx_i32m1(_sum, pA[0], _b, vl); pA++; pB += 4; } - const float ad = pA_descales[0]; - const float* bd = pB_descales; - sum0 += s0 * ad * bd[0]; - sum1 += s1 * ad * bd[1]; - sum2 += s2 * ad * bd[2]; - sum3 += s3 * ad * bd[3]; + + vfloat32m1_t _bd = __riscv_vle32_v_f32m1(pB_descales, vl); + vfloat32m1_t _v = __riscv_vfmul_vf_f32m1(__riscv_vfcvt_f_x_v_f32m1(_sum, vl), pA_descales[0], vl); + _fsum = __riscv_vfmacc_vv_f32m1(_fsum, _bd, _v, vl); pA_descales++; pB_descales += 4; } - outptr[0] = sum0; - outptr[1] = sum1; - outptr[2] = sum2; - outptr[3] = sum3; + + if (k0 != 0) + _fsum = __riscv_vfadd_vv_f32m1(_fsum, __riscv_vle32_v_f32m1(outptr, vl), vl); + __riscv_vse32_v_f32m1(outptr, _fsum, vl); outptr += 4; + pB += (size_t)4 * (B_hstep - k0 - max_kk0); + pB_descales += (size_t)4 * (num_blocks - block_start - block_count); } for (; jj + 1 < max_jj; jj += 2) { - float sum0 = 0.f, sum1 = 0.f; + pB += (size_t)2 * k0; + pB_descales += (size_t)2 * block_start; + const size_t vl = __riscv_vsetvl_e32m1(2); + vfloat32m1_t _fsum = __riscv_vfmv_v_f_f32m1(0.f, vl); + const signed char* pA = pAT; const float* pA_descales = pAT_descales; for (int k = 0; k < K; k += block_size) { const int max_kk = std::min(K - k, block_size); - int s0 = 0, s1 = 0; + vint32m1_t _sum = __riscv_vmv_v_x_i32m1(0, vl); + int kk = 0; for (; kk + 3 < max_kk; kk += 4) { - const signed char* b0 = pB; - const signed char* b1 = b0 + 4; - s0 += pA[0] * b0[0] + pA[1] * b0[1] + pA[2] * b0[2] + pA[3] * b0[3]; - s1 += pA[0] * b1[0] + pA[1] * b1[1] + pA[2] * b1[2] + pA[3] * b1[3]; + vint16mf2_t _b0 = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB, vl), vl); + _sum = __riscv_vwmacc_vx_i32m1(_sum, pA[0], _b0, vl); + _b0 = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB + 2, vl), vl); + _sum = __riscv_vwmacc_vx_i32m1(_sum, pA[1], _b0, vl); + _b0 = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB + 4, vl), vl); + _sum = __riscv_vwmacc_vx_i32m1(_sum, pA[2], _b0, vl); + _b0 = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB + 6, vl), vl); + _sum = __riscv_vwmacc_vx_i32m1(_sum, pA[3], _b0, vl); pA += 4; pB += 8; } for (; kk + 1 < max_kk; kk += 2) { - const signed char* b0 = pB; - const signed char* b1 = b0 + 2; - s0 += pA[0] * b0[0] + pA[1] * b0[1]; - s1 += pA[0] * b1[0] + pA[1] * b1[1]; + vint16mf2_t _b0 = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB, vl), vl); + _sum = __riscv_vwmacc_vx_i32m1(_sum, pA[0], _b0, vl); + _b0 = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB + 2, vl), vl); + _sum = __riscv_vwmacc_vx_i32m1(_sum, pA[1], _b0, vl); + pA += 2; + pB += 4; + } + for (; kk < max_kk; kk++) + { + vint16mf2_t _b = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB, vl), vl); + _sum = __riscv_vwmacc_vx_i32m1(_sum, pA[0], _b, vl); + pA++; + pB += 2; + } + + vfloat32m1_t _bd = __riscv_vle32_v_f32m1(pB_descales, vl); + vfloat32m1_t _v = __riscv_vfmul_vf_f32m1(__riscv_vfcvt_f_x_v_f32m1(_sum, vl), pA_descales[0], vl); + _fsum = __riscv_vfmacc_vv_f32m1(_fsum, _bd, _v, vl); + pA_descales++; + pB_descales += 2; + } + + if (k0 != 0) + _fsum = __riscv_vfadd_vv_f32m1(_fsum, __riscv_vle32_v_f32m1(outptr, vl), vl); + __riscv_vse32_v_f32m1(outptr, _fsum, vl); + outptr += 2; + pB += (size_t)2 * (B_hstep - k0 - max_kk0); + pB_descales += (size_t)2 * (num_blocks - block_start - block_count); + } + for (; jj < max_jj; jj++) + { + pB += k0; + pB_descales += block_start; + const size_t vl = __riscv_vsetvl_e32m1(1); + vfloat32m1_t _fsum = __riscv_vfmv_v_f_f32m1(0.f, vl); + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + const int max_kk = std::min(K - k, block_size); + vint32m1_t _sum = __riscv_vmv_v_x_i32m1(0, vl); + + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + vint16mf2_t _b0 = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB, vl), vl); + _sum = __riscv_vwmacc_vx_i32m1(_sum, pA[0], _b0, vl); + _b0 = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB + 1, vl), vl); + _sum = __riscv_vwmacc_vx_i32m1(_sum, pA[1], _b0, vl); + _b0 = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB + 2, vl), vl); + _sum = __riscv_vwmacc_vx_i32m1(_sum, pA[2], _b0, vl); + _b0 = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB + 3, vl), vl); + _sum = __riscv_vwmacc_vx_i32m1(_sum, pA[3], _b0, vl); + pA += 4; + pB += 4; + } + for (; kk + 1 < max_kk; kk += 2) + { + vint16mf2_t _b0 = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB, vl), vl); + _sum = __riscv_vwmacc_vx_i32m1(_sum, pA[0], _b0, vl); + _b0 = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB + 1, vl), vl); + _sum = __riscv_vwmacc_vx_i32m1(_sum, pA[1], _b0, vl); + pA += 2; + pB += 2; + } + for (; kk < max_kk; kk++) + { + vint16mf2_t _b = __riscv_vsext_vf2_i16mf2(__riscv_vle8_v_i8mf4(pB, vl), vl); + _sum = __riscv_vwmacc_vx_i32m1(_sum, pA[0], _b, vl); + pA++; + pB++; + } + + vfloat32m1_t _bd = __riscv_vle32_v_f32m1(pB_descales, vl); + vfloat32m1_t _v = __riscv_vfmul_vf_f32m1(__riscv_vfcvt_f_x_v_f32m1(_sum, vl), pA_descales[0], vl); + _fsum = __riscv_vfmacc_vv_f32m1(_fsum, _bd, _v, vl); + pA_descales++; + pB_descales++; + } + + if (k0 != 0) + _fsum = __riscv_vfadd_vv_f32m1(_fsum, __riscv_vle32_v_f32m1(outptr, vl), vl); + __riscv_vse32_v_f32m1(outptr, _fsum, vl); + outptr++; + pB += B_hstep - k0 - max_kk0; + pB_descales += num_blocks - block_start - block_count; + } + + pAT += K; + pAT_descales += block_count; + } +#else + for (; ii + 1 < max_ii; ii += 2) + { + const signed char* pB = pBT; + const float* pB_descales = pBT_descales; + int jj = 0; + for (; jj + 3 < max_jj; jj += 4) + { + pB += (size_t)4 * k0; + pB_descales += (size_t)4 * block_start; + float sum00 = 0.f; + float sum01 = 0.f; + float sum02 = 0.f; + float sum03 = 0.f; + float sum10 = 0.f; + float sum11 = 0.f; + float sum12 = 0.f; + float sum13 = 0.f; + + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + const int max_kk = std::min(K - k, block_size); + int s00 = 0, s01 = 0, s02 = 0, s03 = 0; + int s10 = 0, s11 = 0, s12 = 0, s13 = 0; + + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + const signed char* b0 = pB; + const signed char* b1 = b0 + 4; + const signed char* b2 = b1 + 4; + const signed char* b3 = b2 + 4; + s00 += pA[0] * b0[0] + pA[2] * b0[1] + pA[4] * b0[2] + pA[6] * b0[3]; + s01 += pA[0] * b1[0] + pA[2] * b1[1] + pA[4] * b1[2] + pA[6] * b1[3]; + s02 += pA[0] * b2[0] + pA[2] * b2[1] + pA[4] * b2[2] + pA[6] * b2[3]; + s03 += pA[0] * b3[0] + pA[2] * b3[1] + pA[4] * b3[2] + pA[6] * b3[3]; + s10 += pA[1] * b0[0] + pA[3] * b0[1] + pA[5] * b0[2] + pA[7] * b0[3]; + s11 += pA[1] * b1[0] + pA[3] * b1[1] + pA[5] * b1[2] + pA[7] * b1[3]; + s12 += pA[1] * b2[0] + pA[3] * b2[1] + pA[5] * b2[2] + pA[7] * b2[3]; + s13 += pA[1] * b3[0] + pA[3] * b3[1] + pA[5] * b3[2] + pA[7] * b3[3]; + pA += 8; + pB += 16; + } + for (; kk + 1 < max_kk; kk += 2) + { + const signed char* b0 = pB; + const signed char* b1 = b0 + 2; + const signed char* b2 = b1 + 2; + const signed char* b3 = b2 + 2; + s00 += pA[0] * b0[0] + pA[2] * b0[1]; + s01 += pA[0] * b1[0] + pA[2] * b1[1]; + s02 += pA[0] * b2[0] + pA[2] * b2[1]; + s03 += pA[0] * b3[0] + pA[2] * b3[1]; + s10 += pA[1] * b0[0] + pA[3] * b0[1]; + s11 += pA[1] * b1[0] + pA[3] * b1[1]; + s12 += pA[1] * b2[0] + pA[3] * b2[1]; + s13 += pA[1] * b3[0] + pA[3] * b3[1]; + pA += 4; + pB += 8; + } + for (; kk < max_kk; kk++) + { + s00 += pA[0] * pB[0]; + s01 += pA[0] * pB[1]; + s02 += pA[0] * pB[2]; + s03 += pA[0] * pB[3]; + s10 += pA[1] * pB[0]; + s11 += pA[1] * pB[1]; + s12 += pA[1] * pB[2]; + s13 += pA[1] * pB[3]; + pA += 2; + pB += 4; + } + + const float ad0 = pA_descales[0]; + const float ad1 = pA_descales[1]; + const float* bd = pB_descales; + sum00 += s00 * ad0 * bd[0]; + sum01 += s01 * ad0 * bd[1]; + sum02 += s02 * ad0 * bd[2]; + sum03 += s03 * ad0 * bd[3]; + sum10 += s10 * ad1 * bd[0]; + sum11 += s11 * ad1 * bd[1]; + sum12 += s12 * ad1 * bd[2]; + sum13 += s13 * ad1 * bd[3]; + pA_descales += 2; + pB_descales += 4; + } + + if (k0 != 0) + { + sum00 += outptr[0]; + sum10 += outptr[1]; + sum01 += outptr[2]; + sum11 += outptr[3]; + sum02 += outptr[4]; + sum12 += outptr[5]; + sum03 += outptr[6]; + sum13 += outptr[7]; + } + outptr[0] = sum00; + outptr[1] = sum10; + outptr[2] = sum01; + outptr[3] = sum11; + outptr[4] = sum02; + outptr[5] = sum12; + outptr[6] = sum03; + outptr[7] = sum13; + outptr += 8; + pB += (size_t)4 * (B_hstep - k0 - max_kk0); + pB_descales += (size_t)4 * (num_blocks - block_start - block_count); + } + for (; jj + 1 < max_jj; jj += 2) + { + pB += (size_t)2 * k0; + pB_descales += (size_t)2 * block_start; + float sum00 = 0.f, sum01 = 0.f, sum10 = 0.f, sum11 = 0.f; + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + const int max_kk = std::min(K - k, block_size); + int s00 = 0, s01 = 0, s10 = 0, s11 = 0; + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + const signed char* b0 = pB; + const signed char* b1 = b0 + 4; + s00 += pA[0] * b0[0] + pA[2] * b0[1] + pA[4] * b0[2] + pA[6] * b0[3]; + s01 += pA[0] * b1[0] + pA[2] * b1[1] + pA[4] * b1[2] + pA[6] * b1[3]; + s10 += pA[1] * b0[0] + pA[3] * b0[1] + pA[5] * b0[2] + pA[7] * b0[3]; + s11 += pA[1] * b1[0] + pA[3] * b1[1] + pA[5] * b1[2] + pA[7] * b1[3]; + pA += 8; + pB += 8; + } + for (; kk + 1 < max_kk; kk += 2) + { + const signed char* b0 = pB; + const signed char* b1 = b0 + 2; + s00 += pA[0] * b0[0] + pA[2] * b0[1]; + s01 += pA[0] * b1[0] + pA[2] * b1[1]; + s10 += pA[1] * b0[0] + pA[3] * b0[1]; + s11 += pA[1] * b1[0] + pA[3] * b1[1]; + pA += 4; + pB += 4; + } + for (; kk < max_kk; kk++) + { + s00 += pA[0] * pB[0]; + s01 += pA[0] * pB[1]; + s10 += pA[1] * pB[0]; + s11 += pA[1] * pB[1]; + pA += 2; + pB += 2; + } + const float ad0 = pA_descales[0]; + const float ad1 = pA_descales[1]; + const float* bd = pB_descales; + sum00 += s00 * ad0 * bd[0]; + sum01 += s01 * ad0 * bd[1]; + sum10 += s10 * ad1 * bd[0]; + sum11 += s11 * ad1 * bd[1]; + pA_descales += 2; + pB_descales += 2; + } + if (k0 != 0) + { + sum00 += outptr[0]; + sum10 += outptr[1]; + sum01 += outptr[2]; + sum11 += outptr[3]; + } + outptr[0] = sum00; + outptr[1] = sum10; + outptr[2] = sum01; + outptr[3] = sum11; + outptr += 4; + pB += (size_t)2 * (B_hstep - k0 - max_kk0); + pB_descales += (size_t)2 * (num_blocks - block_start - block_count); + } + for (; jj < max_jj; jj++) + { + pB += k0; + pB_descales += block_start; + float sum0 = 0.f, sum1 = 0.f; + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + const int max_kk = std::min(K - k, block_size); + int s0 = 0, s1 = 0; + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + const signed char* b = pB; + s0 += pA[0] * b[0] + pA[2] * b[1] + pA[4] * b[2] + pA[6] * b[3]; + s1 += pA[1] * b[0] + pA[3] * b[1] + pA[5] * b[2] + pA[7] * b[3]; + pA += 8; + pB += 4; + } + for (; kk + 1 < max_kk; kk += 2) + { + const signed char* b = pB; + s0 += pA[0] * b[0] + pA[2] * b[1]; + s1 += pA[1] * b[0] + pA[3] * b[1]; + pA += 4; + pB += 2; + } + for (; kk < max_kk; kk++) + { + s0 += pA[0] * pB[0]; + s1 += pA[1] * pB[0]; + pA += 2; + pB++; + } + const float bd = pB_descales[0]; + sum0 += s0 * pA_descales[0] * bd; + sum1 += s1 * pA_descales[1] * bd; + pA_descales += 2; + pB_descales++; + } + if (k0 != 0) + { + sum0 += outptr[0]; + sum1 += outptr[1]; + } + outptr[0] = sum0; + outptr[1] = sum1; + outptr += 2; + pB += B_hstep - k0 - max_kk0; + pB_descales += num_blocks - block_start - block_count; + } + pAT += K * 2; + pAT_descales += block_count * 2; + } + for (; ii < max_ii; ii++) + { + const signed char* pB = pBT; + const float* pB_descales = pBT_descales; + int jj = 0; + for (; jj + 3 < max_jj; jj += 4) + { + pB += (size_t)4 * k0; + pB_descales += (size_t)4 * block_start; + float sum0 = 0.f, sum1 = 0.f, sum2 = 0.f, sum3 = 0.f; + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + const int max_kk = std::min(K - k, block_size); + int s0 = 0, s1 = 0, s2 = 0, s3 = 0; + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + const signed char* b0 = pB; + const signed char* b1 = b0 + 4; + const signed char* b2 = b1 + 4; + const signed char* b3 = b2 + 4; + s0 += pA[0] * b0[0] + pA[1] * b0[1] + pA[2] * b0[2] + pA[3] * b0[3]; + s1 += pA[0] * b1[0] + pA[1] * b1[1] + pA[2] * b1[2] + pA[3] * b1[3]; + s2 += pA[0] * b2[0] + pA[1] * b2[1] + pA[2] * b2[2] + pA[3] * b2[3]; + s3 += pA[0] * b3[0] + pA[1] * b3[1] + pA[2] * b3[2] + pA[3] * b3[3]; + pA += 4; + pB += 16; + } + for (; kk + 1 < max_kk; kk += 2) + { + const signed char* b0 = pB; + const signed char* b1 = b0 + 2; + const signed char* b2 = b1 + 2; + const signed char* b3 = b2 + 2; + s0 += pA[0] * b0[0] + pA[1] * b0[1]; + s1 += pA[0] * b1[0] + pA[1] * b1[1]; + s2 += pA[0] * b2[0] + pA[1] * b2[1]; + s3 += pA[0] * b3[0] + pA[1] * b3[1]; + pA += 2; + pB += 8; + } + for (; kk < max_kk; kk++) + { + s0 += pA[0] * pB[0]; + s1 += pA[0] * pB[1]; + s2 += pA[0] * pB[2]; + s3 += pA[0] * pB[3]; + pA++; + pB += 4; + } + const float ad = pA_descales[0]; + const float* bd = pB_descales; + sum0 += s0 * ad * bd[0]; + sum1 += s1 * ad * bd[1]; + sum2 += s2 * ad * bd[2]; + sum3 += s3 * ad * bd[3]; + pA_descales++; + pB_descales += 4; + } + if (k0 != 0) + { + sum0 += outptr[0]; + sum1 += outptr[1]; + sum2 += outptr[2]; + sum3 += outptr[3]; + } + outptr[0] = sum0; + outptr[1] = sum1; + outptr[2] = sum2; + outptr[3] = sum3; + outptr += 4; + pB += (size_t)4 * (B_hstep - k0 - max_kk0); + pB_descales += (size_t)4 * (num_blocks - block_start - block_count); + } + for (; jj + 1 < max_jj; jj += 2) + { + pB += (size_t)2 * k0; + pB_descales += (size_t)2 * block_start; + float sum0 = 0.f, sum1 = 0.f; + const signed char* pA = pAT; + const float* pA_descales = pAT_descales; + for (int k = 0; k < K; k += block_size) + { + const int max_kk = std::min(K - k, block_size); + int s0 = 0, s1 = 0; + int kk = 0; + for (; kk + 3 < max_kk; kk += 4) + { + const signed char* b0 = pB; + const signed char* b1 = b0 + 4; + s0 += pA[0] * b0[0] + pA[1] * b0[1] + pA[2] * b0[2] + pA[3] * b0[3]; + s1 += pA[0] * b1[0] + pA[1] * b1[1] + pA[2] * b1[2] + pA[3] * b1[3]; + pA += 4; + pB += 8; + } + for (; kk + 1 < max_kk; kk += 2) + { + const signed char* b0 = pB; + const signed char* b1 = b0 + 2; + s0 += pA[0] * b0[0] + pA[1] * b0[1]; + s1 += pA[0] * b1[0] + pA[1] * b1[1]; pA += 2; pB += 4; } @@ -1184,12 +1890,21 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pA_descales++; pB_descales += 2; } + if (k0 != 0) + { + sum0 += outptr[0]; + sum1 += outptr[1]; + } outptr[0] = sum0; outptr[1] = sum1; outptr += 2; + pB += (size_t)2 * (B_hstep - k0 - max_kk0); + pB_descales += (size_t)2 * (num_blocks - block_start - block_count); } for (; jj < max_jj; jj++) { + pB += k0; + pB_descales += block_start; float sum = 0.f; const signed char* pA = pAT; const float* pA_descales = pAT_descales; @@ -1222,8 +1937,12 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pA_descales++; pB_descales++; } + if (k0 != 0) + sum += outptr[0]; outptr[0] = sum; outptr++; + pB += B_hstep - k0 - max_kk0; + pB_descales += num_blocks - block_start - block_count; } pAT += K; pAT_descales += block_count; @@ -1231,8 +1950,9 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #endif } -static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta) +static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta) { + const float* pp = topT; beta *= alpha; (void)N; const size_t c_hstep = C.dims == 3 ? C.cstep : (size_t)C.w; @@ -1265,6 +1985,7 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b if (pC && (broadcast_type_C == 1 || broadcast_type_C == 2)) _c = __riscv_vfmul_vf_f32m1(__riscv_vle32_v_f32m1(pC, vl_packn), beta, vl_packn); + float* out0 = outptr; for (int jj = 0; jj < max_jj; jj++) { vfloat32m1_t _sum = __riscv_vle32_v_f32m1(pp, vl_packn); @@ -1278,22 +1999,30 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b _sum = __riscv_vfadd_vv_f32m1(_sum, _c, vl_packn); if (broadcast_type_C == 3) { - const vfloat32m1_t _c0 = __riscv_vlse32_v_f32m1(pC + jj, c_stride, vl_packn); + vfloat32m1_t _c0 = __riscv_vlse32_v_f32m1(pC, c_stride, vl_packn); + pC++; if (beta == 1.f) _sum = __riscv_vfadd_vv_f32m1(_sum, _c0, vl_packn); else _sum = __riscv_vfmacc_vf_f32m1(_sum, beta, _c0, vl_packn); } if (broadcast_type_C == 4) - _sum = __riscv_vfadd_vf_f32m1(_sum, pC[jj] * beta, vl_packn); + { + _sum = __riscv_vfadd_vf_f32m1(_sum, *pC * beta, vl_packn); + pC++; + } } - __riscv_vsse32_v_f32m1(outptr + jj, out_stride, _sum, vl_packn); + __riscv_vsse32_v_f32m1(out0, out_stride, _sum, vl_packn); + out0++; pp += packn; } outptr += out_hstep * packn; } +#endif // __riscv_vector for (; ii + 1 < max_ii; ii += 2) { + float* out0 = outptr; + float* out1 = out0 + out_hstep; const float* pC = C; if (pC) { @@ -1310,13 +2039,12 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b if (pC && (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2)) { c0 = pC[0] * beta; - c1 = pC[broadcast_type_C == 0 ? 0 : 1] * beta; - } - - float* out0 = outptr; - float* out1 = out0 + out_hstep; + c1 = pC[broadcast_type_C == 0 ? 0 : 1] * beta; + } + int jj = 0; +#if __riscv_vector const size_t vl = __riscv_vsetvl_e32m4(max_jj); - const vfloat32m4x2_t _s = __riscv_vlseg2e32_v_f32m4x2(pp, vl); + vfloat32m4x2_t _s = __riscv_vlseg2e32_v_f32m4x2(pp, vl); vfloat32m4_t _sum0 = __riscv_vget_v_f32m4x2_f32m4(_s, 0); vfloat32m4_t _sum1 = __riscv_vget_v_f32m4x2_f32m4(_s, 1); if (alpha != 1.f) @@ -1334,8 +2062,8 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b } if (broadcast_type_C == 3) { - const vfloat32m4_t _c0 = __riscv_vle32_v_f32m4(pC, vl); - const vfloat32m4_t _c1 = __riscv_vle32_v_f32m4(pC + c_hstep, vl); + vfloat32m4_t _c0 = __riscv_vle32_v_f32m4(pC, vl); + vfloat32m4_t _c1 = __riscv_vle32_v_f32m4(pC + c_hstep, vl); if (beta == 1.f) { _sum0 = __riscv_vfadd_vv_f32m4(_sum0, _c0, vl); @@ -1349,7 +2077,7 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b } if (broadcast_type_C == 4) { - const vfloat32m4_t _c = __riscv_vle32_v_f32m4(pC, vl); + vfloat32m4_t _c = __riscv_vle32_v_f32m4(pC, vl); if (beta == 1.f) { _sum0 = __riscv_vfadd_vv_f32m4(_sum0, _c, vl); @@ -1365,82 +2093,13 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b __riscv_vse32_v_f32m4(out0, _sum0, vl); __riscv_vse32_v_f32m4(out1, _sum1, vl); + jj += (int)vl; pp += vl * 2; - outptr += out_hstep * 2; - } - for (; ii < max_ii; ii++) - { - const float* pC = C; - if (pC) - { - if (broadcast_type_C == 1 || broadcast_type_C == 2) - pC += i + ii; - if (broadcast_type_C == 3) - pC += (size_t)(i + ii) * c_hstep + j; - if (broadcast_type_C == 4) - pC += j; - } - - float c0 = 0.f; - if (pC && (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2)) - c0 = pC[0] * beta; - - float* out0 = outptr; - const size_t vl = __riscv_vsetvl_e32m4(max_jj); - vfloat32m4_t _sum = __riscv_vle32_v_f32m4(pp, vl); - if (alpha != 1.f) - _sum = __riscv_vfmul_vf_f32m4(_sum, alpha, vl); - - if (pC) - { - if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) - _sum = __riscv_vfadd_vf_f32m4(_sum, c0, vl); - if (broadcast_type_C == 3) - { - const vfloat32m4_t _c = __riscv_vle32_v_f32m4(pC, vl); - if (beta == 1.f) - _sum = __riscv_vfadd_vv_f32m4(_sum, _c, vl); - else - _sum = __riscv_vfmacc_vf_f32m4(_sum, beta, _c, vl); - } - if (broadcast_type_C == 4) - { - const vfloat32m4_t _c = __riscv_vle32_v_f32m4(pC, vl); - if (beta == 1.f) - _sum = __riscv_vfadd_vv_f32m4(_sum, _c, vl); - else - _sum = __riscv_vfmacc_vf_f32m4(_sum, beta, _c, vl); - } - } - - __riscv_vse32_v_f32m4(out0, _sum, vl); - pp += vl; - outptr += out_hstep; - } -#else - for (; ii + 1 < max_ii; ii += 2) - { - float* out0 = outptr; - float* out1 = out0 + out_hstep; - const float* pC = C; - if (pC) - { - if (broadcast_type_C == 1 || broadcast_type_C == 2) - pC += i + ii; - if (broadcast_type_C == 3) - pC += (size_t)(i + ii) * c_hstep + j; - if (broadcast_type_C == 4) - pC += j; - } - - float c0 = 0.f; - float c1 = 0.f; - if (pC && (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2)) - { - c0 = pC[0] * beta; - c1 = pC[broadcast_type_C == 0 ? 0 : 1] * beta; - } - int jj = 0; + out0 += vl; + out1 += vl; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += vl; +#endif // __riscv_vector for (; jj + 3 < max_jj; jj += 4) { float sum00 = pp[0] * alpha; @@ -1451,7 +2110,6 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b float sum12 = pp[5] * alpha; float sum03 = pp[6] * alpha; float sum13 = pp[7] * alpha; - pp += 8; if (pC) { @@ -1487,6 +2145,7 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b sum11 += pC[c_hstep + 1] * beta; sum12 += pC[c_hstep + 2] * beta; sum13 += pC[c_hstep + 3] * beta; + pC += 4; } if (broadcast_type_C == 4) { @@ -1498,6 +2157,7 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b sum11 += pC[1] * beta; sum12 += pC[2] * beta; sum13 += pC[3] * beta; + pC += 4; } } @@ -1511,8 +2171,7 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b out1[3] = sum13; out0 += 4; out1 += 4; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 4; + pp += 8; } for (; jj + 1 < max_jj; jj += 2) { @@ -1520,7 +2179,6 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b float sum10 = pp[1] * alpha; float sum01 = pp[2] * alpha; float sum11 = pp[3] * alpha; - pp += 4; if (pC) { @@ -1544,6 +2202,7 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b sum01 += pC[1] * beta; sum10 += pC[c_hstep] * beta; sum11 += pC[c_hstep + 1] * beta; + pC += 2; } if (broadcast_type_C == 4) { @@ -1551,6 +2210,7 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b sum01 += pC[1] * beta; sum10 += pC[0] * beta; sum11 += pC[1] * beta; + pC += 2; } } @@ -1560,14 +2220,12 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b out1[1] = sum11; out0 += 2; out1 += 2; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 2; + pp += 4; } for (; jj < max_jj; jj++) { float sum0 = pp[0] * alpha; float sum1 = pp[1] * alpha; - pp += 2; if (pC) { if (broadcast_type_C == 0) @@ -1584,19 +2242,20 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b { sum0 += pC[0] * beta; sum1 += pC[c_hstep] * beta; + pC++; } if (broadcast_type_C == 4) { sum0 += pC[0] * beta; sum1 += pC[0] * beta; + pC++; } } out0[0] = sum0; out1[0] = sum1; out0++; out1++; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC++; + pp += 2; } outptr += out_hstep * 2; } @@ -1618,13 +2277,47 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b if (pC && (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2)) c0 = pC[0] * beta; int jj = 0; +#if __riscv_vector + const size_t vl = __riscv_vsetvl_e32m4(max_jj); + vfloat32m4_t _sum = __riscv_vle32_v_f32m4(pp, vl); + if (alpha != 1.f) + _sum = __riscv_vfmul_vf_f32m4(_sum, alpha, vl); + + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + _sum = __riscv_vfadd_vf_f32m4(_sum, c0, vl); + if (broadcast_type_C == 3) + { + vfloat32m4_t _c = __riscv_vle32_v_f32m4(pC, vl); + if (beta == 1.f) + _sum = __riscv_vfadd_vv_f32m4(_sum, _c, vl); + else + _sum = __riscv_vfmacc_vf_f32m4(_sum, beta, _c, vl); + } + if (broadcast_type_C == 4) + { + vfloat32m4_t _c = __riscv_vle32_v_f32m4(pC, vl); + if (beta == 1.f) + _sum = __riscv_vfadd_vv_f32m4(_sum, _c, vl); + else + _sum = __riscv_vfmacc_vf_f32m4(_sum, beta, _c, vl); + } + } + + __riscv_vse32_v_f32m4(out0, _sum, vl); + jj += (int)vl; + pp += vl; + out0 += vl; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += vl; +#endif // __riscv_vector for (; jj + 3 < max_jj; jj += 4) { float sum0 = pp[0] * alpha; float sum1 = pp[1] * alpha; float sum2 = pp[2] * alpha; float sum3 = pp[3] * alpha; - pp += 4; if (pC) { if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) @@ -1640,6 +2333,7 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b sum1 += pC[1] * beta; sum2 += pC[2] * beta; sum3 += pC[3] * beta; + pC += 4; } if (broadcast_type_C == 4) { @@ -1647,6 +2341,7 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b sum1 += pC[1] * beta; sum2 += pC[2] * beta; sum3 += pC[3] * beta; + pC += 4; } } out0[0] = sum0; @@ -1654,14 +2349,12 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b out0[2] = sum2; out0[3] = sum3; out0 += 4; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 4; + pp += 4; } for (; jj + 1 < max_jj; jj += 2) { float sum0 = pp[0] * alpha; float sum1 = pp[1] * alpha; - pp += 2; if (pC) { if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) @@ -1673,22 +2366,23 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b { sum0 += pC[0] * beta; sum1 += pC[1] * beta; + pC += 2; } if (broadcast_type_C == 4) { sum0 += pC[0] * beta; sum1 += pC[1] * beta; + pC += 2; } } out0[0] = sum0; out0[1] = sum1; out0 += 2; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 2; + pp += 2; } for (; jj < max_jj; jj++) { - float sum = *pp++ * alpha; + float sum = *pp * alpha; if (pC) { if (broadcast_type_C == 0) @@ -1696,33 +2390,36 @@ static void unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_b if (broadcast_type_C == 1 || broadcast_type_C == 2) sum += c0; if (broadcast_type_C == 3) + { sum += pC[0] * beta; + pC++; + } if (broadcast_type_C == 4) + { sum += pC[0] * beta; + pC++; + } } out0[0] = sum; out0++; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC++; + pp++; } outptr += out_hstep; } -#endif } -static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta) +static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta) { + const float* pp = topT; beta *= alpha; (void)N; const size_t c_hstep = C.dims == 3 ? C.cstep : (size_t)C.w; const size_t out_hstep = top_blob.dims == 3 ? top_blob.cstep : (size_t)top_blob.w; float* outptr = (float*)top_blob + (size_t)j * out_hstep + i; -#if __riscv_vector - const ptrdiff_t out_stride = (ptrdiff_t)out_hstep * sizeof(float); -#endif int ii = 0; #if __riscv_vector + const ptrdiff_t out_stride = (ptrdiff_t)out_hstep * sizeof(float); const int packn = csrr_vlenb() / 4; const size_t vl_packn = __riscv_vsetvl_e32m1(packn); const ptrdiff_t c_stride = (ptrdiff_t)c_hstep * sizeof(float); @@ -1746,6 +2443,7 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, if (pC && (broadcast_type_C == 1 || broadcast_type_C == 2)) _c = __riscv_vfmul_vf_f32m1(__riscv_vle32_v_f32m1(pC, vl_packn), beta, vl_packn); + float* out0 = outptr; for (int jj = 0; jj < max_jj; jj++) { vfloat32m1_t _sum = __riscv_vle32_v_f32m1(pp, vl_packn); @@ -1759,20 +2457,26 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, _sum = __riscv_vfadd_vv_f32m1(_sum, _c, vl_packn); if (broadcast_type_C == 3) { - const vfloat32m1_t _c0 = __riscv_vlse32_v_f32m1(pC + jj, c_stride, vl_packn); + vfloat32m1_t _c0 = __riscv_vlse32_v_f32m1(pC, c_stride, vl_packn); + pC++; if (beta == 1.f) _sum = __riscv_vfadd_vv_f32m1(_sum, _c0, vl_packn); else _sum = __riscv_vfmacc_vf_f32m1(_sum, beta, _c0, vl_packn); } if (broadcast_type_C == 4) - _sum = __riscv_vfadd_vf_f32m1(_sum, pC[jj] * beta, vl_packn); + { + _sum = __riscv_vfadd_vf_f32m1(_sum, *pC * beta, vl_packn); + pC++; + } } - __riscv_vse32_v_f32m1(outptr + (size_t)jj * out_hstep, _sum, vl_packn); + __riscv_vse32_v_f32m1(out0, _sum, vl_packn); + out0 += out_hstep; pp += packn; } outptr += packn; } +#endif // __riscv_vector for (; ii + 1 < max_ii; ii += 2) { float* out0 = outptr; @@ -1794,9 +2498,10 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, c0 = pC[0] * beta; c1 = pC[broadcast_type_C == 0 ? 0 : 1] * beta; } - + int jj = 0; +#if __riscv_vector const size_t vl = __riscv_vsetvl_e32m4(max_jj); - const vfloat32m4x2_t _s = __riscv_vlseg2e32_v_f32m4x2(pp, vl); + vfloat32m4x2_t _s = __riscv_vlseg2e32_v_f32m4x2(pp, vl); vfloat32m4_t _sum0 = __riscv_vget_v_f32m4x2_f32m4(_s, 0); vfloat32m4_t _sum1 = __riscv_vget_v_f32m4x2_f32m4(_s, 1); if (alpha != 1.f) @@ -1813,8 +2518,8 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, } if (broadcast_type_C == 3) { - const vfloat32m4_t _c0 = __riscv_vle32_v_f32m4(pC, vl); - const vfloat32m4_t _c1 = __riscv_vle32_v_f32m4(pC + c_hstep, vl); + vfloat32m4_t _c0 = __riscv_vle32_v_f32m4(pC, vl); + vfloat32m4_t _c1 = __riscv_vle32_v_f32m4(pC + c_hstep, vl); if (beta == 1.f) { _sum0 = __riscv_vfadd_vv_f32m4(_sum0, _c0, vl); @@ -1828,7 +2533,7 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, } if (broadcast_type_C == 4) { - const vfloat32m4_t _c = __riscv_vle32_v_f32m4(pC, vl); + vfloat32m4_t _c = __riscv_vle32_v_f32m4(pC, vl); if (beta == 1.f) { _sum0 = __riscv_vfadd_vv_f32m4(_sum0, _c, vl); @@ -1841,87 +2546,17 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, } } } - const vfloat32m4x2_t _sum = __riscv_vcreate_v_f32m4x2(_sum0, _sum1); + vfloat32m4x2_t _sum = __riscv_vcreate_v_f32m4x2(_sum0, _sum1); if (out_hstep == 2) __riscv_vsseg2e32_v_f32m4x2(out0, _sum, vl); else __riscv_vssseg2e32_v_f32m4x2(out0, out_stride, _sum, vl); + jj += (int)vl; pp += vl * 2; - outptr += 2; - } - for (; ii < max_ii; ii++) - { - float* out0 = outptr; - const float* pC = C; - if (pC) - { - if (broadcast_type_C == 1 || broadcast_type_C == 2) - pC += i + ii; - if (broadcast_type_C == 3) - pC += (size_t)(i + ii) * c_hstep + j; - if (broadcast_type_C == 4) - pC += j; - } - - float c0 = 0.f; - if (pC && (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2)) - c0 = pC[0] * beta; - - const size_t vl = __riscv_vsetvl_e32m4(max_jj); - vfloat32m4_t _sum = __riscv_vle32_v_f32m4(pp, vl); - if (alpha != 1.f) - _sum = __riscv_vfmul_vf_f32m4(_sum, alpha, vl); - if (pC) - { - if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) - _sum = __riscv_vfadd_vf_f32m4(_sum, c0, vl); - if (broadcast_type_C == 3) - { - const vfloat32m4_t _c = __riscv_vle32_v_f32m4(pC, vl); - if (beta == 1.f) - _sum = __riscv_vfadd_vv_f32m4(_sum, _c, vl); - else - _sum = __riscv_vfmacc_vf_f32m4(_sum, beta, _c, vl); - } - if (broadcast_type_C == 4) - { - const vfloat32m4_t _c = __riscv_vle32_v_f32m4(pC, vl); - if (beta == 1.f) - _sum = __riscv_vfadd_vv_f32m4(_sum, _c, vl); - else - _sum = __riscv_vfmacc_vf_f32m4(_sum, beta, _c, vl); - } - } - if (out_hstep == 1) - __riscv_vse32_v_f32m4(out0, _sum, vl); - else - __riscv_vsse32_v_f32m4(out0, out_stride, _sum, vl); - pp += vl; - outptr++; - } -#else - for (; ii + 1 < max_ii; ii += 2) - { - float* out0 = outptr; - const float* pC = C; - if (pC) - { - if (broadcast_type_C == 1 || broadcast_type_C == 2) - pC += i + ii; - if (broadcast_type_C == 3) - pC += (size_t)(i + ii) * c_hstep + j; - if (broadcast_type_C == 4) - pC += j; - } - - float c0 = 0.f; - float c1 = 0.f; - if (pC && (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2)) - { - c0 = pC[0] * beta; - c1 = pC[broadcast_type_C == 0 ? 0 : 1] * beta; - } - int jj = 0; + out0 += out_hstep * vl; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += vl; +#endif // __riscv_vector for (; jj + 3 < max_jj; jj += 4) { float sum00 = pp[0] * alpha; @@ -1932,7 +2567,6 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, float sum12 = pp[5] * alpha; float sum03 = pp[6] * alpha; float sum13 = pp[7] * alpha; - pp += 8; if (pC) { if (broadcast_type_C == 0) @@ -1979,6 +2613,8 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, sum12 += pC[2] * beta; sum13 += pC[3] * beta; } + if (broadcast_type_C == 3 || broadcast_type_C == 4) + pC += 4; } out0[0] = sum00; out0[1] = sum10; @@ -1989,8 +2625,7 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, out0[out_hstep * 3] = sum03; out0[out_hstep * 3 + 1] = sum13; out0 += out_hstep * 4; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 4; + pp += 8; } for (; jj + 1 < max_jj; jj += 2) { @@ -1998,7 +2633,6 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, float sum10 = pp[1] * alpha; float sum01 = pp[2] * alpha; float sum11 = pp[3] * alpha; - pp += 4; if (pC) { if (broadcast_type_C == 0) @@ -2029,20 +2663,20 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, sum10 += pC[0] * beta; sum11 += pC[1] * beta; } + if (broadcast_type_C == 3 || broadcast_type_C == 4) + pC += 2; } out0[0] = sum00; out0[1] = sum10; out0[out_hstep] = sum01; out0[out_hstep + 1] = sum11; out0 += out_hstep * 2; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 2; + pp += 4; } for (; jj < max_jj; jj++) { float sum0 = pp[0] * alpha; float sum1 = pp[1] * alpha; - pp += 2; if (pC) { if (broadcast_type_C == 0) @@ -2065,12 +2699,13 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, sum0 += pC[0] * beta; sum1 += pC[0] * beta; } + if (broadcast_type_C == 3 || broadcast_type_C == 4) + pC++; } out0[0] = sum0; out0[1] = sum1; out0 += out_hstep; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC++; + pp += 2; } outptr += 2; } @@ -2092,13 +2727,48 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, if (pC && (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2)) c0 = pC[0] * beta; int jj = 0; +#if __riscv_vector + const size_t vl = __riscv_vsetvl_e32m4(max_jj); + vfloat32m4_t _sum = __riscv_vle32_v_f32m4(pp, vl); + if (alpha != 1.f) + _sum = __riscv_vfmul_vf_f32m4(_sum, alpha, vl); + if (pC) + { + if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) + _sum = __riscv_vfadd_vf_f32m4(_sum, c0, vl); + if (broadcast_type_C == 3) + { + vfloat32m4_t _c = __riscv_vle32_v_f32m4(pC, vl); + if (beta == 1.f) + _sum = __riscv_vfadd_vv_f32m4(_sum, _c, vl); + else + _sum = __riscv_vfmacc_vf_f32m4(_sum, beta, _c, vl); + } + if (broadcast_type_C == 4) + { + vfloat32m4_t _c = __riscv_vle32_v_f32m4(pC, vl); + if (beta == 1.f) + _sum = __riscv_vfadd_vv_f32m4(_sum, _c, vl); + else + _sum = __riscv_vfmacc_vf_f32m4(_sum, beta, _c, vl); + } + } + if (out_hstep == 1) + __riscv_vse32_v_f32m4(out0, _sum, vl); + else + __riscv_vsse32_v_f32m4(out0, out_stride, _sum, vl); + jj += (int)vl; + pp += vl; + out0 += out_hstep * vl; + if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) + pC += vl; +#endif // __riscv_vector for (; jj + 3 < max_jj; jj += 4) { float sum0 = pp[0] * alpha; float sum1 = pp[1] * alpha; float sum2 = pp[2] * alpha; float sum3 = pp[3] * alpha; - pp += 4; if (pC) { if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) @@ -2122,20 +2792,20 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, sum2 += pC[2] * beta; sum3 += pC[3] * beta; } + if (broadcast_type_C == 3 || broadcast_type_C == 4) + pC += 4; } out0[0] = sum0; out0[out_hstep] = sum1; out0[out_hstep * 2] = sum2; out0[out_hstep * 3] = sum3; out0 += out_hstep * 4; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 4; + pp += 4; } for (; jj + 1 < max_jj; jj += 2) { float sum0 = pp[0] * alpha; float sum1 = pp[1] * alpha; - pp += 2; if (pC) { if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) @@ -2153,16 +2823,17 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, sum0 += pC[0] * beta; sum1 += pC[1] * beta; } + if (broadcast_type_C == 3 || broadcast_type_C == 4) + pC += 2; } out0[0] = sum0; out0[out_hstep] = sum1; out0 += out_hstep * 2; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC += 2; + pp += 2; } for (; jj < max_jj; jj++) { - float sum = *pp++ * alpha; + float sum = *pp * alpha; if (pC) { if (broadcast_type_C == 0) @@ -2173,19 +2844,23 @@ static void transpose_unpack_output_tile_wq_int8(const float* pp, const Mat& C, sum += pC[0] * beta; if (broadcast_type_C == 4) sum += pC[0] * beta; + if (broadcast_type_C == 3 || broadcast_type_C == 4) + pC++; } out0[0] = sum; out0 += out_hstep; - if (pC && (broadcast_type_C == 3 || broadcast_type_C == 4)) - pC++; + pp++; } outptr++; } -#endif } -static void get_optimal_tile_mnk_wq_int8(int M, int N, int K, int constant_TILE_M, int constant_TILE_N, int constant_TILE_K, int& TILE_M, int& TILE_N, int& TILE_K, int nT) +static void get_optimal_tile_mnk_wq_int8(int M, int N, int K, int block_size, int constant_TILE_M, int constant_TILE_N, int constant_TILE_K, int& TILE_M, int& TILE_N, int& TILE_K, int nT) { + // resolve optimal tile size from cache size + const size_t l2_cache_size = get_cpu_level2_cache_size(); + int tile_k = (int)sqrtf((float)l2_cache_size / (2 * sizeof(signed char) + sizeof(float))); + #if __riscv_vector const int packm = std::max(8, csrr_vlenb() / 4); const int packn = csrr_vlenb(); @@ -2196,7 +2871,21 @@ static void get_optimal_tile_mnk_wq_int8(int M, int N, int K, int constant_TILE_ TILE_M = packm; TILE_N = packn; - TILE_K = K; + TILE_K = std::max(block_size, tile_k / block_size * block_size); + + if (K > 0) + { + if (TILE_K >= K) + { + TILE_K = K; + } + else + { + const int nn_K = (K + TILE_K - 1) / TILE_K; + tile_k = (K + nn_K - 1) / nn_K; + TILE_K = std::max(block_size, tile_k / block_size * block_size); + } + } // take constant TILE_M/N value when provided if (constant_TILE_M > 0) @@ -2213,8 +2902,14 @@ static void get_optimal_tile_mnk_wq_int8(int M, int N, int K, int constant_TILE_ TILE_M = std::min(TILE_M, packm); TILE_N = std::min(TILE_N, packn); + if (constant_TILE_K > 0) + { + TILE_K = std::max(block_size, constant_TILE_K / block_size * block_size); + if (K > 0) + TILE_K = std::min(TILE_K, K); + } + (void)M; (void)N; - (void)constant_TILE_K; (void)nT; } diff --git a/src/layer/x86/gemm_wq_int8.h b/src/layer/x86/gemm_wq_int8.h index f098e6dbcc08..53f1b34b3d1f 100644 --- a/src/layer/x86/gemm_wq_int8.h +++ b/src/layer/x86/gemm_wq_int8.h @@ -5,7 +5,7 @@ void pack_B_tile_wq_int8_avx512vnni(const Mat& B, const Mat& B_scales, unsigned char* pp, float* pd, int j, int max_jj, int K, int block_size); void quantize_A_tile_wq_int8_avx512vnni(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr); void transpose_quantize_A_tile_wq_int8_avx512vnni(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr); -void gemm_transB_packed_tile_wq_int8_avx512vnni(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int block_size); +void gemm_transB_packed_tile_wq_int8_avx512vnni(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int k, int max_kk, int K, int block_size); void unpack_output_tile_wq_int8_avx512vnni(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta); void transpose_unpack_output_tile_wq_int8_avx512vnni(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta); #endif @@ -14,7 +14,7 @@ void transpose_unpack_output_tile_wq_int8_avx512vnni(const Mat& topT, const Mat& void pack_B_tile_wq_int8_avxvnniint8(const Mat& B, const Mat& B_scales, unsigned char* pp, float* pd, int j, int max_jj, int K, int block_size); void quantize_A_tile_wq_int8_avxvnniint8(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr); void transpose_quantize_A_tile_wq_int8_avxvnniint8(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr); -void gemm_transB_packed_tile_wq_int8_avxvnniint8(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int block_size); +void gemm_transB_packed_tile_wq_int8_avxvnniint8(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int k, int max_kk, int K, int block_size); void unpack_output_tile_wq_int8_avxvnniint8(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta); void transpose_unpack_output_tile_wq_int8_avxvnniint8(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta); #endif @@ -23,7 +23,7 @@ void transpose_unpack_output_tile_wq_int8_avxvnniint8(const Mat& topT, const Mat void pack_B_tile_wq_int8_avxvnni(const Mat& B, const Mat& B_scales, unsigned char* pp, float* pd, int j, int max_jj, int K, int block_size); void quantize_A_tile_wq_int8_avxvnni(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr); void transpose_quantize_A_tile_wq_int8_avxvnni(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr); -void gemm_transB_packed_tile_wq_int8_avxvnni(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int block_size); +void gemm_transB_packed_tile_wq_int8_avxvnni(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int k, int max_kk, int K, int block_size); void unpack_output_tile_wq_int8_avxvnni(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta); void transpose_unpack_output_tile_wq_int8_avxvnni(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta); #endif @@ -32,13 +32,13 @@ void transpose_unpack_output_tile_wq_int8_avxvnni(const Mat& topT, const Mat& C, void pack_B_tile_wq_int8_avx2(const Mat& B, const Mat& B_scales, unsigned char* pp, float* pd, int j, int max_jj, int K, int block_size); void quantize_A_tile_wq_int8_avx2(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr); void transpose_quantize_A_tile_wq_int8_avx2(const Mat& A, Mat& AT_tile, Mat& AT_descales_tile, int i, int max_ii, int K, int block_size, const float* input_scale_ptr); -void gemm_transB_packed_tile_wq_int8_avx2(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int block_size); +void gemm_transB_packed_tile_wq_int8_avx2(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int k, int max_kk, int K, int block_size); void unpack_output_tile_wq_int8_avx2(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta); void transpose_unpack_output_tile_wq_int8_avx2(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta); #endif #if NCNN_RUNTIME_CPU && NCNN_XOP && __SSE2__ && !__XOP__ && !__AVX2__ && !__AVXVNNI__ && !__AVXVNNIINT8__ && !__AVX512VNNI__ -void gemm_transB_packed_tile_wq_int8_xop(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int block_size); +void gemm_transB_packed_tile_wq_int8_xop(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int k, int max_kk, int K, int block_size); #endif static void pack_B_tile_wq_int8(const Mat& B, const Mat& B_scales, unsigned char* pp, float* pd, int j, int max_jj, int K, int block_size) @@ -80,26 +80,30 @@ static void pack_B_tile_wq_int8(const Mat& B, const Mat& B_scales, unsigned char #if __AVX512F__ for (; jj + 7 < max_jj; jj += 8) { + const signed char* p0 = B.row(j + jj); + const float* ps = B_scales.row(j + jj); + + __m256i _vindex = _mm256_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7); + _vindex = _mm256_mullo_epi32(_vindex, _mm256_set1_epi32(B.w)); + __m256i _sindex = _mm256_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7); + _sindex = _mm256_mullo_epi32(_sindex, _mm256_set1_epi32(B_scales.w)); + for (int g = 0; g < block_count; g++) { - const int k0 = g * block_size; - const int max_kk = std::min(K - k0, block_size); + const int max_kk = std::min(K - g * block_size, block_size); - const signed char* p0 = B.row(j + jj) + k0; - __m256i _vindex = _mm256_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7); - _vindex = _mm256_mullo_epi32(_vindex, _mm256_set1_epi32(B.w)); int kk = 0; #if __AVX512VNNI__ || __AVXVNNI__ #if __AVXVNNIINT8__ for (; kk + 3 < max_kk; kk += 4) { - const __m256i _p = _mm256_i32gather_epi32((const int*)p0, _vindex, sizeof(signed char)); + __m256i _p = _mm256_i32gather_epi32((const int*)p0, _vindex, sizeof(signed char)); _mm256_storeu_si256((__m256i*)pp, _p); pp += 32; p0 += 4; } #else // __AVXVNNIINT8__ - const __m256i _v127 = _mm256_set1_epi8(127); + __m256i _v127 = _mm256_set1_epi8(127); for (; kk + 3 < max_kk; kk += 4) { __m256i _p = _mm256_i32gather_epi32((const int*)p0, _vindex, sizeof(signed char)); @@ -112,89 +116,96 @@ static void pack_B_tile_wq_int8(const Mat& B, const Mat& B_scales, unsigned char #endif // __AVX512VNNI__ || __AVXVNNI__ for (; kk + 1 < max_kk; kk += 2) { - const __m128i _p = _mm256_comp_cvtepi32_epi16(_mm256_i32gather_epi32((const int*)p0, _vindex, sizeof(signed char))); + __m128i _p = _mm256_comp_cvtepi32_epi16(_mm256_i32gather_epi32((const int*)p0, _vindex, sizeof(signed char))); _mm_storeu_si128((__m128i*)pp, _p); pp += 16; p0 += 2; } for (; kk < max_kk; kk++) { - const __m128i _p = _mm256_comp_cvtepi32_epi8(_mm256_i32gather_epi32((const int*)p0, _vindex, sizeof(signed char))); + __m128i _p = _mm256_comp_cvtepi32_epi8(_mm256_i32gather_epi32((const int*)p0, _vindex, sizeof(signed char))); _mm_storel_epi64((__m128i*)pp, _p); pp += 8; p0++; } - __m256i _sindex = _mm256_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7); - _sindex = _mm256_mullo_epi32(_sindex, _mm256_set1_epi32(B_scales.w)); - const __m256 _scale = _mm256_i32gather_ps(B_scales.row(j + jj) + g, _sindex, sizeof(float)); + __m256 _scale = _mm256_i32gather_ps(ps, _sindex, sizeof(float)); _mm256_storeu_ps(pd, _mm256_div_ps(_mm256_set1_ps(1.f), _scale)); pd += 8; + ps++; } } #endif // __AVX512F__ for (; jj + 3 < max_jj; jj += 4) { + const signed char* p0 = B.row(j + jj); + const signed char* p1 = B.row(j + jj + 1); + const signed char* p2 = B.row(j + jj + 2); + const signed char* p3 = B.row(j + jj + 3); + const float* ps0 = B_scales.row(j + jj); + const float* ps1 = B_scales.row(j + jj + 1); + const float* ps2 = B_scales.row(j + jj + 2); + const float* ps3 = B_scales.row(j + jj + 3); + for (int g = 0; g < block_count; g++) { - const int k0 = g * block_size; - const int max_kk = std::min(K - k0, block_size); + const int max_kk = std::min(K - g * block_size, block_size); int kk = 0; #if __AVX512VNNI__ || __AVXVNNI__ // VNNI consumes one contiguous K4 dword per output lane. for (; kk + 3 < max_kk; kk += 4) { - const signed char* p0 = B.row(j + jj) + k0 + kk; - const signed char* p1 = B.row(j + jj + 1) + k0 + kk; - const signed char* p2 = B.row(j + jj + 2) + k0 + kk; - const signed char* p3 = B.row(j + jj + 3) + k0 + kk; - __m128i _p = _mm_setr_epi32(*(const int*)p0, *(const int*)p1, *(const int*)p2, *(const int*)p3); + __m128i _p = _mm_setr_epi32(_mm_cvtsi128_si32(_mm_loadu_si32(p0)), _mm_cvtsi128_si32(_mm_loadu_si32(p1)), _mm_cvtsi128_si32(_mm_loadu_si32(p2)), _mm_cvtsi128_si32(_mm_loadu_si32(p3))); #if !__AVXVNNIINT8__ _p = _mm_add_epi8(_p, _mm_set1_epi8(127)); #endif // __AVXVNNIINT8__ _mm_storeu_si128((__m128i*)pp, _p); pp += 16; + p0 += 4; + p1 += 4; + p2 += 4; + p3 += 4; } #else // AVX2/SSE2 consumes two K2 vectors for each real K4 region. for (; kk + 3 < max_kk; kk += 4) { - const signed char* p0 = B.row(j + jj) + k0 + kk; - const signed char* p1 = B.row(j + jj + 1) + k0 + kk; - const signed char* p2 = B.row(j + jj + 2) + k0 + kk; - const signed char* p3 = B.row(j + jj + 3) + k0 + kk; - const __m128i _p01 = _mm_setr_epi16((short)*(const unsigned short*)p0, (short)*(const unsigned short*)p1, (short)*(const unsigned short*)p2, (short)*(const unsigned short*)p3, 0, 0, 0, 0); - const __m128i _p23 = _mm_setr_epi16((short)*(const unsigned short*)(p0 + 2), (short)*(const unsigned short*)(p1 + 2), (short)*(const unsigned short*)(p2 + 2), (short)*(const unsigned short*)(p3 + 2), 0, 0, 0, 0); + __m128i _p01 = _mm_setr_epi16((short)_mm_cvtsi128_si32(_mm_loadu_si16(p0)), (short)_mm_cvtsi128_si32(_mm_loadu_si16(p1)), (short)_mm_cvtsi128_si32(_mm_loadu_si16(p2)), (short)_mm_cvtsi128_si32(_mm_loadu_si16(p3)), 0, 0, 0, 0); + __m128i _p23 = _mm_setr_epi16((short)_mm_cvtsi128_si32(_mm_loadu_si16(p0 + 2)), (short)_mm_cvtsi128_si32(_mm_loadu_si16(p1 + 2)), (short)_mm_cvtsi128_si32(_mm_loadu_si16(p2 + 2)), (short)_mm_cvtsi128_si32(_mm_loadu_si16(p3 + 2)), 0, 0, 0, 0); _mm_storel_epi64((__m128i*)pp, _p01); _mm_storel_epi64((__m128i*)(pp + 8), _p23); pp += 16; + p0 += 4; + p1 += 4; + p2 += 4; + p3 += 4; } -#endif // __AVX512VNNI__ || __AVXVNNI__ \ -// K2/K1 are always signed and compact, including classic VNNI. +#endif // __AVX512VNNI__ || __AVXVNNI__ + // K2/K1 are always signed and compact, including classic VNNI. for (; kk + 1 < max_kk; kk += 2) { - const signed char* p0 = B.row(j + jj) + k0 + kk; - const signed char* p1 = B.row(j + jj + 1) + k0 + kk; - const signed char* p2 = B.row(j + jj + 2) + k0 + kk; - const signed char* p3 = B.row(j + jj + 3) + k0 + kk; - const __m128i _p = _mm_setr_epi16((short)*(const unsigned short*)p0, (short)*(const unsigned short*)p1, (short)*(const unsigned short*)p2, (short)*(const unsigned short*)p3, 0, 0, 0, 0); + __m128i _p = _mm_setr_epi16((short)_mm_cvtsi128_si32(_mm_loadu_si16(p0)), (short)_mm_cvtsi128_si32(_mm_loadu_si16(p1)), (short)_mm_cvtsi128_si32(_mm_loadu_si16(p2)), (short)_mm_cvtsi128_si32(_mm_loadu_si16(p3)), 0, 0, 0, 0); _mm_storel_epi64((__m128i*)pp, _p); pp += 8; + p0 += 2; + p1 += 2; + p2 += 2; + p3 += 2; } for (; kk < max_kk; kk++) { - pp[0] = (unsigned char)B.row(j + jj)[k0 + kk]; - pp[1] = (unsigned char)B.row(j + jj + 1)[k0 + kk]; - pp[2] = (unsigned char)B.row(j + jj + 2)[k0 + kk]; - pp[3] = (unsigned char)B.row(j + jj + 3)[k0 + kk]; + pp[0] = (unsigned char)*p0++; + pp[1] = (unsigned char)*p1++; + pp[2] = (unsigned char)*p2++; + pp[3] = (unsigned char)*p3++; pp += 4; } - pd[0] = 1.f / B_scales.row(j + jj)[g]; - pd[1] = 1.f / B_scales.row(j + jj + 1)[g]; - pd[2] = 1.f / B_scales.row(j + jj + 2)[g]; - pd[3] = 1.f / B_scales.row(j + jj + 3)[g]; + pd[0] = 1.f / *ps0++; + pd[1] = 1.f / *ps1++; + pd[2] = 1.f / *ps2++; + pd[3] = 1.f / *ps3++; pd += 4; } } @@ -202,107 +213,125 @@ static void pack_B_tile_wq_int8(const Mat& B, const Mat& B_scales, unsigned char #endif // __SSE2__ for (; jj + 1 < max_jj; jj += 2) { + const signed char* p0 = B.row(j + jj); + const signed char* p1 = B.row(j + jj + 1); + const float* ps0 = B_scales.row(j + jj); + const float* ps1 = B_scales.row(j + jj + 1); + for (int g = 0; g < block_count; g++) { - const int k0 = g * block_size; - const int max_kk = std::min(K - k0, block_size); + const int max_kk = std::min(K - g * block_size, block_size); int kk = 0; #if __AVX512VNNI__ || __AVXVNNI__ // VNNI consumes one contiguous K4 dword per output lane. for (; kk + 3 < max_kk; kk += 4) { - const signed char* p0 = B.row(j + jj) + k0 + kk; - const signed char* p1 = B.row(j + jj + 1) + k0 + kk; - __m128i _p = _mm_setr_epi32(*(const int*)p0, *(const int*)p1, 0, 0); + __m128i _p = _mm_setr_epi32(_mm_cvtsi128_si32(_mm_loadu_si32(p0)), _mm_cvtsi128_si32(_mm_loadu_si32(p1)), 0, 0); #if !__AVXVNNIINT8__ _p = _mm_add_epi8(_p, _mm_set1_epi8(127)); #endif // __AVXVNNIINT8__ _mm_storel_epi64((__m128i*)pp, _p); pp += 8; + p0 += 4; + p1 += 4; } #else #if __SSE2__ // AVX2/SSE2 consumes two K2 vectors for each real K4 region. for (; kk + 3 < max_kk; kk += 4) { - const signed char* p0 = B.row(j + jj) + k0 + kk; - const signed char* p1 = B.row(j + jj + 1) + k0 + kk; - *(unsigned short*)pp = *(const unsigned short*)p0; - *(unsigned short*)(pp + 2) = *(const unsigned short*)p1; - *(unsigned short*)(pp + 4) = *(const unsigned short*)(p0 + 2); - *(unsigned short*)(pp + 6) = *(const unsigned short*)(p1 + 2); + _mm_storeu_si16(pp, _mm_loadu_si16(p0)); + _mm_storeu_si16(pp + 2, _mm_loadu_si16(p1)); + _mm_storeu_si16(pp + 4, _mm_loadu_si16(p0 + 2)); + _mm_storeu_si16(pp + 6, _mm_loadu_si16(p1 + 2)); pp += 8; + p0 += 4; + p1 += 4; } #endif // __SSE2__ -#endif // __AVX512VNNI__ || __AVXVNNI__ \ -// K2/K1 are always signed and compact, including classic VNNI. +#endif // __AVX512VNNI__ || __AVXVNNI__ + // K2/K1 are always signed and compact, including classic VNNI. #if __SSE2__ for (; kk + 1 < max_kk; kk += 2) { - const signed char* p0 = B.row(j + jj) + k0 + kk; - const signed char* p1 = B.row(j + jj + 1) + k0 + kk; - *(unsigned short*)pp = *(const unsigned short*)p0; - *(unsigned short*)(pp + 2) = *(const unsigned short*)p1; + _mm_storeu_si16(pp, _mm_loadu_si16(p0)); + _mm_storeu_si16(pp + 2, _mm_loadu_si16(p1)); pp += 4; + p0 += 2; + p1 += 2; } #endif // __SSE2__ for (; kk < max_kk; kk++) { - pp[0] = (unsigned char)B.row(j + jj)[k0 + kk]; - pp[1] = (unsigned char)B.row(j + jj + 1)[k0 + kk]; + pp[0] = (unsigned char)*p0++; + pp[1] = (unsigned char)*p1++; pp += 2; } - pd[0] = 1.f / B_scales.row(j + jj)[g]; - pd[1] = 1.f / B_scales.row(j + jj + 1)[g]; + pd[0] = 1.f / *ps0++; + pd[1] = 1.f / *ps1++; pd += 2; } } for (; jj < max_jj; jj++) { + const signed char* p0 = B.row(j + jj); + const float* ps0 = B_scales.row(j + jj); + for (int g = 0; g < block_count; g++) { - const int k0 = g * block_size; - const int max_kk = std::min(K - k0, block_size); + const int max_kk = std::min(K - g * block_size, block_size); int kk = 0; #if __AVX512VNNI__ || __AVXVNNI__ // VNNI consumes one contiguous K4 dword per output lane. for (; kk + 3 < max_kk; kk += 4) { - const signed char* p0 = B.row(j + jj) + k0 + kk; #if !__AVXVNNIINT8__ - __m128i _p = _mm_cvtsi32_si128(*(const int*)p0); + __m128i _p = _mm_loadu_si32(p0); _p = _mm_add_epi8(_p, _mm_set1_epi8(127)); - *(int*)pp = _mm_cvtsi128_si32(_p); + _mm_storeu_si32(pp, _p); #else // __AVXVNNIINT8__ - *(int*)pp = *(const int*)p0; + _mm_storeu_si32(pp, _mm_loadu_si32(p0)); #endif // __AVXVNNIINT8__ pp += 4; + p0 += 4; } #else // AVX2/SSE2 consumes two K2 vectors for each real K4 region. for (; kk + 3 < max_kk; kk += 4) { - const signed char* p0 = B.row(j + jj) + k0 + kk; - *(int*)pp = *(const int*)p0; +#if __SSE2__ + _mm_storeu_si32(pp, _mm_loadu_si32(p0)); +#else + pp[0] = p0[0]; + pp[1] = p0[1]; + pp[2] = p0[2]; + pp[3] = p0[3]; +#endif // __SSE2__ pp += 4; + p0 += 4; } -#endif // __AVX512VNNI__ || __AVXVNNI__ \ -// K2/K1 are always signed and compact, including classic VNNI. +#endif // __AVX512VNNI__ || __AVXVNNI__ + // K2/K1 are always signed and compact, including classic VNNI. for (; kk + 1 < max_kk; kk += 2) { - const signed char* p0 = B.row(j + jj) + k0 + kk; - *(unsigned short*)pp = *(const unsigned short*)p0; +#if __SSE2__ + _mm_storeu_si16(pp, _mm_loadu_si16(p0)); +#else + pp[0] = p0[0]; + pp[1] = p0[1]; +#endif // __SSE2__ pp += 2; + p0 += 2; } for (; kk < max_kk; kk++) { - *pp++ = (unsigned char)B.row(j + jj)[k0 + kk]; + *pp++ = (unsigned char)*p0++; } - pd[0] = 1.f / B_scales.row(j + jj)[g]; + pd[0] = 1.f / *ps0++; pd += 1; } } @@ -353,15 +382,15 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales _vindex = _mm512_mullo_epi32(_vindex, _mm512_set1_epi32((int)A_hstep)); for (; ii + 15 < max_ii; ii += 16) { - signed char* outptr0 = outptr + ii * out_hstep; - float* descale_ptr0 = descale_ptr + ii * block_count; + const float* p0 = (const float*)A + (i + ii) * A_hstep; + signed char* pp0 = outptr + ii * out_hstep; + float* pd = descale_ptr + ii * block_count; for (int g = 0; g < block_count; g++) { const int k0 = g * block_size; const int max_kk = std::min(K - k0, block_size); - const float* p0 = (const float*)A + (i + ii) * A_hstep + k0; __m512 _absmax0 = _mm512_setzero_ps(); __m512 _absmax1 = _mm512_setzero_ps(); __m512 _absmax2 = _mm512_setzero_ps(); @@ -399,7 +428,7 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales __m512 _pf = _mm512_loadu_ps(p0 + A_hstep * 15 + kk_absmax); if (input_scale_ptr) { - const __m512 _s = _mm512_loadu_ps(input_scale_ptr + k0 + kk_absmax); + __m512 _s = _mm512_loadu_ps(input_scale_ptr + k0 + kk_absmax); _p0 = _mm512_mul_ps(_p0, _s); _p1 = _mm512_mul_ps(_p1, _s); _p2 = _mm512_mul_ps(_p2, _s); @@ -453,7 +482,7 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales float absmaxf = _mm512_reduce_max_ps(_absmaxf); for (; kk_absmax + 3 < max_kk; kk_absmax += 4) { - const __m128 _s = input_scale_ptr ? _mm_loadu_ps(input_scale_ptr + k0 + kk_absmax) : _mm_set1_ps(1.f); + __m128 _s = input_scale_ptr ? _mm_loadu_ps(input_scale_ptr + k0 + kk_absmax) : _mm_set1_ps(1.f); absmax0 = std::max(absmax0, _mm_reduce_max_ps(abs_ps(_mm_mul_ps(_mm_loadu_ps(p0 + kk_absmax), _s)))); absmax1 = std::max(absmax1, _mm_reduce_max_ps(abs_ps(_mm_mul_ps(_mm_loadu_ps(p0 + A_hstep + kk_absmax), _s)))); absmax2 = std::max(absmax2, _mm_reduce_max_ps(abs_ps(_mm_mul_ps(_mm_loadu_ps(p0 + A_hstep * 2 + kk_absmax), _s)))); @@ -491,28 +520,17 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales absmaxe = std::max(absmaxe, fabsf(p0[A_hstep * 14 + kk_absmax] * s)); absmaxf = std::max(absmaxf, fabsf(p0[A_hstep * 15 + kk_absmax] * s)); } - const __m512 _absmax = _mm512_setr_ps(absmax0, absmax1, absmax2, absmax3, absmax4, absmax5, absmax6, absmax7, absmax8, absmax9, absmaxa, absmaxb, absmaxc, absmaxd, absmaxe, absmaxf); + __m512 _absmax = _mm512_setr_ps(absmax0, absmax1, absmax2, absmax3, absmax4, absmax5, absmax6, absmax7, absmax8, absmax9, absmaxa, absmaxb, absmaxc, absmaxd, absmaxe, absmaxf); - const __m512 _descale = _mm512_div_ps(_absmax, _mm512_set1_ps(127.f)); - const __m256 _absmax0_fp32 = _mm512_castps512_ps256(_absmax); - const __m256 _absmax1_fp32 = _mm256_castpd_ps(_mm512_extractf64x4_pd(_mm512_castps_pd(_absmax), 1)); - const __m512d _absmax0_fp64 = _mm512_cvtps_pd(_absmax0_fp32); - const __m512d _absmax1_fp64 = _mm512_cvtps_pd(_absmax1_fp32); - const __mmask8 _nonzero0 = _mm512_cmp_pd_mask(_absmax0_fp64, _mm512_setzero_pd(), _CMP_NEQ_OQ); - const __mmask8 _nonzero1 = _mm512_cmp_pd_mask(_absmax1_fp64, _mm512_setzero_pd(), _CMP_NEQ_OQ); - const __m256 _scale0 = _mm512_cvtpd_ps(_mm512_maskz_div_pd(_nonzero0, _mm512_set1_pd(127.0), _absmax0_fp64)); - const __m256 _scale1 = _mm512_cvtpd_ps(_mm512_maskz_div_pd(_nonzero1, _mm512_set1_pd(127.0), _absmax1_fp64)); - const __m512 _scale = combine8x2_ps(_scale0, _scale1); - _mm512_storeu_ps(descale_ptr0 + g * 16, _descale); + __m512 _descale = _mm512_div_ps(_absmax, _mm512_set1_ps(127.f)); + __mmask16 _nonzero = _mm512_cmp_ps_mask(_absmax, _mm512_setzero_ps(), _CMP_NEQ_OQ); + __m512 _scale = _mm512_maskz_div_ps(_nonzero, _mm512_set1_ps(127.f), _absmax); + _mm512_storeu_ps(pd, _descale); #if __AVX512VNNI__ __m512i _w_shift = _mm512_setzero_si512(); #endif -#if __AVX512VNNI__ - signed char* pp = outptr0 + (k0 + g * 4) * 16; -#else - signed char* pp = outptr0 + k0 * 16; -#endif + signed char* pp = pp0; int kk = 0; for (; kk + 15 < max_kk; kk += 16) { @@ -534,7 +552,7 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales __m512 _pf = _mm512_loadu_ps(p0 + A_hstep * 15 + kk); if (input_scale_ptr) { - const __m512 _s = _mm512_loadu_ps(input_scale_ptr + k0 + kk); + __m512 _s = _mm512_loadu_ps(input_scale_ptr + k0 + kk); _p0 = _mm512_mul_ps(_p0, _s); _p1 = _mm512_mul_ps(_p1, _s); _p2 = _mm512_mul_ps(_p2, _s); @@ -551,45 +569,6 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales _pd = _mm512_mul_ps(_pd, _s); _pe = _mm512_mul_ps(_pe, _s); _pf = _mm512_mul_ps(_pf, _s); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+x"(_p0), "+x"(_p1), "+x"(_p2), "+x"(_p3), "+x"(_p4), "+x"(_p5), "+x"(_p6), "+x"(_p7)); - asm volatile("" - : "+x"(_p8), "+x"(_p9), "+x"(_pa), "+x"(_pb), "+x"(_pc), "+x"(_pd), "+x"(_pe), "+x"(_pf)); -#else - volatile __m512 _p0_ordered = _p0; - volatile __m512 _p1_ordered = _p1; - volatile __m512 _p2_ordered = _p2; - volatile __m512 _p3_ordered = _p3; - volatile __m512 _p4_ordered = _p4; - volatile __m512 _p5_ordered = _p5; - volatile __m512 _p6_ordered = _p6; - volatile __m512 _p7_ordered = _p7; - volatile __m512 _p8_ordered = _p8; - volatile __m512 _p9_ordered = _p9; - volatile __m512 _pa_ordered = _pa; - volatile __m512 _pb_ordered = _pb; - volatile __m512 _pc_ordered = _pc; - volatile __m512 _pd_ordered = _pd; - volatile __m512 _pe_ordered = _pe; - volatile __m512 _pf_ordered = _pf; - _p0 = _p0_ordered; - _p1 = _p1_ordered; - _p2 = _p2_ordered; - _p3 = _p3_ordered; - _p4 = _p4_ordered; - _p5 = _p5_ordered; - _p6 = _p6_ordered; - _p7 = _p7_ordered; - _p8 = _p8_ordered; - _p9 = _p9_ordered; - _pa = _pa_ordered; - _pb = _pb_ordered; - _pc = _pc_ordered; - _pd = _pd_ordered; - _pe = _pe_ordered; - _pf = _pf_ordered; -#endif } transpose16x16_ps(_p0, _p1, _p2, _p3, _p4, _p5, _p6, _p7, _p8, _p9, _pa, _pb, _pc, _pd, _pe, _pf); __m128i _q0 = float2int8_avx512(_mm512_mul_ps(_p0, _scale)); @@ -680,7 +659,7 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales __m128 _pf = _mm_loadu_ps(p0 + A_hstep * 15 + kk); if (input_scale_ptr) { - const __m128 _s = _mm_loadu_ps(input_scale_ptr + k0 + kk); + __m128 _s = _mm_loadu_ps(input_scale_ptr + k0 + kk); _p0 = _mm_mul_ps(_p0, _s); _p1 = _mm_mul_ps(_p1, _s); _p2 = _mm_mul_ps(_p2, _s); @@ -697,45 +676,6 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales _pd = _mm_mul_ps(_pd, _s); _pe = _mm_mul_ps(_pe, _s); _pf = _mm_mul_ps(_pf, _s); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+x"(_p0), "+x"(_p1), "+x"(_p2), "+x"(_p3), "+x"(_p4), "+x"(_p5), "+x"(_p6), "+x"(_p7)); - asm volatile("" - : "+x"(_p8), "+x"(_p9), "+x"(_pa), "+x"(_pb), "+x"(_pc), "+x"(_pd), "+x"(_pe), "+x"(_pf)); -#else - volatile __m128 _p0_ordered = _p0; - volatile __m128 _p1_ordered = _p1; - volatile __m128 _p2_ordered = _p2; - volatile __m128 _p3_ordered = _p3; - volatile __m128 _p4_ordered = _p4; - volatile __m128 _p5_ordered = _p5; - volatile __m128 _p6_ordered = _p6; - volatile __m128 _p7_ordered = _p7; - volatile __m128 _p8_ordered = _p8; - volatile __m128 _p9_ordered = _p9; - volatile __m128 _pa_ordered = _pa; - volatile __m128 _pb_ordered = _pb; - volatile __m128 _pc_ordered = _pc; - volatile __m128 _pd_ordered = _pd; - volatile __m128 _pe_ordered = _pe; - volatile __m128 _pf_ordered = _pf; - _p0 = _p0_ordered; - _p1 = _p1_ordered; - _p2 = _p2_ordered; - _p3 = _p3_ordered; - _p4 = _p4_ordered; - _p5 = _p5_ordered; - _p6 = _p6_ordered; - _p7 = _p7_ordered; - _p8 = _p8_ordered; - _p9 = _p9_ordered; - _pa = _pa_ordered; - _pb = _pb_ordered; - _pc = _pc_ordered; - _pd = _pd_ordered; - _pe = _pe_ordered; - _pf = _pf_ordered; -#endif } __m512 _t0 = combine4x4_ps(_p0, _p4, _p8, _pc); __m512 _t1 = combine4x4_ps(_p1, _p5, _p9, _pd); @@ -755,7 +695,7 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales __m128i _q3 = float2int8_avx512(_mm512_mul_ps(_t3, _scale)); #if __AVX512VNNI__ transpose16x4_epi8(_q0, _q1, _q2, _q3); - const __m512i _q = combine4x4_epi32(_q0, _q1, _q2, _q3); + __m512i _q = combine4x4_epi32(_q0, _q1, _q2, _q3); _mm512_storeu_si512((__m512i*)pp, _q); _w_shift = _mm512_dpbusd_epi32(_w_shift, _mm512_set1_epi8(127), _q); #else @@ -775,57 +715,49 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales #endif // __AVX512VNNI__ for (; kk + 1 < max_kk; kk += 2) { - __m512 _p0 = _mm512_i32gather_ps(_vindex, (const float*)A + (i + ii) * A_hstep + k0 + kk, sizeof(float)); - __m512 _p1 = _mm512_i32gather_ps(_vindex, (const float*)A + (i + ii) * A_hstep + k0 + kk + 1, sizeof(float)); + __m512 _p0 = _mm512_i32gather_ps(_vindex, p0 + kk, sizeof(float)); + __m512 _p1 = _mm512_i32gather_ps(_vindex, p0 + kk + 1, sizeof(float)); if (input_scale_ptr) { _p0 = _mm512_mul_ps(_p0, _mm512_set1_ps(input_scale_ptr[k0 + kk])); _p1 = _mm512_mul_ps(_p1, _mm512_set1_ps(input_scale_ptr[k0 + kk + 1])); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+x"(_p0), "+x"(_p1)); -#else - volatile __m512 _p0_ordered = _p0; - volatile __m512 _p1_ordered = _p1; - _p0 = _p0_ordered; - _p1 = _p1_ordered; -#endif } - const __m128i _q0 = float2int8_avx512(_mm512_mul_ps(_p0, _scale)); - const __m128i _q1 = float2int8_avx512(_mm512_mul_ps(_p1, _scale)); + __m128i _q0 = float2int8_avx512(_mm512_mul_ps(_p0, _scale)); + __m128i _q1 = float2int8_avx512(_mm512_mul_ps(_p1, _scale)); _mm_storeu_si128((__m128i*)pp, _mm_unpacklo_epi8(_q0, _q1)); _mm_storeu_si128((__m128i*)(pp + 16), _mm_unpackhi_epi8(_q0, _q1)); pp += 32; } if (kk < max_kk) { - __m512 _p = _mm512_i32gather_ps(_vindex, (const float*)A + (i + ii) * A_hstep + k0 + kk, sizeof(float)); + __m512 _p = _mm512_i32gather_ps(_vindex, p0 + kk, sizeof(float)); if (input_scale_ptr) { _p = _mm512_mul_ps(_p, _mm512_set1_ps(input_scale_ptr[k0 + kk])); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+x"(_p)); -#else - volatile __m512 _p_ordered = _p; - _p = _p_ordered; -#endif } _mm_storeu_si128((__m128i*)pp, float2int8_avx512(_mm512_mul_ps(_p, _scale))); } + + p0 += max_kk; + pd += 16; +#if __AVX512VNNI__ + pp0 += (max_kk + (max_kk >= 4 ? 4 : 0)) * 16; +#else + pp0 += max_kk * 16; +#endif } } #endif // __AVX512F__ for (; ii + 7 < max_ii; ii += 8) { - signed char* outptr0 = outptr + ii * out_hstep; - float* descale_ptr0 = descale_ptr + ii * block_count; + const float* p0 = (const float*)A + (i + ii) * A_hstep; + signed char* pp0 = outptr + ii * out_hstep; + float* pd = descale_ptr + ii * block_count; for (int g = 0; g < block_count; g++) { const int k0 = g * block_size; const int max_kk = std::min(K - k0, block_size); - const float* p0 = (const float*)A + (i + ii) * A_hstep + k0; __m256 _absmax0 = _mm256_setzero_ps(); __m256 _absmax1 = _mm256_setzero_ps(); @@ -848,7 +780,7 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales __m256 _p7 = _mm256_loadu_ps(p0 + A_hstep * 7 + kk); if (input_scale_ptr) { - const __m256 _s = _mm256_loadu_ps(input_scale_ptr + k0 + kk); + __m256 _s = _mm256_loadu_ps(input_scale_ptr + k0 + kk); _p0 = _mm256_mul_ps(_p0, _s); _p1 = _mm256_mul_ps(_p1, _s); _p2 = _mm256_mul_ps(_p2, _s); @@ -889,37 +821,31 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales absmax7 = std::max(absmax7, fabsf(p0[A_hstep * 7 + kk] * s)); } - const __m256 _absmax = _mm256_setr_ps(absmax0, absmax1, absmax2, absmax3, absmax4, absmax5, absmax6, absmax7); - const __m256 _descale = _mm256_div_ps(_absmax, _mm256_set1_ps(127.f)); - const __m256 _nonzero = _mm256_cmp_ps(_absmax, _mm256_setzero_ps(), _CMP_NEQ_OQ); - const __m256 _absmax_nonzero = _mm256_blendv_ps(_mm256_set1_ps(1.f), _absmax, _nonzero); - const __m256d _absmax0_fp64 = _mm256_cvtps_pd(_mm256_castps256_ps128(_absmax_nonzero)); - const __m256d _absmax1_fp64 = _mm256_cvtps_pd(_mm256_extractf128_ps(_absmax_nonzero, 1)); - const __m128 _scale0 = _mm256_cvtpd_ps(_mm256_div_pd(_mm256_set1_pd(127.0), _absmax0_fp64)); - const __m128 _scale1 = _mm256_cvtpd_ps(_mm256_div_pd(_mm256_set1_pd(127.0), _absmax1_fp64)); - const __m256 _scale = _mm256_and_ps(combine4x2_ps(_scale0, _scale1), _nonzero); - _mm256_storeu_ps(descale_ptr0 + g * 8, _descale); + __m256 _absmax = _mm256_setr_ps(absmax0, absmax1, absmax2, absmax3, absmax4, absmax5, absmax6, absmax7); + __m256 _descale = _mm256_div_ps(_absmax, _mm256_set1_ps(127.f)); + __m256 _nonzero = _mm256_cmp_ps(_absmax, _mm256_setzero_ps(), _CMP_NEQ_OQ); + __m256 _absmax_nonzero = _mm256_blendv_ps(_mm256_set1_ps(1.f), _absmax, _nonzero); + __m256 _scale = _mm256_and_ps(_mm256_div_ps(_mm256_set1_ps(127.f), _absmax_nonzero), _nonzero); + _mm256_storeu_ps(pd, _descale); + signed char* pp = pp0; #if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) __m256i _w_shift = _mm256_setzero_si256(); - signed char* pp = outptr0 + (k0 + g * 4) * 8; -#else - signed char* pp = outptr0 + k0 * 8; #endif kk = 0; for (; kk + 3 < max_kk; kk += 4) { - __m128 _p0 = _mm_loadu_ps((const float*)A + (i + ii) * A_hstep + k0 + kk); - __m128 _p1 = _mm_loadu_ps((const float*)A + (i + ii + 1) * A_hstep + k0 + kk); - __m128 _p2 = _mm_loadu_ps((const float*)A + (i + ii + 2) * A_hstep + k0 + kk); - __m128 _p3 = _mm_loadu_ps((const float*)A + (i + ii + 3) * A_hstep + k0 + kk); - __m128 _p4 = _mm_loadu_ps((const float*)A + (i + ii + 4) * A_hstep + k0 + kk); - __m128 _p5 = _mm_loadu_ps((const float*)A + (i + ii + 5) * A_hstep + k0 + kk); - __m128 _p6 = _mm_loadu_ps((const float*)A + (i + ii + 6) * A_hstep + k0 + kk); - __m128 _p7 = _mm_loadu_ps((const float*)A + (i + ii + 7) * A_hstep + k0 + kk); + __m128 _p0 = _mm_loadu_ps(p0 + kk); + __m128 _p1 = _mm_loadu_ps(p0 + A_hstep + kk); + __m128 _p2 = _mm_loadu_ps(p0 + A_hstep * 2 + kk); + __m128 _p3 = _mm_loadu_ps(p0 + A_hstep * 3 + kk); + __m128 _p4 = _mm_loadu_ps(p0 + A_hstep * 4 + kk); + __m128 _p5 = _mm_loadu_ps(p0 + A_hstep * 5 + kk); + __m128 _p6 = _mm_loadu_ps(p0 + A_hstep * 6 + kk); + __m128 _p7 = _mm_loadu_ps(p0 + A_hstep * 7 + kk); if (input_scale_ptr) { - const __m128 _s = _mm_loadu_ps(input_scale_ptr + k0 + kk); + __m128 _s = _mm_loadu_ps(input_scale_ptr + k0 + kk); _p0 = _mm_mul_ps(_p0, _s); _p1 = _mm_mul_ps(_p1, _s); _p2 = _mm_mul_ps(_p2, _s); @@ -928,27 +854,6 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales _p5 = _mm_mul_ps(_p5, _s); _p6 = _mm_mul_ps(_p6, _s); _p7 = _mm_mul_ps(_p7, _s); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+x"(_p0), "+x"(_p1), "+x"(_p2), "+x"(_p3), "+x"(_p4), "+x"(_p5), "+x"(_p6), "+x"(_p7)); -#else - volatile __m128 _p0_ordered = _p0; - volatile __m128 _p1_ordered = _p1; - volatile __m128 _p2_ordered = _p2; - volatile __m128 _p3_ordered = _p3; - volatile __m128 _p4_ordered = _p4; - volatile __m128 _p5_ordered = _p5; - volatile __m128 _p6_ordered = _p6; - volatile __m128 _p7_ordered = _p7; - _p0 = _p0_ordered; - _p1 = _p1_ordered; - _p2 = _p2_ordered; - _p3 = _p3_ordered; - _p4 = _p4_ordered; - _p5 = _p5_ordered; - _p6 = _p6_ordered; - _p7 = _p7_ordered; -#endif } __m256 _t0 = combine4x2_ps(_p0, _p4); @@ -975,7 +880,7 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales #if __AVX512VNNI__ || __AVXVNNI__ _q0 = _mm_unpacklo_epi16(_q01, _q23); _q1 = _mm_unpackhi_epi16(_q01, _q23); - const __m256i _q = combine4x2_epi32(_q0, _q1); + __m256i _q = combine4x2_epi32(_q0, _q1); _mm256_storeu_si256((__m256i*)pp, _q); #if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) _w_shift = _mm256_comp_dpbusd_epi32(_w_shift, _mm256_set1_epi8(127), _q); @@ -997,26 +902,17 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales { __m256i _vindex = _mm256_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7); _vindex = _mm256_mullo_epi32(_vindex, _mm256_set1_epi32((int)A_hstep)); - __m256 _p0 = _mm256_i32gather_ps((const float*)A + (i + ii) * A_hstep + k0 + kk, _vindex, sizeof(float)); - __m256 _p1 = _mm256_i32gather_ps((const float*)A + (i + ii) * A_hstep + k0 + kk + 1, _vindex, sizeof(float)); + __m256 _p0 = _mm256_i32gather_ps(p0 + kk, _vindex, sizeof(float)); + __m256 _p1 = _mm256_i32gather_ps(p0 + kk + 1, _vindex, sizeof(float)); if (input_scale_ptr) { _p0 = _mm256_mul_ps(_p0, _mm256_set1_ps(input_scale_ptr[k0 + kk])); _p1 = _mm256_mul_ps(_p1, _mm256_set1_ps(input_scale_ptr[k0 + kk + 1])); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+x"(_p0), "+x"(_p1)); -#else - volatile __m256 _p0_ordered = _p0; - volatile __m256 _p1_ordered = _p1; - _p0 = _p0_ordered; - _p1 = _p1_ordered; -#endif } _p0 = _mm256_mul_ps(_p0, _scale); _p1 = _mm256_mul_ps(_p1, _scale); __m128i _q = float2int8_avx(_p0, _p1); - const __m128i _si = _mm_setr_epi8(0, 8, 1, 9, 2, 10, 3, 11, 4, 12, 5, 13, 6, 14, 7, 15); + __m128i _si = _mm_setr_epi8(0, 8, 1, 9, 2, 10, 3, 11, 4, 12, 5, 13, 6, 14, 7, 15); _q = _mm_shuffle_epi8(_q, _si); _mm_storeu_si128((__m128i*)pp, _q); pp += 16; @@ -1025,34 +921,35 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales { __m256i _vindex = _mm256_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7); _vindex = _mm256_mullo_epi32(_vindex, _mm256_set1_epi32((int)A_hstep)); - __m256 _p = _mm256_i32gather_ps((const float*)A + (i + ii) * A_hstep + k0 + kk, _vindex, sizeof(float)); + __m256 _p = _mm256_i32gather_ps(p0 + kk, _vindex, sizeof(float)); if (input_scale_ptr) { _p = _mm256_mul_ps(_p, _mm256_set1_ps(input_scale_ptr[k0 + kk])); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+x"(_p)); -#else - volatile __m256 _p_ordered = _p; - _p = _p_ordered; -#endif } - *(int64_t*)pp = float2int8_avx(_mm256_mul_ps(_p, _scale)); + _mm_storeu_si64(pp, _mm_cvtsi64_si128(float2int8_avx(_mm256_mul_ps(_p, _scale)))); pp += 8; } + + p0 += max_kk; + pd += 8; +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + pp0 += (max_kk + (max_kk >= 4 ? 4 : 0)) * 8; +#else + pp0 += max_kk * 8; +#endif } } #endif // __AVX2__ for (; ii + 3 < max_ii; ii += 4) { - signed char* outptr0 = outptr + ii * out_hstep; - float* descale_ptr0 = descale_ptr + ii * block_count; + const float* p0 = (const float*)A + (i + ii) * A_hstep; + signed char* pp0 = outptr + ii * out_hstep; + float* pd = descale_ptr + ii * block_count; for (int g = 0; g < block_count; g++) { const int k0 = g * block_size; const int max_kk = std::min(K - k0, block_size); - const float* p0 = (const float*)A + (i + ii) * A_hstep + k0; __m128 _absmax0 = _mm_setzero_ps(); __m128 _absmax1 = _mm_setzero_ps(); @@ -1067,7 +964,7 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales __m128 _p3 = _mm_loadu_ps(p0 + A_hstep * 3 + kk); if (input_scale_ptr) { - const __m128 _s = _mm_loadu_ps(input_scale_ptr + k0 + kk); + __m128 _s = _mm_loadu_ps(input_scale_ptr + k0 + kk); _p0 = _mm_mul_ps(_p0, _s); _p1 = _mm_mul_ps(_p1, _s); _p2 = _mm_mul_ps(_p2, _s); @@ -1092,22 +989,16 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales absmax3 = std::max(absmax3, fabsf(p0[A_hstep * 3 + kk] * s)); } - const __m128 _absmax = _mm_setr_ps(absmax0, absmax1, absmax2, absmax3); - const __m128 _descale = _mm_div_ps(_absmax, _mm_set1_ps(127.f)); - const __m128 _nonzero = _mm_cmpneq_ps(_absmax, _mm_setzero_ps()); - const __m128 _absmax_nonzero = _mm_or_ps(_mm_and_ps(_absmax, _nonzero), _mm_andnot_ps(_nonzero, _mm_set1_ps(1.f))); - const __m128d _absmax01_fp64 = _mm_cvtps_pd(_absmax_nonzero); - const __m128d _absmax23_fp64 = _mm_cvtps_pd(_mm_movehl_ps(_absmax_nonzero, _absmax_nonzero)); - const __m128 _scale01 = _mm_cvtpd_ps(_mm_div_pd(_mm_set1_pd(127.0), _absmax01_fp64)); - const __m128 _scale23 = _mm_cvtpd_ps(_mm_div_pd(_mm_set1_pd(127.0), _absmax23_fp64)); - const __m128 _scale = _mm_and_ps(_mm_movelh_ps(_scale01, _scale23), _nonzero); - _mm_storeu_ps(descale_ptr0 + g * 4, _descale); + __m128 _absmax = _mm_setr_ps(absmax0, absmax1, absmax2, absmax3); + __m128 _descale = _mm_div_ps(_absmax, _mm_set1_ps(127.f)); + __m128 _nonzero = _mm_cmpneq_ps(_absmax, _mm_setzero_ps()); + __m128 _absmax_nonzero = _mm_or_ps(_mm_and_ps(_absmax, _nonzero), _mm_andnot_ps(_nonzero, _mm_set1_ps(1.f))); + __m128 _scale = _mm_and_ps(_mm_div_ps(_mm_set1_ps(127.f), _absmax_nonzero), _nonzero); + _mm_storeu_ps(pd, _descale); + signed char* pp = pp0; #if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) __m128i _w_shift = _mm_setzero_si128(); - signed char* pp = outptr0 + (k0 + g * 4) * 4; -#else - signed char* pp = outptr0 + k0 * 4; #endif kk = 0; for (; kk + 3 < max_kk; kk += 4) @@ -1118,38 +1009,25 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales __m128 _p3 = _mm_loadu_ps(p0 + A_hstep * 3 + kk); if (input_scale_ptr) { - const __m128 _s = _mm_loadu_ps(input_scale_ptr + k0 + kk); + __m128 _s = _mm_loadu_ps(input_scale_ptr + k0 + kk); _p0 = _mm_mul_ps(_p0, _s); _p1 = _mm_mul_ps(_p1, _s); _p2 = _mm_mul_ps(_p2, _s); _p3 = _mm_mul_ps(_p3, _s); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+x"(_p0), "+x"(_p1), "+x"(_p2), "+x"(_p3)); -#else - volatile __m128 _p0_ordered = _p0; - volatile __m128 _p1_ordered = _p1; - volatile __m128 _p2_ordered = _p2; - volatile __m128 _p3_ordered = _p3; - _p0 = _p0_ordered; - _p1 = _p1_ordered; - _p2 = _p2_ordered; - _p3 = _p3_ordered; -#endif - } - const __m128i _q0 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p0, _mm_shuffle_ps(_scale, _scale, _MM_SHUFFLE(0, 0, 0, 0))))); - const __m128i _q1 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p1, _mm_shuffle_ps(_scale, _scale, _MM_SHUFFLE(1, 1, 1, 1))))); - const __m128i _q2 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p2, _mm_shuffle_ps(_scale, _scale, _MM_SHUFFLE(2, 2, 2, 2))))); - const __m128i _q3 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p3, _mm_shuffle_ps(_scale, _scale, _MM_SHUFFLE(3, 3, 3, 3))))); + } + __m128i _q0 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p0, _mm_shuffle_ps(_scale, _scale, _MM_SHUFFLE(0, 0, 0, 0))))); + __m128i _q1 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p1, _mm_shuffle_ps(_scale, _scale, _MM_SHUFFLE(1, 1, 1, 1))))); + __m128i _q2 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p2, _mm_shuffle_ps(_scale, _scale, _MM_SHUFFLE(2, 2, 2, 2))))); + __m128i _q3 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p3, _mm_shuffle_ps(_scale, _scale, _MM_SHUFFLE(3, 3, 3, 3))))); #if __AVX512VNNI__ || __AVXVNNI__ - const __m128i _q = _mm_unpacklo_epi64(_mm_unpacklo_epi32(_q0, _q1), _mm_unpacklo_epi32(_q2, _q3)); + __m128i _q = _mm_unpacklo_epi64(_mm_unpacklo_epi32(_q0, _q1), _mm_unpacklo_epi32(_q2, _q3)); _mm_storeu_si128((__m128i*)pp, _q); #if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) _w_shift = _mm_comp_dpbusd_epi32(_w_shift, _mm_set1_epi8(127), _q); #endif #else - const __m128i _q01 = _mm_unpacklo_epi16(_q0, _q1); - const __m128i _q23 = _mm_unpacklo_epi16(_q2, _q3); + __m128i _q01 = _mm_unpacklo_epi16(_q0, _q1); + __m128i _q23 = _mm_unpacklo_epi16(_q2, _q3); _mm_storeu_si128((__m128i*)pp, _mm_unpacklo_epi32(_q01, _q23)); #endif // __AVX512VNNI__ || __AVXVNNI__ pp += 16; @@ -1169,18 +1047,9 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales { _p0 = _mm_mul_ps(_p0, _mm_set1_ps(input_scale_ptr[k0 + kk])); _p1 = _mm_mul_ps(_p1, _mm_set1_ps(input_scale_ptr[k0 + kk + 1])); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+x"(_p0), "+x"(_p1)); -#else - volatile __m128 _p0_ordered = _p0; - volatile __m128 _p1_ordered = _p1; - _p0 = _p0_ordered; - _p1 = _p1_ordered; -#endif } - const __m128i _q0 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p0, _scale))); - const __m128i _q1 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p1, _scale))); + __m128i _q0 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p0, _scale))); + __m128i _q1 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p1, _scale))); _mm_storel_epi64((__m128i*)pp, _mm_unpacklo_epi8(_q0, _q1)); pp += 8; } @@ -1190,31 +1059,32 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales if (input_scale_ptr) { _p = _mm_mul_ps(_p, _mm_set1_ps(input_scale_ptr[k0 + kk])); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+x"(_p)); -#else - volatile __m128 _p_ordered = _p; - _p = _p_ordered; -#endif } - *(int*)pp = float2int8_sse(_mm_mul_ps(_p, _scale)); + _mm_storeu_si32(pp, _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p, _scale)))); pp += 4; } + + p0 += max_kk; + pd += 4; +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + pp0 += (max_kk + (max_kk >= 4 ? 4 : 0)) * 4; +#else + pp0 += max_kk * 4; +#endif } } #endif // __SSE2__ for (; ii + 1 < max_ii; ii += 2) { #if __SSE2__ - signed char* outptr0 = outptr + ii * out_hstep; - float* descale_ptr0 = descale_ptr + ii * block_count; + const float* p0 = (const float*)A + (i + ii) * A_hstep; + signed char* pp0 = outptr + ii * out_hstep; + float* pd = descale_ptr + ii * block_count; for (int g = 0; g < block_count; g++) { const int k0 = g * block_size; const int max_kk = std::min(K - k0, block_size); - const float* p0 = (const float*)A + (i + ii) * A_hstep + k0; __m128 _absmax0 = _mm_setzero_ps(); __m128 _absmax1 = _mm_setzero_ps(); @@ -1225,7 +1095,7 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales __m128 _p1 = _mm_loadu_ps(p0 + A_hstep + kk); if (input_scale_ptr) { - const __m128 _s = _mm_loadu_ps(input_scale_ptr + k0 + kk); + __m128 _s = _mm_loadu_ps(input_scale_ptr + k0 + kk); _p0 = _mm_mul_ps(_p0, _s); _p1 = _mm_mul_ps(_p1, _s); } @@ -1242,18 +1112,16 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales absmax1 = std::max(absmax1, fabsf(p0[A_hstep + kk] * s)); } - const __m128 _absmax = _mm_setr_ps(absmax0, absmax1, 0.f, 0.f); - const __m128 _descale = _mm_div_ps(_absmax, _mm_set1_ps(127.f)); - const __m128 _nonzero = _mm_cmpneq_ps(_absmax, _mm_setzero_ps()); - const __m128 _absmax_nonzero = _mm_or_ps(_mm_and_ps(_absmax, _nonzero), _mm_andnot_ps(_nonzero, _mm_set1_ps(1.f))); - const __m128 _scale = _mm_and_ps(_mm_cvtpd_ps(_mm_div_pd(_mm_set1_pd(127.0), _mm_cvtps_pd(_absmax_nonzero))), _nonzero); - _mm_storel_pi((__m64*)(descale_ptr0 + g * 2), _descale); + __m128 _absmax = _mm_setr_ps(absmax0, absmax1, 0.f, 0.f); + __m128 _descale = _mm_div_ps(_absmax, _mm_set1_ps(127.f)); + __m128 _nonzero = _mm_cmpneq_ps(_absmax, _mm_setzero_ps()); + __m128 _absmax_nonzero = _mm_or_ps(_mm_and_ps(_absmax, _nonzero), _mm_andnot_ps(_nonzero, _mm_set1_ps(1.f))); + __m128 _scale = _mm_and_ps(_mm_div_ps(_mm_set1_ps(127.f), _absmax_nonzero), _nonzero); + _mm_storel_pi((__m64*)pd, _descale); + signed char* pp = pp0; #if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) __m128i _w_shift = _mm_setzero_si128(); - signed char* pp = outptr0 + (k0 + g * 4) * 2; -#else - signed char* pp = outptr0 + k0 * 2; #endif kk = 0; for (; kk + 3 < max_kk; kk += 4) @@ -1262,23 +1130,14 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales __m128 _p1 = _mm_loadu_ps(p0 + A_hstep + kk); if (input_scale_ptr) { - const __m128 _s = _mm_loadu_ps(input_scale_ptr + k0 + kk); + __m128 _s = _mm_loadu_ps(input_scale_ptr + k0 + kk); _p0 = _mm_mul_ps(_p0, _s); _p1 = _mm_mul_ps(_p1, _s); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+x"(_p0), "+x"(_p1)); -#else - volatile __m128 _p0_ordered = _p0; - volatile __m128 _p1_ordered = _p1; - _p0 = _p0_ordered; - _p1 = _p1_ordered; -#endif } - const __m128i _q0 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p0, _mm_shuffle_ps(_scale, _scale, _MM_SHUFFLE(0, 0, 0, 0))))); - const __m128i _q1 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p1, _mm_shuffle_ps(_scale, _scale, _MM_SHUFFLE(1, 1, 1, 1))))); + __m128i _q0 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p0, _mm_shuffle_ps(_scale, _scale, _MM_SHUFFLE(0, 0, 0, 0))))); + __m128i _q1 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p1, _mm_shuffle_ps(_scale, _scale, _MM_SHUFFLE(1, 1, 1, 1))))); #if __AVX512VNNI__ || __AVXVNNI__ - const __m128i _q = _mm_unpacklo_epi32(_q0, _q1); + __m128i _q = _mm_unpacklo_epi32(_q0, _q1); _mm_storel_epi64((__m128i*)pp, _q); #if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) _w_shift = _mm_comp_dpbusd_epi32(_w_shift, _mm_set1_epi8(127), _q); @@ -1303,19 +1162,10 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales { _p0 = _mm_mul_ps(_p0, _mm_set1_ps(input_scale_ptr[k0 + kk])); _p1 = _mm_mul_ps(_p1, _mm_set1_ps(input_scale_ptr[k0 + kk + 1])); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+x"(_p0), "+x"(_p1)); -#else - volatile __m128 _p0_ordered = _p0; - volatile __m128 _p1_ordered = _p1; - _p0 = _p0_ordered; - _p1 = _p1_ordered; -#endif } - const __m128i _q0 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p0, _scale))); - const __m128i _q1 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p1, _scale))); - *(int*)pp = _mm_cvtsi128_si32(_mm_unpacklo_epi8(_q0, _q1)); + __m128i _q0 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p0, _scale))); + __m128i _q1 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p1, _scale))); + _mm_storeu_si32(pp, _mm_unpacklo_epi8(_q0, _q1)); pp += 4; } for (; kk < max_kk; kk++) @@ -1324,22 +1174,23 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales if (input_scale_ptr) { _p = _mm_mul_ps(_p, _mm_set1_ps(input_scale_ptr[k0 + kk])); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+x"(_p)); -#else - volatile __m128 _p_ordered = _p; - _p = _p_ordered; -#endif } - *(unsigned short*)pp = (unsigned short)float2int8_sse(_mm_mul_ps(_p, _scale)); + _mm_storeu_si16(pp, _mm_cvtsi32_si128((unsigned short)float2int8_sse(_mm_mul_ps(_p, _scale)))); pp += 2; } + + p0 += max_kk; + pd += 2; +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + pp0 += (max_kk + (max_kk >= 4 ? 4 : 0)) * 2; +#else + pp0 += max_kk * 2; +#endif } #else const float* p0 = (const float*)A + (i + ii) * A_hstep; - signed char* outptr0 = outptr + ii * out_hstep; - float* descale_ptr0 = descale_ptr + ii * block_count; + signed char* pp0 = outptr + ii * out_hstep; + float* pd = descale_ptr + ii * block_count; for (int g = 0; g < block_count; g++) { const int k0 = g * block_size; @@ -1348,8 +1199,8 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales float absmax1 = 0.f; for (int kk = 0; kk < max_kk; kk++) { - float v0 = p0[k0 + kk]; - float v1 = p0[A_hstep + k0 + kk]; + float v0 = p0[kk]; + float v1 = p0[A_hstep + kk]; if (input_scale_ptr) { const float s = input_scale_ptr[k0 + kk]; @@ -1364,54 +1215,48 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales float scale1 = 0.f; if (absmax0 != 0.f) { - volatile double scale_fp64 = 127.0 / (double)absmax0; - scale0 = (float)scale_fp64; + scale0 = 127.f / absmax0; } if (absmax1 != 0.f) { - volatile double scale_fp64 = 127.0 / (double)absmax1; - scale1 = (float)scale_fp64; + scale1 = 127.f / absmax1; } - descale_ptr0[g * 2] = absmax0 / 127.f; - descale_ptr0[g * 2 + 1] = absmax1 / 127.f; + pd[0] = absmax0 / 127.f; + pd[1] = absmax1 / 127.f; - signed char* pp = outptr0 + k0 * 2; + signed char* pp = pp0; for (int kk = 0; kk < max_kk; kk++) { - float v0 = p0[k0 + kk]; - float v1 = p0[A_hstep + k0 + kk]; + float v0 = p0[kk]; + float v1 = p0[A_hstep + kk]; if (input_scale_ptr) { const float s = input_scale_ptr[k0 + kk]; v0 *= s; v1 *= s; - volatile float v0_ordered = v0; - volatile float v1_ordered = v1; - v0 = v0_ordered; - v1 = v1_ordered; } pp[0] = float2int8(v0 * scale0); pp[1] = float2int8(v1 * scale1); pp += 2; } + + p0 += max_kk; + pp0 += max_kk * 2; + pd += 2; } #endif // __SSE2__ } for (; ii < max_ii; ii++) { - const float* ptrA = (const float*)A + (i + ii) * A_hstep; - signed char* outptr0 = outptr + ii * out_hstep; - float* descale_ptr0 = descale_ptr + ii * block_count; + const float* p0 = (const float*)A + (i + ii) * A_hstep; + signed char* pp0 = outptr + ii * out_hstep; + float* pd = descale_ptr + ii * block_count; for (int g = 0; g < block_count; g++) { const int k0 = g * block_size; const int max_kk = std::min(K - k0, block_size); -#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) - signed char* pp = outptr0 + k0 + g * 4; -#else - signed char* pp = outptr0 + k0; -#endif + signed char* pp = pp0; float absmax = 0.f; int kk = 0; @@ -1421,7 +1266,7 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales __m512 _absmax512 = _mm512_setzero_ps(); for (; kk + 15 < max_kk; kk += 16) { - __m512 _p = _mm512_loadu_ps(ptrA + k0 + kk); + __m512 _p = _mm512_loadu_ps(p0 + kk); if (input_scale_ptr) _p = _mm512_mul_ps(_p, _mm512_loadu_ps(input_scale_ptr + k0 + kk)); _absmax512 = _mm512_max_ps(_absmax512, abs512_ps(_p)); @@ -1431,7 +1276,7 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales __m256 _absmax256 = _mm256_setzero_ps(); for (; kk + 7 < max_kk; kk += 8) { - __m256 _p = _mm256_loadu_ps(ptrA + k0 + kk); + __m256 _p = _mm256_loadu_ps(p0 + kk); if (input_scale_ptr) _p = _mm256_mul_ps(_p, _mm256_loadu_ps(input_scale_ptr + k0 + kk)); _absmax256 = _mm256_max_ps(_absmax256, abs256_ps(_p)); @@ -1441,7 +1286,7 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales __m128 _absmax128 = _mm_setzero_ps(); for (; kk + 3 < max_kk; kk += 4) { - __m128 _p = _mm_loadu_ps(ptrA + k0 + kk); + __m128 _p = _mm_loadu_ps(p0 + kk); if (input_scale_ptr) _p = _mm_mul_ps(_p, _mm_loadu_ps(input_scale_ptr + k0 + kk)); _absmax128 = _mm_max_ps(_absmax128, abs_ps(_p)); @@ -1450,7 +1295,7 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales #endif // __SSE2__ for (; kk < max_kk; kk++) { - float v = ptrA[k0 + kk]; + float v = p0[kk]; if (input_scale_ptr) v *= input_scale_ptr[k0 + kk]; absmax = std::max(absmax, (float)fabsf(v)); @@ -1458,22 +1303,21 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales if (absmax == 0.f) { - descale_ptr0[g] = 0.f; + pd[0] = 0.f; #if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) memset(pp, 0, max_kk >= 4 ? max_kk + 4 : max_kk); + pp0 += max_kk + (max_kk >= 4 ? 4 : 0); #else memset(pp, 0, max_kk); + pp0 += max_kk; #endif + p0 += max_kk; + pd++; continue; } -#if __SSE2__ - const float scale = (float)(127.0 / (double)absmax); -#else - volatile double scale_fp64 = 127.0 / (double)absmax; - const float scale = (float)scale_fp64; -#endif - descale_ptr0[g] = absmax / 127.f; + const float scale = 127.f / absmax; + pd[0] = absmax / 127.f; #if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) int w_shift = 0; #endif @@ -1481,82 +1325,61 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales #if __SSE2__ #if __AVX__ #if __AVX512F__ - const __m512 _scale512 = _mm512_set1_ps(scale); + __m512 _scale512 = _mm512_set1_ps(scale); for (; kk + 15 < max_kk; kk += 16) { - __m512 _p = _mm512_loadu_ps(ptrA + k0 + kk); + __m512 _p = _mm512_loadu_ps(p0 + kk); if (input_scale_ptr) { _p = _mm512_mul_ps(_p, _mm512_loadu_ps(input_scale_ptr + k0 + kk)); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+x"(_p)); -#else - volatile __m512 _p_ordered = _p; - _p = _p_ordered; -#endif } - const __m128i _q = float2int8_avx512(_mm512_mul_ps(_p, _scale512)); + __m128i _q = float2int8_avx512(_mm512_mul_ps(_p, _scale512)); _mm_storeu_si128((__m128i*)pp, _q); pp += 16; #if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) - const __m256i _q16 = _mm256_cvtepi8_epi16(_q); - const __m256i _q32 = _mm256_madd_epi16(_q16, _mm256_set1_epi16(1)); + __m256i _q16 = _mm256_cvtepi8_epi16(_q); + __m256i _q32 = _mm256_madd_epi16(_q16, _mm256_set1_epi16(1)); w_shift += _mm_reduce_add_epi32(_mm256_castsi256_si128(_q32)); w_shift += _mm_reduce_add_epi32(_mm256_extracti128_si256(_q32, 1)); #endif } #endif // __AVX512F__ - const __m256 _scale256 = _mm256_set1_ps(scale); + __m256 _scale256 = _mm256_set1_ps(scale); for (; kk + 7 < max_kk; kk += 8) { - __m256 _p = _mm256_loadu_ps(ptrA + k0 + kk); + __m256 _p = _mm256_loadu_ps(p0 + kk); if (input_scale_ptr) { _p = _mm256_mul_ps(_p, _mm256_loadu_ps(input_scale_ptr + k0 + kk)); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+x"(_p)); -#else - volatile __m256 _p_ordered = _p; - _p = _p_ordered; -#endif } const int64_t q = float2int8_avx(_mm256_mul_ps(_p, _scale256)); - *(int64_t*)pp = q; + _mm_storeu_si64(pp, _mm_cvtsi64_si128(q)); pp += 8; #if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) #if defined(__x86_64__) || defined(_M_X64) - const __m128i _q8 = _mm_cvtsi64_si128(q); + __m128i _q8 = _mm_cvtsi64_si128(q); #else - const __m128i _q8 = _mm_loadl_epi64((const __m128i*)(pp - 8)); + __m128i _q8 = _mm_loadl_epi64((const __m128i*)(pp - 8)); #endif - const __m128i _q16 = _mm_unpacklo_epi8(_q8, _mm_cmpgt_epi8(_mm_setzero_si128(), _q8)); + __m128i _q16 = _mm_unpacklo_epi8(_q8, _mm_cmpgt_epi8(_mm_setzero_si128(), _q8)); w_shift += _mm_reduce_add_epi32(_mm_madd_epi16(_q16, _mm_set1_epi16(1))); #endif } #endif // __AVX__ - const __m128 _scale128 = _mm_set1_ps(scale); + __m128 _scale128 = _mm_set1_ps(scale); for (; kk + 3 < max_kk; kk += 4) { - __m128 _p = _mm_loadu_ps(ptrA + k0 + kk); + __m128 _p = _mm_loadu_ps(p0 + kk); if (input_scale_ptr) { _p = _mm_mul_ps(_p, _mm_loadu_ps(input_scale_ptr + k0 + kk)); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+x"(_p)); -#else - volatile __m128 _p_ordered = _p; - _p = _p_ordered; -#endif } const int32_t q = float2int8_sse(_mm_mul_ps(_p, _scale128)); - *(int32_t*)pp = q; + _mm_storeu_si32(pp, _mm_cvtsi32_si128(q)); pp += 4; #if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) - const __m128i _q8 = _mm_cvtsi32_si128(q); - const __m128i _q16 = _mm_unpacklo_epi8(_q8, _mm_cmpgt_epi8(_mm_setzero_si128(), _q8)); + __m128i _q8 = _mm_cvtsi32_si128(q); + __m128i _q16 = _mm_unpacklo_epi8(_q8, _mm_cmpgt_epi8(_mm_setzero_si128(), _q8)); w_shift += _mm_reduce_add_epi32(_mm_madd_epi16(_q16, _mm_set1_epi16(1))); #endif } @@ -1564,26 +1387,27 @@ static void quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& AT_descales #if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) if (max_kk >= 4) { - ((int*)pp)[0] = w_shift * 127; + _mm_storeu_si32(pp, _mm_cvtsi32_si128(w_shift * 127)); pp += 4; } #endif for (; kk < max_kk; kk++) { - float v = ptrA[k0 + kk]; + float v = p0[kk]; if (input_scale_ptr) { v *= input_scale_ptr[k0 + kk]; -#if NCNN_GNU_INLINE_ASM && __SSE2__ - asm volatile("" - : "+x"(v)); -#else - volatile float v_ordered = v; - v = v_ordered; -#endif } *pp++ = float2int8(v * scale); } + + p0 += max_kk; + pd++; +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + pp0 += max_kk + (max_kk >= 4 ? 4 : 0); +#else + pp0 += max_kk; +#endif } } } @@ -1631,8 +1455,9 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A #if __AVX512F__ for (; ii + 15 < max_ii; ii += 16) { - signed char* outptr0 = outptr + ii * out_hstep; - float* descale_ptr0 = descale_ptr + ii * block_count; + const float* p0 = (const float*)A + i + ii; + signed char* pp0 = outptr + ii * out_hstep; + float* pd = descale_ptr + ii * block_count; for (int g = 0; g < block_count; g++) { @@ -1642,66 +1467,42 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A __m512 _absmax = _mm512_setzero_ps(); for (int kk = 0; kk < max_kk; kk++) { - __m512 _p = _mm512_loadu_ps((const float*)A + (k0 + kk) * A_hstep + i + ii); + __m512 _p = _mm512_loadu_ps(p0 + kk * A_hstep); if (input_scale_ptr) _p = _mm512_mul_ps(_p, _mm512_set1_ps(input_scale_ptr[k0 + kk])); _absmax = _mm512_max_ps(_absmax, abs512_ps(_p)); } - const __m512 _descale = _mm512_div_ps(_absmax, _mm512_set1_ps(127.f)); - const __m256 _absmax0_fp32 = _mm512_castps512_ps256(_absmax); - const __m256 _absmax1_fp32 = _mm256_castpd_ps(_mm512_extractf64x4_pd(_mm512_castps_pd(_absmax), 1)); - const __m512d _absmax0_fp64 = _mm512_cvtps_pd(_absmax0_fp32); - const __m512d _absmax1_fp64 = _mm512_cvtps_pd(_absmax1_fp32); - const __mmask8 _nonzero0 = _mm512_cmp_pd_mask(_absmax0_fp64, _mm512_setzero_pd(), _CMP_NEQ_OQ); - const __mmask8 _nonzero1 = _mm512_cmp_pd_mask(_absmax1_fp64, _mm512_setzero_pd(), _CMP_NEQ_OQ); - const __m256 _scale0 = _mm512_cvtpd_ps(_mm512_maskz_div_pd(_nonzero0, _mm512_set1_pd(127.0), _absmax0_fp64)); - const __m256 _scale1 = _mm512_cvtpd_ps(_mm512_maskz_div_pd(_nonzero1, _mm512_set1_pd(127.0), _absmax1_fp64)); - const __m512 _scale = combine8x2_ps(_scale0, _scale1); - _mm512_storeu_ps(descale_ptr0 + g * 16, _descale); + __m512 _descale = _mm512_div_ps(_absmax, _mm512_set1_ps(127.f)); + __mmask16 _nonzero = _mm512_cmp_ps_mask(_absmax, _mm512_setzero_ps(), _CMP_NEQ_OQ); + __m512 _scale = _mm512_maskz_div_ps(_nonzero, _mm512_set1_ps(127.f), _absmax); + _mm512_storeu_ps(pd, _descale); #if __AVX512VNNI__ __m512i _w_shift = _mm512_setzero_si512(); #endif -#if __AVX512VNNI__ - signed char* pp = outptr0 + (k0 + g * 4) * 16; -#else - signed char* pp = outptr0 + k0 * 16; -#endif + signed char* pp = pp0; int kk = 0; #if __AVX512VNNI__ for (; kk + 3 < max_kk; kk += 4) { - __m512 _p0 = _mm512_loadu_ps((const float*)A + (k0 + kk) * A_hstep + i + ii); - __m512 _p1 = _mm512_loadu_ps((const float*)A + (k0 + kk + 1) * A_hstep + i + ii); - __m512 _p2 = _mm512_loadu_ps((const float*)A + (k0 + kk + 2) * A_hstep + i + ii); - __m512 _p3 = _mm512_loadu_ps((const float*)A + (k0 + kk + 3) * A_hstep + i + ii); + __m512 _p0 = _mm512_loadu_ps(p0 + kk * A_hstep); + __m512 _p1 = _mm512_loadu_ps(p0 + (kk + 1) * A_hstep); + __m512 _p2 = _mm512_loadu_ps(p0 + (kk + 2) * A_hstep); + __m512 _p3 = _mm512_loadu_ps(p0 + (kk + 3) * A_hstep); if (input_scale_ptr) { _p0 = _mm512_mul_ps(_p0, _mm512_set1_ps(input_scale_ptr[k0 + kk])); _p1 = _mm512_mul_ps(_p1, _mm512_set1_ps(input_scale_ptr[k0 + kk + 1])); _p2 = _mm512_mul_ps(_p2, _mm512_set1_ps(input_scale_ptr[k0 + kk + 2])); _p3 = _mm512_mul_ps(_p3, _mm512_set1_ps(input_scale_ptr[k0 + kk + 3])); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+x"(_p0), "+x"(_p1), "+x"(_p2), "+x"(_p3)); -#else - volatile __m512 _p0_ordered = _p0; - volatile __m512 _p1_ordered = _p1; - volatile __m512 _p2_ordered = _p2; - volatile __m512 _p3_ordered = _p3; - _p0 = _p0_ordered; - _p1 = _p1_ordered; - _p2 = _p2_ordered; - _p3 = _p3_ordered; -#endif } __m128i _q0 = float2int8_avx512(_mm512_mul_ps(_p0, _scale)); __m128i _q1 = float2int8_avx512(_mm512_mul_ps(_p1, _scale)); __m128i _q2 = float2int8_avx512(_mm512_mul_ps(_p2, _scale)); __m128i _q3 = float2int8_avx512(_mm512_mul_ps(_p3, _scale)); transpose16x4_epi8(_q0, _q1, _q2, _q3); - const __m512i _q = combine4x4_epi32(_q0, _q1, _q2, _q3); + __m512i _q = combine4x4_epi32(_q0, _q1, _q2, _q3); _mm512_storeu_si512((__m512i*)pp, _q); _w_shift = _mm512_dpbusd_epi32(_w_shift, _mm512_set1_epi8(127), _q); pp += 64; @@ -1716,51 +1517,44 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A #endif // __AVX512VNNI__ for (; kk + 1 < max_kk; kk += 2) { - __m512 _p0 = _mm512_loadu_ps((const float*)A + (k0 + kk) * A_hstep + i + ii); - __m512 _p1 = _mm512_loadu_ps((const float*)A + (k0 + kk + 1) * A_hstep + i + ii); + __m512 _p0 = _mm512_loadu_ps(p0 + kk * A_hstep); + __m512 _p1 = _mm512_loadu_ps(p0 + (kk + 1) * A_hstep); if (input_scale_ptr) { _p0 = _mm512_mul_ps(_p0, _mm512_set1_ps(input_scale_ptr[k0 + kk])); _p1 = _mm512_mul_ps(_p1, _mm512_set1_ps(input_scale_ptr[k0 + kk + 1])); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+x"(_p0), "+x"(_p1)); -#else - volatile __m512 _p0_ordered = _p0; - volatile __m512 _p1_ordered = _p1; - _p0 = _p0_ordered; - _p1 = _p1_ordered; -#endif } - const __m128i _q0 = float2int8_avx512(_mm512_mul_ps(_p0, _scale)); - const __m128i _q1 = float2int8_avx512(_mm512_mul_ps(_p1, _scale)); + __m128i _q0 = float2int8_avx512(_mm512_mul_ps(_p0, _scale)); + __m128i _q1 = float2int8_avx512(_mm512_mul_ps(_p1, _scale)); _mm_storeu_si128((__m128i*)pp, _mm_unpacklo_epi8(_q0, _q1)); _mm_storeu_si128((__m128i*)(pp + 16), _mm_unpackhi_epi8(_q0, _q1)); pp += 32; } if (kk < max_kk) { - __m512 _p = _mm512_loadu_ps((const float*)A + (k0 + kk) * A_hstep + i + ii); + __m512 _p = _mm512_loadu_ps(p0 + kk * A_hstep); if (input_scale_ptr) { _p = _mm512_mul_ps(_p, _mm512_set1_ps(input_scale_ptr[k0 + kk])); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+x"(_p)); -#else - volatile __m512 _p_ordered = _p; - _p = _p_ordered; -#endif } _mm_storeu_si128((__m128i*)pp, float2int8_avx512(_mm512_mul_ps(_p, _scale))); } + + p0 += max_kk * A_hstep; + pd += 16; +#if __AVX512VNNI__ + pp0 += (max_kk + (max_kk >= 4 ? 4 : 0)) * 16; +#else + pp0 += max_kk * 16; +#endif } } #endif // __AVX512F__ for (; ii + 7 < max_ii; ii += 8) { - signed char* outptr0 = outptr + ii * out_hstep; - float* descale_ptr0 = descale_ptr + ii * block_count; + const float* p0 = (const float*)A + i + ii; + signed char* pp0 = outptr + ii * out_hstep; + float* pd = descale_ptr + ii * block_count; for (int g = 0; g < block_count; g++) { @@ -1770,54 +1564,35 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A __m256 _absmax = _mm256_setzero_ps(); for (int kk = 0; kk < max_kk; kk++) { - __m256 _p = _mm256_loadu_ps((const float*)A + (k0 + kk) * A_hstep + i + ii); + __m256 _p = _mm256_loadu_ps(p0 + kk * A_hstep); if (input_scale_ptr) _p = _mm256_mul_ps(_p, _mm256_set1_ps(input_scale_ptr[k0 + kk])); _absmax = _mm256_max_ps(_absmax, abs256_ps(_p)); } - const __m256 _descale = _mm256_div_ps(_absmax, _mm256_set1_ps(127.f)); - const __m256 _nonzero = _mm256_cmp_ps(_absmax, _mm256_setzero_ps(), _CMP_NEQ_OQ); - const __m256 _absmax_nonzero = _mm256_blendv_ps(_mm256_set1_ps(1.f), _absmax, _nonzero); - const __m256d _absmax0_fp64 = _mm256_cvtps_pd(_mm256_castps256_ps128(_absmax_nonzero)); - const __m256d _absmax1_fp64 = _mm256_cvtps_pd(_mm256_extractf128_ps(_absmax_nonzero, 1)); - const __m128 _scale0 = _mm256_cvtpd_ps(_mm256_div_pd(_mm256_set1_pd(127.0), _absmax0_fp64)); - const __m128 _scale1 = _mm256_cvtpd_ps(_mm256_div_pd(_mm256_set1_pd(127.0), _absmax1_fp64)); - const __m256 _scale = _mm256_and_ps(combine4x2_ps(_scale0, _scale1), _nonzero); - _mm256_storeu_ps(descale_ptr0 + g * 8, _descale); + __m256 _descale = _mm256_div_ps(_absmax, _mm256_set1_ps(127.f)); + __m256 _nonzero = _mm256_cmp_ps(_absmax, _mm256_setzero_ps(), _CMP_NEQ_OQ); + __m256 _absmax_nonzero = _mm256_blendv_ps(_mm256_set1_ps(1.f), _absmax, _nonzero); + __m256 _scale = _mm256_and_ps(_mm256_div_ps(_mm256_set1_ps(127.f), _absmax_nonzero), _nonzero); + _mm256_storeu_ps(pd, _descale); + signed char* pp = pp0; #if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) __m256i _w_shift = _mm256_setzero_si256(); - signed char* pp = outptr0 + (k0 + g * 4) * 8; -#else - signed char* pp = outptr0 + k0 * 8; #endif int kk = 0; for (; kk + 3 < max_kk; kk += 4) { - __m256 _p0 = _mm256_loadu_ps((const float*)A + (k0 + kk) * A_hstep + i + ii); - __m256 _p1 = _mm256_loadu_ps((const float*)A + (k0 + kk + 1) * A_hstep + i + ii); - __m256 _p2 = _mm256_loadu_ps((const float*)A + (k0 + kk + 2) * A_hstep + i + ii); - __m256 _p3 = _mm256_loadu_ps((const float*)A + (k0 + kk + 3) * A_hstep + i + ii); + __m256 _p0 = _mm256_loadu_ps(p0 + kk * A_hstep); + __m256 _p1 = _mm256_loadu_ps(p0 + (kk + 1) * A_hstep); + __m256 _p2 = _mm256_loadu_ps(p0 + (kk + 2) * A_hstep); + __m256 _p3 = _mm256_loadu_ps(p0 + (kk + 3) * A_hstep); if (input_scale_ptr) { _p0 = _mm256_mul_ps(_p0, _mm256_set1_ps(input_scale_ptr[k0 + kk])); _p1 = _mm256_mul_ps(_p1, _mm256_set1_ps(input_scale_ptr[k0 + kk + 1])); _p2 = _mm256_mul_ps(_p2, _mm256_set1_ps(input_scale_ptr[k0 + kk + 2])); _p3 = _mm256_mul_ps(_p3, _mm256_set1_ps(input_scale_ptr[k0 + kk + 3])); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+x"(_p0), "+x"(_p1), "+x"(_p2), "+x"(_p3)); -#else - volatile __m256 _p0_ordered = _p0; - volatile __m256 _p1_ordered = _p1; - volatile __m256 _p2_ordered = _p2; - volatile __m256 _p3_ordered = _p3; - _p0 = _p0_ordered; - _p1 = _p1_ordered; - _p2 = _p2_ordered; - _p3 = _p3_ordered; -#endif } _p0 = _mm256_mul_ps(_p0, _scale); _p1 = _mm256_mul_ps(_p1, _scale); @@ -1831,7 +1606,7 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A #if __AVX512VNNI__ || __AVXVNNI__ _q0 = _mm_unpacklo_epi16(_q01, _q23); _q1 = _mm_unpackhi_epi16(_q01, _q23); - const __m256i _q = combine4x2_epi32(_q0, _q1); + __m256i _q = combine4x2_epi32(_q0, _q1); _mm256_storeu_si256((__m256i*)pp, _q); #if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) _w_shift = _mm256_comp_dpbusd_epi32(_w_shift, _mm256_set1_epi8(127), _q); @@ -1851,54 +1626,47 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A #endif for (; kk + 1 < max_kk; kk += 2) { - __m256 _p0 = _mm256_loadu_ps((const float*)A + (k0 + kk) * A_hstep + i + ii); - __m256 _p1 = _mm256_loadu_ps((const float*)A + (k0 + kk + 1) * A_hstep + i + ii); + __m256 _p0 = _mm256_loadu_ps(p0 + kk * A_hstep); + __m256 _p1 = _mm256_loadu_ps(p0 + (kk + 1) * A_hstep); if (input_scale_ptr) { _p0 = _mm256_mul_ps(_p0, _mm256_set1_ps(input_scale_ptr[k0 + kk])); _p1 = _mm256_mul_ps(_p1, _mm256_set1_ps(input_scale_ptr[k0 + kk + 1])); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+x"(_p0), "+x"(_p1)); -#else - volatile __m256 _p0_ordered = _p0; - volatile __m256 _p1_ordered = _p1; - _p0 = _p0_ordered; - _p1 = _p1_ordered; -#endif } _p0 = _mm256_mul_ps(_p0, _scale); _p1 = _mm256_mul_ps(_p1, _scale); __m128i _q = float2int8_avx(_p0, _p1); - const __m128i _si = _mm_setr_epi8(0, 8, 1, 9, 2, 10, 3, 11, 4, 12, 5, 13, 6, 14, 7, 15); + __m128i _si = _mm_setr_epi8(0, 8, 1, 9, 2, 10, 3, 11, 4, 12, 5, 13, 6, 14, 7, 15); _q = _mm_shuffle_epi8(_q, _si); _mm_storeu_si128((__m128i*)pp, _q); pp += 16; } for (; kk < max_kk; kk++) { - __m256 _p = _mm256_loadu_ps((const float*)A + (k0 + kk) * A_hstep + i + ii); + __m256 _p = _mm256_loadu_ps(p0 + kk * A_hstep); if (input_scale_ptr) { _p = _mm256_mul_ps(_p, _mm256_set1_ps(input_scale_ptr[k0 + kk])); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+x"(_p)); -#else - volatile __m256 _p_ordered = _p; - _p = _p_ordered; -#endif } - *(int64_t*)pp = float2int8_avx(_mm256_mul_ps(_p, _scale)); + _mm_storeu_si64(pp, _mm_cvtsi64_si128(float2int8_avx(_mm256_mul_ps(_p, _scale)))); pp += 8; } + + p0 += max_kk * A_hstep; + pd += 8; +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + pp0 += (max_kk + (max_kk >= 4 ? 4 : 0)) * 8; +#else + pp0 += max_kk * 8; +#endif } } #endif // __AVX2__ for (; ii + 3 < max_ii; ii += 4) { - signed char* outptr0 = outptr + ii * out_hstep; - float* descale_ptr0 = descale_ptr + ii * block_count; + const float* p0 = (const float*)A + i + ii; + signed char* pp0 = outptr + ii * out_hstep; + float* pd = descale_ptr + ii * block_count; for (int g = 0; g < block_count; g++) { @@ -1908,63 +1676,44 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A __m128 _absmax = _mm_setzero_ps(); for (int kk = 0; kk < max_kk; kk++) { - __m128 _p = _mm_loadu_ps((const float*)A + (k0 + kk) * A_hstep + i + ii); + __m128 _p = _mm_loadu_ps(p0 + kk * A_hstep); if (input_scale_ptr) _p = _mm_mul_ps(_p, _mm_set1_ps(input_scale_ptr[k0 + kk])); _absmax = _mm_max_ps(_absmax, abs_ps(_p)); } - const __m128 _descale = _mm_div_ps(_absmax, _mm_set1_ps(127.f)); - const __m128 _nonzero = _mm_cmpneq_ps(_absmax, _mm_setzero_ps()); - const __m128 _absmax_nonzero = _mm_or_ps(_mm_and_ps(_absmax, _nonzero), _mm_andnot_ps(_nonzero, _mm_set1_ps(1.f))); - const __m128d _absmax01_fp64 = _mm_cvtps_pd(_absmax_nonzero); - const __m128d _absmax23_fp64 = _mm_cvtps_pd(_mm_movehl_ps(_absmax_nonzero, _absmax_nonzero)); - const __m128 _scale01 = _mm_cvtpd_ps(_mm_div_pd(_mm_set1_pd(127.0), _absmax01_fp64)); - const __m128 _scale23 = _mm_cvtpd_ps(_mm_div_pd(_mm_set1_pd(127.0), _absmax23_fp64)); - const __m128 _scale = _mm_and_ps(_mm_movelh_ps(_scale01, _scale23), _nonzero); - _mm_storeu_ps(descale_ptr0 + g * 4, _descale); + __m128 _descale = _mm_div_ps(_absmax, _mm_set1_ps(127.f)); + __m128 _nonzero = _mm_cmpneq_ps(_absmax, _mm_setzero_ps()); + __m128 _absmax_nonzero = _mm_or_ps(_mm_and_ps(_absmax, _nonzero), _mm_andnot_ps(_nonzero, _mm_set1_ps(1.f))); + __m128 _scale = _mm_and_ps(_mm_div_ps(_mm_set1_ps(127.f), _absmax_nonzero), _nonzero); + _mm_storeu_ps(pd, _descale); + signed char* pp = pp0; #if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) __m128i _w_shift = _mm_setzero_si128(); - signed char* pp = outptr0 + (k0 + g * 4) * 4; -#else - signed char* pp = outptr0 + k0 * 4; #endif int kk = 0; for (; kk + 3 < max_kk; kk += 4) { - __m128 _p0 = _mm_loadu_ps((const float*)A + (k0 + kk) * A_hstep + i + ii); - __m128 _p1 = _mm_loadu_ps((const float*)A + (k0 + kk + 1) * A_hstep + i + ii); - __m128 _p2 = _mm_loadu_ps((const float*)A + (k0 + kk + 2) * A_hstep + i + ii); - __m128 _p3 = _mm_loadu_ps((const float*)A + (k0 + kk + 3) * A_hstep + i + ii); + __m128 _p0 = _mm_loadu_ps(p0 + kk * A_hstep); + __m128 _p1 = _mm_loadu_ps(p0 + (kk + 1) * A_hstep); + __m128 _p2 = _mm_loadu_ps(p0 + (kk + 2) * A_hstep); + __m128 _p3 = _mm_loadu_ps(p0 + (kk + 3) * A_hstep); if (input_scale_ptr) { _p0 = _mm_mul_ps(_p0, _mm_set1_ps(input_scale_ptr[k0 + kk])); _p1 = _mm_mul_ps(_p1, _mm_set1_ps(input_scale_ptr[k0 + kk + 1])); _p2 = _mm_mul_ps(_p2, _mm_set1_ps(input_scale_ptr[k0 + kk + 2])); _p3 = _mm_mul_ps(_p3, _mm_set1_ps(input_scale_ptr[k0 + kk + 3])); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+x"(_p0), "+x"(_p1), "+x"(_p2), "+x"(_p3)); -#else - volatile __m128 _p0_ordered = _p0; - volatile __m128 _p1_ordered = _p1; - volatile __m128 _p2_ordered = _p2; - volatile __m128 _p3_ordered = _p3; - _p0 = _p0_ordered; - _p1 = _p1_ordered; - _p2 = _p2_ordered; - _p3 = _p3_ordered; -#endif - } - const __m128i _q0 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p0, _scale))); - const __m128i _q1 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p1, _scale))); - const __m128i _q2 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p2, _scale))); - const __m128i _q3 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p3, _scale))); - const __m128i _q01 = _mm_unpacklo_epi8(_q0, _q1); - const __m128i _q23 = _mm_unpacklo_epi8(_q2, _q3); + } + __m128i _q0 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p0, _scale))); + __m128i _q1 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p1, _scale))); + __m128i _q2 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p2, _scale))); + __m128i _q3 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p3, _scale))); + __m128i _q01 = _mm_unpacklo_epi8(_q0, _q1); + __m128i _q23 = _mm_unpacklo_epi8(_q2, _q3); #if __AVX512VNNI__ || __AVXVNNI__ - const __m128i _q = _mm_unpacklo_epi16(_q01, _q23); + __m128i _q = _mm_unpacklo_epi16(_q01, _q23); _mm_storeu_si128((__m128i*)pp, _q); #if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) _w_shift = _mm_comp_dpbusd_epi32(_w_shift, _mm_set1_epi8(127), _q); @@ -1983,52 +1732,45 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A #endif for (; kk + 1 < max_kk; kk += 2) { - __m128 _p0 = _mm_loadu_ps((const float*)A + (k0 + kk) * A_hstep + i + ii); - __m128 _p1 = _mm_loadu_ps((const float*)A + (k0 + kk + 1) * A_hstep + i + ii); + __m128 _p0 = _mm_loadu_ps(p0 + kk * A_hstep); + __m128 _p1 = _mm_loadu_ps(p0 + (kk + 1) * A_hstep); if (input_scale_ptr) { _p0 = _mm_mul_ps(_p0, _mm_set1_ps(input_scale_ptr[k0 + kk])); _p1 = _mm_mul_ps(_p1, _mm_set1_ps(input_scale_ptr[k0 + kk + 1])); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+x"(_p0), "+x"(_p1)); -#else - volatile __m128 _p0_ordered = _p0; - volatile __m128 _p1_ordered = _p1; - _p0 = _p0_ordered; - _p1 = _p1_ordered; -#endif } - const __m128i _q0 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p0, _scale))); - const __m128i _q1 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p1, _scale))); + __m128i _q0 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p0, _scale))); + __m128i _q1 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p1, _scale))); _mm_storel_epi64((__m128i*)pp, _mm_unpacklo_epi8(_q0, _q1)); pp += 8; } for (; kk < max_kk; kk++) { - __m128 _p = _mm_loadu_ps((const float*)A + (k0 + kk) * A_hstep + i + ii); + __m128 _p = _mm_loadu_ps(p0 + kk * A_hstep); if (input_scale_ptr) { _p = _mm_mul_ps(_p, _mm_set1_ps(input_scale_ptr[k0 + kk])); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+x"(_p)); -#else - volatile __m128 _p_ordered = _p; - _p = _p_ordered; -#endif } - *(int*)pp = float2int8_sse(_mm_mul_ps(_p, _scale)); + _mm_storeu_si32(pp, _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p, _scale)))); pp += 4; } + + p0 += max_kk * A_hstep; + pd += 4; +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + pp0 += (max_kk + (max_kk >= 4 ? 4 : 0)) * 4; +#else + pp0 += max_kk * 4; +#endif } } #endif // __SSE2__ for (; ii + 1 < max_ii; ii += 2) { #if __SSE2__ - signed char* outptr0 = outptr + ii * out_hstep; - float* descale_ptr0 = descale_ptr + ii * block_count; + const float* p0 = (const float*)A + i + ii; + signed char* pp0 = outptr + ii * out_hstep; + float* pd = descale_ptr + ii * block_count; for (int g = 0; g < block_count; g++) { @@ -2038,59 +1780,44 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A __m128 _absmax = _mm_setzero_ps(); for (int kk = 0; kk < max_kk; kk++) { - __m128 _p = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)((const float*)A + (k0 + kk) * A_hstep + i + ii)); + __m128 _p = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(p0 + kk * A_hstep)); if (input_scale_ptr) _p = _mm_mul_ps(_p, _mm_set1_ps(input_scale_ptr[k0 + kk])); _absmax = _mm_max_ps(_absmax, abs_ps(_p)); } - const __m128 _descale = _mm_div_ps(_absmax, _mm_set1_ps(127.f)); - const __m128 _nonzero = _mm_cmpneq_ps(_absmax, _mm_setzero_ps()); - const __m128 _absmax_nonzero = _mm_or_ps(_mm_and_ps(_absmax, _nonzero), _mm_andnot_ps(_nonzero, _mm_set1_ps(1.f))); - const __m128 _scale = _mm_and_ps(_mm_cvtpd_ps(_mm_div_pd(_mm_set1_pd(127.0), _mm_cvtps_pd(_absmax_nonzero))), _nonzero); - _mm_storel_pi((__m64*)(descale_ptr0 + g * 2), _descale); + __m128 _descale = _mm_div_ps(_absmax, _mm_set1_ps(127.f)); + __m128 _nonzero = _mm_cmpneq_ps(_absmax, _mm_setzero_ps()); + __m128 _absmax_nonzero = _mm_or_ps(_mm_and_ps(_absmax, _nonzero), _mm_andnot_ps(_nonzero, _mm_set1_ps(1.f))); + __m128 _scale = _mm_and_ps(_mm_div_ps(_mm_set1_ps(127.f), _absmax_nonzero), _nonzero); + _mm_storel_pi((__m64*)pd, _descale); + signed char* pp = pp0; #if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) __m128i _w_shift = _mm_setzero_si128(); - signed char* pp = outptr0 + (k0 + g * 4) * 2; -#else - signed char* pp = outptr0 + k0 * 2; #endif int kk = 0; for (; kk + 3 < max_kk; kk += 4) { - __m128 _p0 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)((const float*)A + (k0 + kk) * A_hstep + i + ii)); - __m128 _p1 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)((const float*)A + (k0 + kk + 1) * A_hstep + i + ii)); - __m128 _p2 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)((const float*)A + (k0 + kk + 2) * A_hstep + i + ii)); - __m128 _p3 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)((const float*)A + (k0 + kk + 3) * A_hstep + i + ii)); + __m128 _p0 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(p0 + kk * A_hstep)); + __m128 _p1 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(p0 + (kk + 1) * A_hstep)); + __m128 _p2 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(p0 + (kk + 2) * A_hstep)); + __m128 _p3 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(p0 + (kk + 3) * A_hstep)); if (input_scale_ptr) { _p0 = _mm_mul_ps(_p0, _mm_set1_ps(input_scale_ptr[k0 + kk])); _p1 = _mm_mul_ps(_p1, _mm_set1_ps(input_scale_ptr[k0 + kk + 1])); _p2 = _mm_mul_ps(_p2, _mm_set1_ps(input_scale_ptr[k0 + kk + 2])); _p3 = _mm_mul_ps(_p3, _mm_set1_ps(input_scale_ptr[k0 + kk + 3])); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+x"(_p0), "+x"(_p1), "+x"(_p2), "+x"(_p3)); -#else - volatile __m128 _p0_ordered = _p0; - volatile __m128 _p1_ordered = _p1; - volatile __m128 _p2_ordered = _p2; - volatile __m128 _p3_ordered = _p3; - _p0 = _p0_ordered; - _p1 = _p1_ordered; - _p2 = _p2_ordered; - _p3 = _p3_ordered; -#endif - } - const __m128i _q0 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p0, _scale))); - const __m128i _q1 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p1, _scale))); - const __m128i _q2 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p2, _scale))); - const __m128i _q3 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p3, _scale))); - const __m128i _q01 = _mm_unpacklo_epi8(_q0, _q1); - const __m128i _q23 = _mm_unpacklo_epi8(_q2, _q3); + } + __m128i _q0 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p0, _scale))); + __m128i _q1 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p1, _scale))); + __m128i _q2 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p2, _scale))); + __m128i _q3 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p3, _scale))); + __m128i _q01 = _mm_unpacklo_epi8(_q0, _q1); + __m128i _q23 = _mm_unpacklo_epi8(_q2, _q3); #if __AVX512VNNI__ || __AVXVNNI__ - const __m128i _q = _mm_unpacklo_epi16(_q01, _q23); + __m128i _q = _mm_unpacklo_epi16(_q01, _q23); _mm_storel_epi64((__m128i*)pp, _q); #if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) _w_shift = _mm_comp_dpbusd_epi32(_w_shift, _mm_set1_epi8(127), _q); @@ -2109,48 +1836,41 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A #endif for (; kk + 1 < max_kk; kk += 2) { - __m128 _p0 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)((const float*)A + (k0 + kk) * A_hstep + i + ii)); - __m128 _p1 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)((const float*)A + (k0 + kk + 1) * A_hstep + i + ii)); + __m128 _p0 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(p0 + kk * A_hstep)); + __m128 _p1 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(p0 + (kk + 1) * A_hstep)); if (input_scale_ptr) { _p0 = _mm_mul_ps(_p0, _mm_set1_ps(input_scale_ptr[k0 + kk])); _p1 = _mm_mul_ps(_p1, _mm_set1_ps(input_scale_ptr[k0 + kk + 1])); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+x"(_p0), "+x"(_p1)); -#else - volatile __m128 _p0_ordered = _p0; - volatile __m128 _p1_ordered = _p1; - _p0 = _p0_ordered; - _p1 = _p1_ordered; -#endif } - const __m128i _q0 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p0, _scale))); - const __m128i _q1 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p1, _scale))); - *(int*)pp = _mm_cvtsi128_si32(_mm_unpacklo_epi8(_q0, _q1)); + __m128i _q0 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p0, _scale))); + __m128i _q1 = _mm_cvtsi32_si128(float2int8_sse(_mm_mul_ps(_p1, _scale))); + _mm_storeu_si32(pp, _mm_unpacklo_epi8(_q0, _q1)); pp += 4; } for (; kk < max_kk; kk++) { - __m128 _p = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)((const float*)A + (k0 + kk) * A_hstep + i + ii)); + __m128 _p = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(p0 + kk * A_hstep)); if (input_scale_ptr) { _p = _mm_mul_ps(_p, _mm_set1_ps(input_scale_ptr[k0 + kk])); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+x"(_p)); -#else - volatile __m128 _p_ordered = _p; - _p = _p_ordered; -#endif } - *(unsigned short*)pp = (unsigned short)float2int8_sse(_mm_mul_ps(_p, _scale)); + _mm_storeu_si16(pp, _mm_cvtsi32_si128((unsigned short)float2int8_sse(_mm_mul_ps(_p, _scale)))); pp += 2; } + + p0 += max_kk * A_hstep; + pd += 2; +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + pp0 += (max_kk + (max_kk >= 4 ? 4 : 0)) * 2; +#else + pp0 += max_kk * 2; +#endif } #else - signed char* outptr0 = outptr + ii * out_hstep; - float* descale_ptr0 = descale_ptr + ii * block_count; + const float* p0 = (const float*)A + i + ii; + signed char* pp0 = outptr + ii * out_hstep; + float* pd = descale_ptr + ii * block_count; for (int g = 0; g < block_count; g++) { @@ -2160,7 +1880,7 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A float absmax1 = 0.f; for (int kk = 0; kk < max_kk; kk++) { - const float* ptrA = (const float*)A + (k0 + kk) * A_hstep + i + ii; + const float* ptrA = p0 + kk * A_hstep; float v0 = ptrA[0]; float v1 = ptrA[1]; if (input_scale_ptr) @@ -2177,21 +1897,19 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A float scale1 = 0.f; if (absmax0 != 0.f) { - volatile double scale_fp64 = 127.0 / (double)absmax0; - scale0 = (float)scale_fp64; + scale0 = 127.f / absmax0; } if (absmax1 != 0.f) { - volatile double scale_fp64 = 127.0 / (double)absmax1; - scale1 = (float)scale_fp64; + scale1 = 127.f / absmax1; } - descale_ptr0[g * 2] = absmax0 / 127.f; - descale_ptr0[g * 2 + 1] = absmax1 / 127.f; + pd[0] = absmax0 / 127.f; + pd[1] = absmax1 / 127.f; - signed char* pp = outptr0 + k0 * 2; + signed char* pp = pp0; for (int kk = 0; kk < max_kk; kk++) { - const float* ptrA = (const float*)A + (k0 + kk) * A_hstep + i + ii; + const float* ptrA = p0 + kk * A_hstep; float v0 = ptrA[0]; float v1 = ptrA[1]; if (input_scale_ptr) @@ -2199,15 +1917,15 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A const float s = input_scale_ptr[k0 + kk]; v0 *= s; v1 *= s; - volatile float v0_ordered = v0; - volatile float v1_ordered = v1; - v0 = v0_ordered; - v1 = v1_ordered; } pp[0] = float2int8(v0 * scale0); pp[1] = float2int8(v1 * scale1); pp += 2; } + + p0 += max_kk * A_hstep; + pp0 += max_kk * 2; + pd += 2; } #endif // __SSE2__ } @@ -2225,18 +1943,15 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A for (; ii < max_ii; ii++) { - signed char* outptr0 = outptr + ii * out_hstep; - float* descale_ptr0 = descale_ptr + ii * block_count; + const float* p0 = (const float*)A + i + ii; + signed char* pp0 = outptr + ii * out_hstep; + float* pd = descale_ptr + ii * block_count; for (int g = 0; g < block_count; g++) { const int k0 = g * block_size; const int max_kk = std::min(K - k0, block_size); -#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) - signed char* pp = outptr0 + k0 + g * 4; -#else - signed char* pp = outptr0 + k0; -#endif + signed char* pp = pp0; float absmax = 0.f; int kk = 0; @@ -2246,7 +1961,7 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A __m512 _absmax512 = _mm512_setzero_ps(); for (; kk + 15 < max_kk; kk += 16) { - const float* ptrA = (const float*)A + (k0 + kk) * A_hstep + i + ii; + const float* ptrA = p0 + kk * A_hstep; __m512 _p = _mm512_i32gather_ps(_vindex512, ptrA, sizeof(float)); if (input_scale_ptr) _p = _mm512_mul_ps(_p, _mm512_loadu_ps(input_scale_ptr + k0 + kk)); @@ -2257,7 +1972,7 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A __m256 _absmax256 = _mm256_setzero_ps(); for (; kk + 7 < max_kk; kk += 8) { - const float* ptrA = (const float*)A + (k0 + kk) * A_hstep + i + ii; + const float* ptrA = p0 + kk * A_hstep; __m256 _p = _mm256_i32gather_ps(ptrA, _vindex256, sizeof(float)); if (input_scale_ptr) _p = _mm256_mul_ps(_p, _mm256_loadu_ps(input_scale_ptr + k0 + kk)); @@ -2268,7 +1983,7 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A __m128 _absmax128 = _mm_setzero_ps(); for (; kk + 3 < max_kk; kk += 4) { - const float* ptrA = (const float*)A + (k0 + kk) * A_hstep + i + ii; + const float* ptrA = p0 + kk * A_hstep; __m128 _p = _mm_setr_ps(ptrA[0], ptrA[A_hstep], ptrA[A_hstep * 2], ptrA[A_hstep * 3]); if (input_scale_ptr) _p = _mm_mul_ps(_p, _mm_loadu_ps(input_scale_ptr + k0 + kk)); @@ -2278,31 +1993,29 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A #endif // __SSE2__ for (; kk < max_kk; kk++) { - const int k = k0 + kk; - float v = ((const float*)A)[k * A_hstep + i + ii]; + float v = p0[kk * A_hstep]; if (input_scale_ptr) - v *= input_scale_ptr[k]; + v *= input_scale_ptr[k0 + kk]; absmax = std::max(absmax, (float)fabsf(v)); } if (absmax == 0.f) { - descale_ptr0[g] = 0.f; + pd[0] = 0.f; #if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) memset(pp, 0, max_kk >= 4 ? max_kk + 4 : max_kk); + pp0 += max_kk + (max_kk >= 4 ? 4 : 0); #else memset(pp, 0, max_kk); + pp0 += max_kk; #endif + p0 += max_kk * A_hstep; + pd++; continue; } -#if __SSE2__ - const float scale = (float)(127.0 / (double)absmax); -#else - volatile double scale_fp64 = 127.0 / (double)absmax; - const float scale = (float)scale_fp64; -#endif - descale_ptr0[g] = absmax / 127.f; + const float scale = 127.f / absmax; + pd[0] = absmax / 127.f; #if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) int w_shift = 0; #endif @@ -2310,85 +2023,64 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A #if __SSE2__ #if __AVX2__ #if __AVX512F__ - const __m512 _scale512 = _mm512_set1_ps(scale); + __m512 _scale512 = _mm512_set1_ps(scale); for (; kk + 15 < max_kk; kk += 16) { - const float* ptrA = (const float*)A + (k0 + kk) * A_hstep + i + ii; + const float* ptrA = p0 + kk * A_hstep; __m512 _p = _mm512_i32gather_ps(_vindex512, ptrA, sizeof(float)); if (input_scale_ptr) { _p = _mm512_mul_ps(_p, _mm512_loadu_ps(input_scale_ptr + k0 + kk)); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+x"(_p)); -#else - volatile __m512 _p_ordered = _p; - _p = _p_ordered; -#endif } - const __m128i _q = float2int8_avx512(_mm512_mul_ps(_p, _scale512)); + __m128i _q = float2int8_avx512(_mm512_mul_ps(_p, _scale512)); _mm_storeu_si128((__m128i*)pp, _q); pp += 16; #if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) - const __m256i _q16 = _mm256_cvtepi8_epi16(_q); - const __m256i _q32 = _mm256_madd_epi16(_q16, _mm256_set1_epi16(1)); + __m256i _q16 = _mm256_cvtepi8_epi16(_q); + __m256i _q32 = _mm256_madd_epi16(_q16, _mm256_set1_epi16(1)); w_shift += _mm_reduce_add_epi32(_mm256_castsi256_si128(_q32)); w_shift += _mm_reduce_add_epi32(_mm256_extracti128_si256(_q32, 1)); #endif } #endif // __AVX512F__ - const __m256 _scale256 = _mm256_set1_ps(scale); + __m256 _scale256 = _mm256_set1_ps(scale); for (; kk + 7 < max_kk; kk += 8) { - const float* ptrA = (const float*)A + (k0 + kk) * A_hstep + i + ii; + const float* ptrA = p0 + kk * A_hstep; __m256 _p = _mm256_i32gather_ps(ptrA, _vindex256, sizeof(float)); if (input_scale_ptr) { _p = _mm256_mul_ps(_p, _mm256_loadu_ps(input_scale_ptr + k0 + kk)); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+x"(_p)); -#else - volatile __m256 _p_ordered = _p; - _p = _p_ordered; -#endif } const int64_t q = float2int8_avx(_mm256_mul_ps(_p, _scale256)); - *(int64_t*)pp = q; + _mm_storeu_si64(pp, _mm_cvtsi64_si128(q)); pp += 8; #if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) #if defined(__x86_64__) || defined(_M_X64) - const __m128i _q8 = _mm_cvtsi64_si128(q); + __m128i _q8 = _mm_cvtsi64_si128(q); #else - const __m128i _q8 = _mm_loadl_epi64((const __m128i*)(pp - 8)); + __m128i _q8 = _mm_loadl_epi64((const __m128i*)(pp - 8)); #endif - const __m128i _q16 = _mm_unpacklo_epi8(_q8, _mm_cmpgt_epi8(_mm_setzero_si128(), _q8)); + __m128i _q16 = _mm_unpacklo_epi8(_q8, _mm_cmpgt_epi8(_mm_setzero_si128(), _q8)); w_shift += _mm_reduce_add_epi32(_mm_madd_epi16(_q16, _mm_set1_epi16(1))); #endif } #endif // __AVX2__ - const __m128 _scale128 = _mm_set1_ps(scale); + __m128 _scale128 = _mm_set1_ps(scale); for (; kk + 3 < max_kk; kk += 4) { - const float* ptrA = (const float*)A + (k0 + kk) * A_hstep + i + ii; + const float* ptrA = p0 + kk * A_hstep; __m128 _p = _mm_setr_ps(ptrA[0], ptrA[A_hstep], ptrA[A_hstep * 2], ptrA[A_hstep * 3]); if (input_scale_ptr) { _p = _mm_mul_ps(_p, _mm_loadu_ps(input_scale_ptr + k0 + kk)); -#if NCNN_GNU_INLINE_ASM - asm volatile("" - : "+x"(_p)); -#else - volatile __m128 _p_ordered = _p; - _p = _p_ordered; -#endif } const int32_t q = float2int8_sse(_mm_mul_ps(_p, _scale128)); - *(int32_t*)pp = q; + _mm_storeu_si32(pp, _mm_cvtsi32_si128(q)); pp += 4; #if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) - const __m128i _q8 = _mm_cvtsi32_si128(q); - const __m128i _q16 = _mm_unpacklo_epi8(_q8, _mm_cmpgt_epi8(_mm_setzero_si128(), _q8)); + __m128i _q8 = _mm_cvtsi32_si128(q); + __m128i _q16 = _mm_unpacklo_epi8(_q8, _mm_cmpgt_epi8(_mm_setzero_si128(), _q8)); w_shift += _mm_reduce_add_epi32(_mm_madd_epi16(_q16, _mm_set1_epi16(1))); #endif } @@ -2396,65 +2088,65 @@ static void transpose_quantize_A_tile_wq_int8(const Mat& A, Mat& AT_tile, Mat& A #if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) if (max_kk >= 4) { - ((int*)pp)[0] = w_shift * 127; + _mm_storeu_si32(pp, _mm_cvtsi32_si128(w_shift * 127)); pp += 4; } #endif for (; kk < max_kk; kk++) { - const int k = k0 + kk; - float v = ((const float*)A)[k * A_hstep + i + ii]; + float v = p0[kk * A_hstep]; if (input_scale_ptr) { - v *= input_scale_ptr[k]; -#if NCNN_GNU_INLINE_ASM && __SSE2__ - asm volatile("" - : "+x"(v)); -#else - volatile float v_ordered = v; - v = v_ordered; -#endif + v *= input_scale_ptr[k0 + kk]; } *pp++ = float2int8(v * scale); } + + p0 += max_kk * A_hstep; + pd++; +#if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) + pp0 += max_kk + (max_kk >= 4 ? 4 : 0); +#else + pp0 += max_kk; +#endif } } } -static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int block_size) +static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int k, int max_kk, int K, int block_size) { #if NCNN_RUNTIME_CPU && NCNN_AVX512VNNI && __AVX512F__ && !__AVX512VNNI__ if (ncnn::cpu_support_x86_avx512_vnni()) { - gemm_transB_packed_tile_wq_int8_avx512vnni(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); + gemm_transB_packed_tile_wq_int8_avx512vnni(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, k, max_kk, K, block_size); return; } #endif #if NCNN_RUNTIME_CPU && NCNN_AVXVNNIINT8 && __AVX__ && !__AVXVNNIINT8__ && !__AVX512VNNI__ if (ncnn::cpu_support_x86_avx_vnni_int8()) { - gemm_transB_packed_tile_wq_int8_avxvnniint8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); + gemm_transB_packed_tile_wq_int8_avxvnniint8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, k, max_kk, K, block_size); return; } #endif #if NCNN_RUNTIME_CPU && NCNN_AVXVNNI && __AVX__ && !__AVXVNNI__ && !__AVXVNNIINT8__ && !__AVX512VNNI__ if (ncnn::cpu_support_x86_avx_vnni()) { - gemm_transB_packed_tile_wq_int8_avxvnni(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); + gemm_transB_packed_tile_wq_int8_avxvnni(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, k, max_kk, K, block_size); return; } #endif #if NCNN_RUNTIME_CPU && NCNN_AVX2 && __AVX__ && !__AVX2__ && !__AVXVNNI__ && !__AVXVNNIINT8__ && !__AVX512VNNI__ if (ncnn::cpu_support_x86_avx2()) { - gemm_transB_packed_tile_wq_int8_avx2(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); + gemm_transB_packed_tile_wq_int8_avx2(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, k, max_kk, K, block_size); return; } #endif #if NCNN_RUNTIME_CPU && NCNN_XOP && __SSE2__ && !__XOP__ && !__AVX2__ && !__AVXVNNI__ && !__AVXVNNIINT8__ && !__AVX512VNNI__ if (ncnn::cpu_support_x86_xop()) { - gemm_transB_packed_tile_wq_int8_xop(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); + gemm_transB_packed_tile_wq_int8_xop(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, k, max_kk, K, block_size); return; } #endif @@ -2466,6 +2158,12 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de const signed char* pBT = BT_tile; const float* pBT_descales = BT_descales_tile; float* outptr = topT_tile; + const int tile_K = max_kk; + const int block_count = (K + block_size - 1) / block_size; + const int tile_block_count = (tile_K + block_size - 1) / block_size; + const int block_start = k / block_size; + const int remain_K = K - k - tile_K; + const int remain_blocks = block_count - block_start - tile_block_count; int ii = 0; #if __SSE2__ @@ -2480,18 +2178,20 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #if defined(__x86_64__) || defined(_M_X64) for (; jj + 7 < max_jj; jj += 8) { - __m512 _fsum0 = _mm512_setzero_ps(); - __m512 _fsum1 = _mm512_setzero_ps(); - __m512 _fsum2 = _mm512_setzero_ps(); - __m512 _fsum3 = _mm512_setzero_ps(); - __m512 _fsum4 = _mm512_setzero_ps(); - __m512 _fsum5 = _mm512_setzero_ps(); - __m512 _fsum6 = _mm512_setzero_ps(); - __m512 _fsum7 = _mm512_setzero_ps(); + pB += k * 8; + pB_descales += block_start * 8; + __m512 _fsum0 = k == 0 ? _mm512_setzero_ps() : _mm512_loadu_ps(outptr); + __m512 _fsum1 = k == 0 ? _mm512_setzero_ps() : _mm512_loadu_ps(outptr + 16); + __m512 _fsum2 = k == 0 ? _mm512_setzero_ps() : _mm512_loadu_ps(outptr + 32); + __m512 _fsum3 = k == 0 ? _mm512_setzero_ps() : _mm512_loadu_ps(outptr + 48); + __m512 _fsum4 = k == 0 ? _mm512_setzero_ps() : _mm512_loadu_ps(outptr + 64); + __m512 _fsum5 = k == 0 ? _mm512_setzero_ps() : _mm512_loadu_ps(outptr + 80); + __m512 _fsum6 = k == 0 ? _mm512_setzero_ps() : _mm512_loadu_ps(outptr + 96); + __m512 _fsum7 = k == 0 ? _mm512_setzero_ps() : _mm512_loadu_ps(outptr + 112); const signed char* pA = pAT; const float* pA_descales = pAT_descales; - for (int k = 0; k < K; k += block_size) + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) { __m512i _sum0 = _mm512_setzero_si512(); __m512i _sum1 = _mm512_setzero_si512(); @@ -2501,18 +2201,18 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de __m512i _sum5 = _mm512_setzero_si512(); __m512i _sum6 = _mm512_setzero_si512(); __m512i _sum7 = _mm512_setzero_si512(); - const int max_kk = std::min(K - k, block_size); + const int max_kk = std::min(tile_K - kk0, block_size); int kk = 0; #if __AVX512VNNI__ for (; kk + 3 < max_kk; kk += 4) { - const __m512i _pA0 = _mm512_loadu_si512((const __m512i*)pA); - const __m256i _pB = _mm256_loadu_si256((const __m256i*)pB); - const __m512i _pB0 = combine8x2_epi32(_pB, _pB); - const __m512i _pA1 = _mm512_alignr_epi8(_pA0, _pA0, 8); - const __m512i _pB1 = _mm512_alignr_epi8(_pB0, _pB0, 4); - const __m512i _pB2 = _mm512_permutex_epi64(_pB0, _MM_SHUFFLE(1, 0, 3, 2)); - const __m512i _pB3 = _mm512_alignr_epi8(_pB2, _pB2, 4); + __m512i _pA0 = _mm512_loadu_si512((const __m512i*)pA); + __m256i _pB = _mm256_loadu_si256((const __m256i*)pB); + __m512i _pB0 = combine8x2_epi32(_pB, _pB); + __m512i _pA1 = _mm512_alignr_epi8(_pA0, _pA0, 8); + __m512i _pB1 = _mm512_alignr_epi8(_pB0, _pB0, 4); + __m512i _pB2 = _mm512_permutex_epi64(_pB0, _MM_SHUFFLE(1, 0, 3, 2)); + __m512i _pB3 = _mm512_alignr_epi8(_pB2, _pB2, 4); _sum0 = _mm512_dpbusd_epi32(_sum0, _pB0, _pA0); _sum1 = _mm512_dpbusd_epi32(_sum1, _pB1, _pA0); _sum2 = _mm512_dpbusd_epi32(_sum2, _pB0, _pA1); @@ -2526,8 +2226,8 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } if (max_kk >= 4) { - const __m512i _shift0 = _mm512_loadu_si512((const __m512i*)pA); - const __m512i _shift1 = _mm512_alignr_epi8(_shift0, _shift0, 8); + __m512i _shift0 = _mm512_loadu_si512((const __m512i*)pA); + __m512i _shift1 = _mm512_alignr_epi8(_shift0, _shift0, 8); _sum0 = _mm512_sub_epi32(_sum0, _shift0); _sum1 = _mm512_sub_epi32(_sum1, _shift0); _sum2 = _mm512_sub_epi32(_sum2, _shift1); @@ -2541,15 +2241,15 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #endif // __AVX512VNNI__ for (; kk + 1 < max_kk; kk += 2) { - const __m256i _pA = _mm256_loadu_si256((const __m256i*)pA); - const __m512i _pA0 = _mm512_cvtepi8_epi16(_pA); - const __m512i _pA1 = _mm512_alignr_epi8(_pA0, _pA0, 8); - const __m128i _pB = _mm_loadu_si128((const __m128i*)pB); - const __m256i _pBB = _mm256_cvtepi8_epi16(_pB); - const __m512i _pB0 = combine8x2_epi32(_pBB, _pBB); - const __m512i _pB1 = _mm512_alignr_epi8(_pB0, _pB0, 4); - const __m512i _pB2 = _mm512_permutex_epi64(_pB0, _MM_SHUFFLE(1, 0, 3, 2)); - const __m512i _pB3 = _mm512_alignr_epi8(_pB2, _pB2, 4); + __m256i _pA = _mm256_loadu_si256((const __m256i*)pA); + __m512i _pA0 = _mm512_cvtepi8_epi16(_pA); + __m512i _pA1 = _mm512_alignr_epi8(_pA0, _pA0, 8); + __m128i _pB = _mm_loadu_si128((const __m128i*)pB); + __m256i _pBB = _mm256_cvtepi8_epi16(_pB); + __m512i _pB0 = combine8x2_epi32(_pBB, _pBB); + __m512i _pB1 = _mm512_alignr_epi8(_pB0, _pB0, 4); + __m512i _pB2 = _mm512_permutex_epi64(_pB0, _MM_SHUFFLE(1, 0, 3, 2)); + __m512i _pB3 = _mm512_alignr_epi8(_pB2, _pB2, 4); _sum0 = _mm512_comp_dpwssd_epi32(_sum0, _pA0, _pB0); _sum1 = _mm512_comp_dpwssd_epi32(_sum1, _pA0, _pB1); _sum2 = _mm512_comp_dpwssd_epi32(_sum2, _pA1, _pB0); @@ -2563,15 +2263,15 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } for (; kk < max_kk; kk++) { - const __m128i _pA = _mm_loadu_si128((const __m128i*)pA); - const __m256i _pA0 = _mm256_cvtepi8_epi16(_pA); - const __m256i _pA1 = _mm256_shuffle_epi32(_pA0, _MM_SHUFFLE(2, 3, 0, 1)); + __m128i _pA = _mm_loadu_si128((const __m128i*)pA); + __m256i _pA0 = _mm256_cvtepi8_epi16(_pA); + __m256i _pA1 = _mm256_shuffle_epi32(_pA0, _MM_SHUFFLE(2, 3, 0, 1)); __m128i _pB = _mm_loadl_epi64((const __m128i*)pB); _pB = _mm_cvtepi8_epi16(_pB); - const __m256i _pB0 = combine4x2_epi32(_pB, _pB); - const __m256i _pB1 = _mm256_shufflehi_epi16(_mm256_shufflelo_epi16(_pB0, _MM_SHUFFLE(0, 3, 2, 1)), _MM_SHUFFLE(0, 3, 2, 1)); - const __m256i _pB2 = _mm256_alignr_epi8(_pB0, _pB0, 8); - const __m256i _pB3 = _mm256_shufflehi_epi16(_mm256_shufflelo_epi16(_pB2, _MM_SHUFFLE(0, 3, 2, 1)), _MM_SHUFFLE(0, 3, 2, 1)); + __m256i _pB0 = combine4x2_epi32(_pB, _pB); + __m256i _pB1 = _mm256_shufflehi_epi16(_mm256_shufflelo_epi16(_pB0, _MM_SHUFFLE(0, 3, 2, 1)), _MM_SHUFFLE(0, 3, 2, 1)); + __m256i _pB2 = _mm256_alignr_epi8(_pB0, _pB0, 8); + __m256i _pB3 = _mm256_shufflehi_epi16(_mm256_shufflelo_epi16(_pB2, _MM_SHUFFLE(0, 3, 2, 1)), _MM_SHUFFLE(0, 3, 2, 1)); _sum0 = _mm512_add_epi32(_sum0, _mm512_cvtepi16_epi32(_mm256_mullo_epi16(_pA0, _pB0))); _sum1 = _mm512_add_epi32(_sum1, _mm512_cvtepi16_epi32(_mm256_mullo_epi16(_pA0, _pB1))); _sum2 = _mm512_add_epi32(_sum2, _mm512_cvtepi16_epi32(_mm256_mullo_epi16(_pA1, _pB0))); @@ -2584,13 +2284,13 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pA += 16; } - const __m512 _A0 = _mm512_loadu_ps(pA_descales); - const __m512 _A1 = _mm512_castsi512_ps(_mm512_alignr_epi8(_mm512_castps_si512(_A0), _mm512_castps_si512(_A0), 8)); - const __m256 _b = _mm256_loadu_ps(pB_descales); - const __m512 _B0 = combine8x2_ps(_b, _b); - const __m512 _B1 = _mm512_castsi512_ps(_mm512_alignr_epi8(_mm512_castps_si512(_B0), _mm512_castps_si512(_B0), 4)); - const __m512 _B2 = _mm512_castsi512_ps(_mm512_permutex_epi64(_mm512_castps_si512(_B0), _MM_SHUFFLE(1, 0, 3, 2))); - const __m512 _B3 = _mm512_castsi512_ps(_mm512_alignr_epi8(_mm512_castps_si512(_B2), _mm512_castps_si512(_B2), 4)); + __m512 _A0 = _mm512_loadu_ps(pA_descales); + __m512 _A1 = _mm512_castsi512_ps(_mm512_alignr_epi8(_mm512_castps_si512(_A0), _mm512_castps_si512(_A0), 8)); + __m256 _b = _mm256_loadu_ps(pB_descales); + __m512 _B0 = combine8x2_ps(_b, _b); + __m512 _B1 = _mm512_castsi512_ps(_mm512_alignr_epi8(_mm512_castps_si512(_B0), _mm512_castps_si512(_B0), 4)); + __m512 _B2 = _mm512_castsi512_ps(_mm512_permutex_epi64(_mm512_castps_si512(_B0), _MM_SHUFFLE(1, 0, 3, 2))); + __m512 _B3 = _mm512_castsi512_ps(_mm512_alignr_epi8(_mm512_castps_si512(_B2), _mm512_castps_si512(_B2), 4)); _fsum0 = _mm512_add_ps(_fsum0, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum0), _mm512_mul_ps(_A0, _B0))); _fsum1 = _mm512_add_ps(_fsum1, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum1), _mm512_mul_ps(_A0, _B1))); _fsum2 = _mm512_add_ps(_fsum2, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum2), _mm512_mul_ps(_A1, _B0))); @@ -2603,6 +2303,9 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB_descales += 8; } + pB += remain_K * 8; + pB_descales += remain_blocks * 8; + _mm512_storeu_ps(outptr + 0, _fsum0); _mm512_storeu_ps(outptr + 16, _fsum1); _mm512_storeu_ps(outptr + 32, _fsum2); @@ -2615,28 +2318,30 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } for (; jj + 3 < max_jj; jj += 4) { - __m512 _fsum0 = _mm512_setzero_ps(); - __m512 _fsum1 = _mm512_setzero_ps(); - __m512 _fsum2 = _mm512_setzero_ps(); - __m512 _fsum3 = _mm512_setzero_ps(); + pB += k * 4; + pB_descales += block_start * 4; + __m512 _fsum0 = k == 0 ? _mm512_setzero_ps() : _mm512_loadu_ps(outptr); + __m512 _fsum1 = k == 0 ? _mm512_setzero_ps() : _mm512_loadu_ps(outptr + 16); + __m512 _fsum2 = k == 0 ? _mm512_setzero_ps() : _mm512_loadu_ps(outptr + 32); + __m512 _fsum3 = k == 0 ? _mm512_setzero_ps() : _mm512_loadu_ps(outptr + 48); const signed char* pA = pAT; const float* pA_descales = pAT_descales; - for (int k = 0; k < K; k += block_size) + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) { __m512i _sum0 = _mm512_setzero_si512(); __m512i _sum1 = _mm512_setzero_si512(); __m512i _sum2 = _mm512_setzero_si512(); __m512i _sum3 = _mm512_setzero_si512(); - const int max_kk = std::min(K - k, block_size); + const int max_kk = std::min(tile_K - kk0, block_size); int kk = 0; #if __AVX512VNNI__ for (; kk + 3 < max_kk; kk += 4) { - const __m512i _pA0 = _mm512_loadu_si512((const __m512i*)pA); - const __m512i _pB0 = _mm512_broadcast_i32x4(_mm_loadu_si128((const __m128i*)pB)); - const __m512i _pA1 = _mm512_alignr_epi8(_pA0, _pA0, 8); - const __m512i _pB1 = _mm512_alignr_epi8(_pB0, _pB0, 4); + __m512i _pA0 = _mm512_loadu_si512((const __m512i*)pA); + __m512i _pB0 = _mm512_broadcast_i32x4(_mm_loadu_si128((const __m128i*)pB)); + __m512i _pA1 = _mm512_alignr_epi8(_pA0, _pA0, 8); + __m512i _pB1 = _mm512_alignr_epi8(_pB0, _pB0, 4); _sum0 = _mm512_dpbusd_epi32(_sum0, _pB0, _pA0); _sum1 = _mm512_dpbusd_epi32(_sum1, _pB1, _pA0); _sum2 = _mm512_dpbusd_epi32(_sum2, _pB0, _pA1); @@ -2646,8 +2351,8 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } if (max_kk >= 4) { - const __m512i _shift0 = _mm512_loadu_si512((const __m512i*)pA); - const __m512i _shift1 = _mm512_alignr_epi8(_shift0, _shift0, 8); + __m512i _shift0 = _mm512_loadu_si512((const __m512i*)pA); + __m512i _shift1 = _mm512_alignr_epi8(_shift0, _shift0, 8); _sum0 = _mm512_sub_epi32(_sum0, _shift0); _sum1 = _mm512_sub_epi32(_sum1, _shift0); _sum2 = _mm512_sub_epi32(_sum2, _shift1); @@ -2657,12 +2362,12 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #endif // __AVX512VNNI__ for (; kk + 1 < max_kk; kk += 2) { - const __m256i _pA = _mm256_loadu_si256((const __m256i*)pA); - const __m512i _pA0 = _mm512_cvtepi8_epi16(_pA); - const __m512i _pA1 = _mm512_alignr_epi8(_pA0, _pA0, 8); - const __m256i _pB = _mm256_castpd_si256(_mm256_broadcast_sd((const double*)pB)); - const __m512i _pB0 = _mm512_cvtepi8_epi16(_pB); - const __m512i _pB1 = _mm512_alignr_epi8(_pB0, _pB0, 4); + __m256i _pA = _mm256_loadu_si256((const __m256i*)pA); + __m512i _pA0 = _mm512_cvtepi8_epi16(_pA); + __m512i _pA1 = _mm512_alignr_epi8(_pA0, _pA0, 8); + __m256i _pB = _mm256_castpd_si256(_mm256_broadcast_sd((const double*)pB)); + __m512i _pB0 = _mm512_cvtepi8_epi16(_pB); + __m512i _pB1 = _mm512_alignr_epi8(_pB0, _pB0, 4); _sum0 = _mm512_comp_dpwssd_epi32(_sum0, _pA0, _pB0); _sum1 = _mm512_comp_dpwssd_epi32(_sum1, _pA0, _pB1); _sum2 = _mm512_comp_dpwssd_epi32(_sum2, _pA1, _pB0); @@ -2672,11 +2377,11 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } for (; kk < max_kk; kk++) { - const __m128i _pA = _mm_loadu_si128((const __m128i*)pA); - const __m256i _pA0 = _mm256_cvtepi8_epi16(_pA); - const __m256i _pA1 = _mm256_shuffle_epi32(_pA0, _MM_SHUFFLE(2, 3, 0, 1)); - const __m256i _pB0 = _mm256_cvtepi8_epi16(_mm_castps_si128(_mm_load1_ps((const float*)pB))); - const __m256i _pB1 = _mm256_shufflehi_epi16(_mm256_shufflelo_epi16(_pB0, _MM_SHUFFLE(0, 3, 2, 1)), _MM_SHUFFLE(0, 3, 2, 1)); + __m128i _pA = _mm_loadu_si128((const __m128i*)pA); + __m256i _pA0 = _mm256_cvtepi8_epi16(_pA); + __m256i _pA1 = _mm256_shuffle_epi32(_pA0, _MM_SHUFFLE(2, 3, 0, 1)); + __m256i _pB0 = _mm256_cvtepi8_epi16(_mm_castps_si128(_mm_load1_ps((const float*)pB))); + __m256i _pB1 = _mm256_shufflehi_epi16(_mm256_shufflelo_epi16(_pB0, _MM_SHUFFLE(0, 3, 2, 1)), _MM_SHUFFLE(0, 3, 2, 1)); _sum0 = _mm512_add_epi32(_sum0, _mm512_cvtepi16_epi32(_mm256_mullo_epi16(_pA0, _pB0))); _sum1 = _mm512_add_epi32(_sum1, _mm512_cvtepi16_epi32(_mm256_mullo_epi16(_pA0, _pB1))); _sum2 = _mm512_add_epi32(_sum2, _mm512_cvtepi16_epi32(_mm256_mullo_epi16(_pA1, _pB0))); @@ -2685,10 +2390,10 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pA += 16; } - const __m512 _A0 = _mm512_loadu_ps(pA_descales); - const __m512 _A1 = _mm512_castsi512_ps(_mm512_alignr_epi8(_mm512_castps_si512(_A0), _mm512_castps_si512(_A0), 8)); - const __m512 _B0 = _mm512_broadcast_f32x4(_mm_loadu_ps(pB_descales)); - const __m512 _B1 = _mm512_castsi512_ps(_mm512_alignr_epi8(_mm512_castps_si512(_B0), _mm512_castps_si512(_B0), 4)); + __m512 _A0 = _mm512_loadu_ps(pA_descales); + __m512 _A1 = _mm512_castsi512_ps(_mm512_alignr_epi8(_mm512_castps_si512(_A0), _mm512_castps_si512(_A0), 8)); + __m512 _B0 = _mm512_broadcast_f32x4(_mm_loadu_ps(pB_descales)); + __m512 _B1 = _mm512_castsi512_ps(_mm512_alignr_epi8(_mm512_castps_si512(_B0), _mm512_castps_si512(_B0), 4)); _fsum0 = _mm512_add_ps(_fsum0, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum0), _mm512_mul_ps(_A0, _B0))); _fsum1 = _mm512_add_ps(_fsum1, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum1), _mm512_mul_ps(_A0, _B1))); _fsum2 = _mm512_add_ps(_fsum2, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum2), _mm512_mul_ps(_A1, _B0))); @@ -2697,6 +2402,9 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB_descales += 4; } + pB += remain_K * 4; + pB_descales += remain_blocks * 4; + _mm512_storeu_ps(outptr + 0, _fsum0); _mm512_storeu_ps(outptr + 16, _fsum1); _mm512_storeu_ps(outptr + 32, _fsum2); @@ -2706,23 +2414,25 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #endif // defined(__x86_64__) || defined(_M_X64) for (; jj + 1 < max_jj; jj += 2) { - __m512 _fsum0 = _mm512_setzero_ps(); - __m512 _fsum1 = _mm512_setzero_ps(); + pB += k * 2; + pB_descales += block_start * 2; + __m512 _fsum0 = k == 0 ? _mm512_setzero_ps() : _mm512_loadu_ps(outptr); + __m512 _fsum1 = k == 0 ? _mm512_setzero_ps() : _mm512_loadu_ps(outptr + 16); const signed char* pA = pAT; const float* pA_descales = pAT_descales; - for (int k = 0; k < K; k += block_size) + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) { __m512i _sum0 = _mm512_setzero_si512(); __m512i _sum1 = _mm512_setzero_si512(); - const int max_kk = std::min(K - k, block_size); + const int max_kk = std::min(tile_K - kk0, block_size); int kk = 0; #if __AVX512VNNI__ for (; kk + 3 < max_kk; kk += 4) { - const __m512i _pA0 = _mm512_loadu_si512((const __m512i*)pA); - const __m512i _pB0 = _mm512_castpd_si512(_mm512_set1_pd(*(const double*)pB)); - const __m512i _pB1 = _mm512_alignr_epi8(_pB0, _pB0, 4); + __m512i _pA0 = _mm512_loadu_si512((const __m512i*)pA); + __m512i _pB0 = _mm512_set1_epi64(_mm_cvtsi128_si64(_mm_loadu_si64(pB))); + __m512i _pB1 = _mm512_alignr_epi8(_pB0, _pB0, 4); _sum0 = _mm512_dpbusd_epi32(_sum0, _pB0, _pA0); _sum1 = _mm512_dpbusd_epi32(_sum1, _pB1, _pA0); pB += 8; @@ -2730,7 +2440,7 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } if (max_kk >= 4) { - const __m512i _shift0 = _mm512_loadu_si512((const __m512i*)pA); + __m512i _shift0 = _mm512_loadu_si512((const __m512i*)pA); _sum0 = _mm512_sub_epi32(_sum0, _shift0); _sum1 = _mm512_sub_epi32(_sum1, _shift0); pA += 64; @@ -2738,11 +2448,11 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #endif // __AVX512VNNI__ for (; kk + 1 < max_kk; kk += 2) { - const __m256i _pA = _mm256_loadu_si256((const __m256i*)pA); - const __m512i _pA0 = _mm512_cvtepi8_epi16(_pA); - const __m256i _pB = _mm256_castps_si256(_mm256_broadcast_ss((const float*)pB)); - const __m512i _pB0 = _mm512_cvtepi8_epi16(_pB); - const __m512i _pB1 = _mm512_alignr_epi8(_pB0, _pB0, 4); + __m256i _pA = _mm256_loadu_si256((const __m256i*)pA); + __m512i _pA0 = _mm512_cvtepi8_epi16(_pA); + __m256i _pB = _mm256_castps_si256(_mm256_broadcast_ss((const float*)pB)); + __m512i _pB0 = _mm512_cvtepi8_epi16(_pB); + __m512i _pB1 = _mm512_alignr_epi8(_pB0, _pB0, 4); _sum0 = _mm512_comp_dpwssd_epi32(_sum0, _pA0, _pB0); _sum1 = _mm512_comp_dpwssd_epi32(_sum1, _pA0, _pB1); pB += 4; @@ -2750,81 +2460,89 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } for (; kk < max_kk; kk++) { - const __m128i _pA = _mm_loadu_si128((const __m128i*)pA); - const __m256i _pA0 = _mm256_cvtepi8_epi16(_pA); - const __m128i _pB = _mm_set1_epi16(*(const short*)pB); - const __m256i _pB0 = _mm256_cvtepi8_epi16(_pB); - const __m256i _pB1 = _mm256_shufflehi_epi16(_mm256_shufflelo_epi16(_pB0, _MM_SHUFFLE(0, 1, 0, 1)), _MM_SHUFFLE(0, 1, 0, 1)); + __m128i _pA = _mm_loadu_si128((const __m128i*)pA); + __m256i _pA0 = _mm256_cvtepi8_epi16(_pA); + __m128i _pB = _mm_set1_epi16((short)_mm_cvtsi128_si32(_mm_loadu_si16(pB))); + __m256i _pB0 = _mm256_cvtepi8_epi16(_pB); + __m256i _pB1 = _mm256_shufflehi_epi16(_mm256_shufflelo_epi16(_pB0, _MM_SHUFFLE(0, 1, 0, 1)), _MM_SHUFFLE(0, 1, 0, 1)); _sum0 = _mm512_add_epi32(_sum0, _mm512_cvtepi16_epi32(_mm256_mullo_epi16(_pA0, _pB0))); _sum1 = _mm512_add_epi32(_sum1, _mm512_cvtepi16_epi32(_mm256_mullo_epi16(_pA0, _pB1))); pB += 2; pA += 16; } - const __m512 _A0 = _mm512_loadu_ps(pA_descales); - const __m512 _B0 = _mm512_castpd_ps(_mm512_set1_pd(*(const double*)pB_descales)); - const __m512 _B1 = _mm512_castsi512_ps(_mm512_alignr_epi8(_mm512_castps_si512(_B0), _mm512_castps_si512(_B0), 4)); + __m512 _A0 = _mm512_loadu_ps(pA_descales); + __m512 _B0 = _mm512_castsi512_ps(_mm512_set1_epi64(_mm_cvtsi128_si64(_mm_loadu_si64(pB_descales)))); + __m512 _B1 = _mm512_castsi512_ps(_mm512_alignr_epi8(_mm512_castps_si512(_B0), _mm512_castps_si512(_B0), 4)); _fsum0 = _mm512_add_ps(_fsum0, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum0), _mm512_mul_ps(_A0, _B0))); _fsum1 = _mm512_add_ps(_fsum1, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum1), _mm512_mul_ps(_A0, _B1))); pA_descales += 16; pB_descales += 2; } + pB += remain_K * 2; + pB_descales += remain_blocks * 2; + _mm512_storeu_ps(outptr + 0, _fsum0); _mm512_storeu_ps(outptr + 16, _fsum1); outptr += 32; } for (; jj < max_jj; jj++) { - __m512 _fsum0 = _mm512_setzero_ps(); + pB += k; + pB_descales += block_start; + __m512 _fsum0 = k == 0 ? _mm512_setzero_ps() : _mm512_loadu_ps(outptr); const signed char* pA = pAT; const float* pA_descales = pAT_descales; - for (int k = 0; k < K; k += block_size) + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) { __m512i _sum0 = _mm512_setzero_si512(); - const int max_kk = std::min(K - k, block_size); + const int max_kk = std::min(tile_K - kk0, block_size); int kk = 0; #if __AVX512VNNI__ for (; kk + 3 < max_kk; kk += 4) { - const __m512i _pA0 = _mm512_loadu_si512((const __m512i*)pA); - const __m512i _pB0 = _mm512_set1_epi32(*(const int*)pB); + __m512i _pA0 = _mm512_loadu_si512((const __m512i*)pA); + __m512i _pB0 = _mm512_set1_epi32(_mm_cvtsi128_si32(_mm_loadu_si32(pB))); _sum0 = _mm512_dpbusd_epi32(_sum0, _pB0, _pA0); pB += 4; pA += 64; } if (max_kk >= 4) { - const __m512i _shift0 = _mm512_loadu_si512((const __m512i*)pA); + __m512i _shift0 = _mm512_loadu_si512((const __m512i*)pA); _sum0 = _mm512_sub_epi32(_sum0, _shift0); pA += 64; } #endif // __AVX512VNNI__ for (; kk + 1 < max_kk; kk += 2) { - const __m256i _pA = _mm256_loadu_si256((const __m256i*)pA); - const __m512i _pA0 = _mm512_cvtepi8_epi16(_pA); - const __m512i _pB0 = _mm512_cvtepi8_epi16(_mm256_set1_epi16(*(const short*)pB)); + __m256i _pA = _mm256_loadu_si256((const __m256i*)pA); + __m512i _pA0 = _mm512_cvtepi8_epi16(_pA); + __m512i _pB0 = _mm512_cvtepi8_epi16(_mm256_set1_epi16((short)_mm_cvtsi128_si32(_mm_loadu_si16(pB)))); _sum0 = _mm512_comp_dpwssd_epi32(_sum0, _pA0, _pB0); pB += 2; pA += 32; } for (; kk < max_kk; kk++) { - const __m512i _pA0 = _mm512_cvtepi8_epi32(_mm_loadu_si128((const __m128i*)pA)); + __m512i _pA0 = _mm512_cvtepi8_epi32(_mm_loadu_si128((const __m128i*)pA)); _sum0 = _mm512_add_epi32(_sum0, _mm512_mullo_epi32(_pA0, _mm512_set1_epi32((signed char)pB[0]))); pB += 1; pA += 16; } - const __m512 _A0 = _mm512_loadu_ps(pA_descales); - const __m512 _B0 = _mm512_set1_ps(pB_descales[0]); + __m512 _A0 = _mm512_loadu_ps(pA_descales); + __m512 _B0 = _mm512_set1_ps(pB_descales[0]); _fsum0 = _mm512_add_ps(_fsum0, _mm512_mul_ps(_mm512_cvtepi32_ps(_sum0), _mm512_mul_ps(_A0, _B0))); pA_descales += 16; pB_descales += 1; } + pB += remain_K * 1; + pB_descales += remain_blocks * 1; + _mm512_storeu_ps(outptr + 0, _fsum0); outptr += 16; } @@ -2843,18 +2561,20 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #if __AVX512F__ for (; jj + 7 < max_jj; jj += 8) { - __m256 _fsum0 = _mm256_setzero_ps(); - __m256 _fsum1 = _mm256_setzero_ps(); - __m256 _fsum2 = _mm256_setzero_ps(); - __m256 _fsum3 = _mm256_setzero_ps(); - __m256 _fsum4 = _mm256_setzero_ps(); - __m256 _fsum5 = _mm256_setzero_ps(); - __m256 _fsum6 = _mm256_setzero_ps(); - __m256 _fsum7 = _mm256_setzero_ps(); + pB += k * 8; + pB_descales += block_start * 8; + __m256 _fsum0 = k == 0 ? _mm256_setzero_ps() : _mm256_loadu_ps(outptr); + __m256 _fsum1 = k == 0 ? _mm256_setzero_ps() : _mm256_loadu_ps(outptr + 8); + __m256 _fsum2 = k == 0 ? _mm256_setzero_ps() : _mm256_loadu_ps(outptr + 16); + __m256 _fsum3 = k == 0 ? _mm256_setzero_ps() : _mm256_loadu_ps(outptr + 24); + __m256 _fsum4 = k == 0 ? _mm256_setzero_ps() : _mm256_loadu_ps(outptr + 32); + __m256 _fsum5 = k == 0 ? _mm256_setzero_ps() : _mm256_loadu_ps(outptr + 40); + __m256 _fsum6 = k == 0 ? _mm256_setzero_ps() : _mm256_loadu_ps(outptr + 48); + __m256 _fsum7 = k == 0 ? _mm256_setzero_ps() : _mm256_loadu_ps(outptr + 56); const signed char* pA = pAT; const float* pA_descales = pAT_descales; - for (int k = 0; k < K; k += block_size) + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) { __m256i _sum0 = _mm256_setzero_si256(); __m256i _sum1 = _mm256_setzero_si256(); @@ -2864,17 +2584,17 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de __m256i _sum5 = _mm256_setzero_si256(); __m256i _sum6 = _mm256_setzero_si256(); __m256i _sum7 = _mm256_setzero_si256(); - const int max_kk = std::min(K - k, block_size); + const int max_kk = std::min(tile_K - kk0, block_size); int kk = 0; #if __AVX512VNNI__ for (; kk + 3 < max_kk; kk += 4) { - const __m256i _pA0 = _mm256_loadu_si256((const __m256i*)pA); - const __m256i _pA1 = _mm256_alignr_epi8(_pA0, _pA0, 8); - const __m256i _pB0 = _mm256_loadu_si256((const __m256i*)pB); - const __m256i _pB1 = _mm256_alignr_epi8(_pB0, _pB0, 4); - const __m256i _pB2 = _mm256_permute4x64_epi64(_pB0, _MM_SHUFFLE(1, 0, 3, 2)); - const __m256i _pB3 = _mm256_alignr_epi8(_pB2, _pB2, 4); + __m256i _pA0 = _mm256_loadu_si256((const __m256i*)pA); + __m256i _pA1 = _mm256_alignr_epi8(_pA0, _pA0, 8); + __m256i _pB0 = _mm256_loadu_si256((const __m256i*)pB); + __m256i _pB1 = _mm256_alignr_epi8(_pB0, _pB0, 4); + __m256i _pB2 = _mm256_permute4x64_epi64(_pB0, _MM_SHUFFLE(1, 0, 3, 2)); + __m256i _pB3 = _mm256_alignr_epi8(_pB2, _pB2, 4); _sum0 = _mm256_comp_dpbusd_epi32(_sum0, _pB0, _pA0); _sum1 = _mm256_comp_dpbusd_epi32(_sum1, _pB1, _pA0); _sum2 = _mm256_comp_dpbusd_epi32(_sum2, _pB0, _pA1); @@ -2888,8 +2608,8 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } if (max_kk >= 4) { - const __m256i _shift0 = _mm256_loadu_si256((const __m256i*)pA); - const __m256i _shift1 = _mm256_alignr_epi8(_shift0, _shift0, 8); + __m256i _shift0 = _mm256_loadu_si256((const __m256i*)pA); + __m256i _shift1 = _mm256_alignr_epi8(_shift0, _shift0, 8); _sum0 = _mm256_sub_epi32(_sum0, _shift0); _sum1 = _mm256_sub_epi32(_sum1, _shift0); _sum2 = _mm256_sub_epi32(_sum2, _shift1); @@ -2903,14 +2623,14 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #endif // __AVX512VNNI__ for (; kk + 1 < max_kk; kk += 2) { - const __m128i _pA8 = _mm_loadu_si128((const __m128i*)pA); - const __m128i _pB8 = _mm_loadu_si128((const __m128i*)pB); - const __m256i _pA0 = _mm256_cvtepi8_epi16(_pA8); - const __m256i _pB0 = _mm256_cvtepi8_epi16(_pB8); - const __m256i _pA1 = _mm256_alignr_epi8(_pA0, _pA0, 8); - const __m256i _pB1 = _mm256_alignr_epi8(_pB0, _pB0, 4); - const __m256i _pB2 = _mm256_permute4x64_epi64(_pB0, _MM_SHUFFLE(1, 0, 3, 2)); - const __m256i _pB3 = _mm256_alignr_epi8(_pB2, _pB2, 4); + __m128i _pA8 = _mm_loadu_si128((const __m128i*)pA); + __m128i _pB8 = _mm_loadu_si128((const __m128i*)pB); + __m256i _pA0 = _mm256_cvtepi8_epi16(_pA8); + __m256i _pB0 = _mm256_cvtepi8_epi16(_pB8); + __m256i _pA1 = _mm256_alignr_epi8(_pA0, _pA0, 8); + __m256i _pB1 = _mm256_alignr_epi8(_pB0, _pB0, 4); + __m256i _pB2 = _mm256_permute4x64_epi64(_pB0, _MM_SHUFFLE(1, 0, 3, 2)); + __m256i _pB3 = _mm256_alignr_epi8(_pB2, _pB2, 4); _sum0 = _mm256_comp_dpwssd_epi32(_sum0, _pA0, _pB0); _sum1 = _mm256_comp_dpwssd_epi32(_sum1, _pA0, _pB1); _sum2 = _mm256_comp_dpwssd_epi32(_sum2, _pA1, _pB0); @@ -2928,10 +2648,10 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de __m128i _pB0 = _mm_loadl_epi64((const __m128i*)pB); _pA0 = _mm_cvtepi8_epi16(_pA0); _pB0 = _mm_cvtepi8_epi16(_pB0); - const __m128i _pA1 = _mm_shufflehi_epi16(_mm_shufflelo_epi16(_pA0, _MM_SHUFFLE(1, 0, 3, 2)), _MM_SHUFFLE(1, 0, 3, 2)); - const __m128i _pB1 = _mm_shufflehi_epi16(_mm_shufflelo_epi16(_pB0, _MM_SHUFFLE(0, 3, 2, 1)), _MM_SHUFFLE(0, 3, 2, 1)); - const __m128i _pB2 = _mm_alignr_epi8(_pB0, _pB0, 8); - const __m128i _pB3 = _mm_shufflehi_epi16(_mm_shufflelo_epi16(_pB2, _MM_SHUFFLE(0, 3, 2, 1)), _MM_SHUFFLE(0, 3, 2, 1)); + __m128i _pA1 = _mm_shufflehi_epi16(_mm_shufflelo_epi16(_pA0, _MM_SHUFFLE(1, 0, 3, 2)), _MM_SHUFFLE(1, 0, 3, 2)); + __m128i _pB1 = _mm_shufflehi_epi16(_mm_shufflelo_epi16(_pB0, _MM_SHUFFLE(0, 3, 2, 1)), _MM_SHUFFLE(0, 3, 2, 1)); + __m128i _pB2 = _mm_alignr_epi8(_pB0, _pB0, 8); + __m128i _pB3 = _mm_shufflehi_epi16(_mm_shufflelo_epi16(_pB2, _MM_SHUFFLE(0, 3, 2, 1)), _MM_SHUFFLE(0, 3, 2, 1)); _sum0 = _mm256_add_epi32(_sum0, _mm256_cvtepi16_epi32(_mm_mullo_epi16(_pA0, _pB0))); _sum1 = _mm256_add_epi32(_sum1, _mm256_cvtepi16_epi32(_mm_mullo_epi16(_pA0, _pB1))); _sum2 = _mm256_add_epi32(_sum2, _mm256_cvtepi16_epi32(_mm_mullo_epi16(_pA1, _pB0))); @@ -2943,12 +2663,12 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB += 8; } - const __m256 _A0 = _mm256_loadu_ps(pA_descales); - const __m256 _A1 = _mm256_castsi256_ps(_mm256_alignr_epi8(_mm256_castps_si256(_A0), _mm256_castps_si256(_A0), 8)); - const __m256 _B0 = _mm256_loadu_ps(pB_descales); - const __m256 _B1 = _mm256_castsi256_ps(_mm256_alignr_epi8(_mm256_castps_si256(_B0), _mm256_castps_si256(_B0), 4)); - const __m256 _B2 = _mm256_castsi256_ps(_mm256_permute4x64_epi64(_mm256_castps_si256(_B0), _MM_SHUFFLE(1, 0, 3, 2))); - const __m256 _B3 = _mm256_castsi256_ps(_mm256_alignr_epi8(_mm256_castps_si256(_B2), _mm256_castps_si256(_B2), 4)); + __m256 _A0 = _mm256_loadu_ps(pA_descales); + __m256 _A1 = _mm256_castsi256_ps(_mm256_alignr_epi8(_mm256_castps_si256(_A0), _mm256_castps_si256(_A0), 8)); + __m256 _B0 = _mm256_loadu_ps(pB_descales); + __m256 _B1 = _mm256_castsi256_ps(_mm256_alignr_epi8(_mm256_castps_si256(_B0), _mm256_castps_si256(_B0), 4)); + __m256 _B2 = _mm256_castsi256_ps(_mm256_permute4x64_epi64(_mm256_castps_si256(_B0), _MM_SHUFFLE(1, 0, 3, 2))); + __m256 _B3 = _mm256_castsi256_ps(_mm256_alignr_epi8(_mm256_castps_si256(_B2), _mm256_castps_si256(_B2), 4)); _fsum0 = _mm256_add_ps(_fsum0, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum0), _mm256_mul_ps(_A0, _B0))); _fsum1 = _mm256_add_ps(_fsum1, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum1), _mm256_mul_ps(_A0, _B1))); _fsum2 = _mm256_add_ps(_fsum2, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum2), _mm256_mul_ps(_A1, _B0))); @@ -2961,6 +2681,9 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB_descales += 8; } + pB += remain_K * 8; + pB_descales += remain_blocks * 8; + _mm256_storeu_ps(outptr + 0, _fsum0); _mm256_storeu_ps(outptr + 8, _fsum1); _mm256_storeu_ps(outptr + 16, _fsum2); @@ -2974,29 +2697,31 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #endif // __AVX512F__ for (; jj + 3 < max_jj; jj += 4) { - __m256 _fsum0 = _mm256_setzero_ps(); - __m256 _fsum1 = _mm256_setzero_ps(); - __m256 _fsum2 = _mm256_setzero_ps(); - __m256 _fsum3 = _mm256_setzero_ps(); + pB += k * 4; + pB_descales += block_start * 4; + __m256 _fsum0 = k == 0 ? _mm256_setzero_ps() : _mm256_loadu_ps(outptr); + __m256 _fsum1 = k == 0 ? _mm256_setzero_ps() : _mm256_loadu_ps(outptr + 8); + __m256 _fsum2 = k == 0 ? _mm256_setzero_ps() : _mm256_loadu_ps(outptr + 16); + __m256 _fsum3 = k == 0 ? _mm256_setzero_ps() : _mm256_loadu_ps(outptr + 24); const signed char* pA = pAT; const float* pA_descales = pAT_descales; - for (int k = 0; k < K; k += block_size) + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) { __m256i _sum0 = _mm256_setzero_si256(); __m256i _sum1 = _mm256_setzero_si256(); __m256i _sum2 = _mm256_setzero_si256(); __m256i _sum3 = _mm256_setzero_si256(); - const int max_kk = std::min(K - k, block_size); + const int max_kk = std::min(tile_K - kk0, block_size); int kk = 0; #if __AVX512VNNI__ || __AVXVNNI__ for (; kk + 3 < max_kk; kk += 4) { - const __m256i _pA0 = _mm256_loadu_si256((const __m256i*)pA); - const __m128i _pB = _mm_loadu_si128((const __m128i*)pB); - const __m256i _pB0 = combine4x2_epi32(_pB, _pB); - const __m256i _pA1 = _mm256_alignr_epi8(_pA0, _pA0, 8); - const __m256i _pB1 = _mm256_alignr_epi8(_pB0, _pB0, 4); + __m256i _pA0 = _mm256_loadu_si256((const __m256i*)pA); + __m128i _pB = _mm_loadu_si128((const __m128i*)pB); + __m256i _pB0 = combine4x2_epi32(_pB, _pB); + __m256i _pA1 = _mm256_alignr_epi8(_pA0, _pA0, 8); + __m256i _pB1 = _mm256_alignr_epi8(_pB0, _pB0, 4); #if __AVXVNNIINT8__ _sum0 = _mm256_dpbssd_epi32(_sum0, _pB0, _pA0); _sum1 = _mm256_dpbssd_epi32(_sum1, _pB1, _pA0); @@ -3014,8 +2739,8 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) if (max_kk >= 4) { - const __m256i _shift0 = _mm256_loadu_si256((const __m256i*)pA); - const __m256i _shift1 = _mm256_alignr_epi8(_shift0, _shift0, 8); + __m256i _shift0 = _mm256_loadu_si256((const __m256i*)pA); + __m256i _shift1 = _mm256_alignr_epi8(_shift0, _shift0, 8); _sum0 = _mm256_sub_epi32(_sum0, _shift0); _sum1 = _mm256_sub_epi32(_sum1, _shift0); _sum2 = _mm256_sub_epi32(_sum2, _shift1); @@ -3026,12 +2751,12 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #endif // __AVX512VNNI__ || __AVXVNNI__ for (; kk + 1 < max_kk; kk += 2) { - const __m128i _pA = _mm_loadu_si128((const __m128i*)pA); - const __m128i _pB = _mm_castpd_si128(_mm_load1_pd((const double*)pB)); - const __m256i _pA0 = _mm256_cvtepi8_epi16(_pA); - const __m256i _pA1 = _mm256_alignr_epi8(_pA0, _pA0, 8); - const __m256i _pB0 = _mm256_cvtepi8_epi16(_pB); - const __m256i _pB1 = _mm256_alignr_epi8(_pB0, _pB0, 4); + __m128i _pA = _mm_loadu_si128((const __m128i*)pA); + __m128i _pB = _mm_castpd_si128(_mm_load1_pd((const double*)pB)); + __m256i _pA0 = _mm256_cvtepi8_epi16(_pA); + __m256i _pA1 = _mm256_alignr_epi8(_pA0, _pA0, 8); + __m256i _pB0 = _mm256_cvtepi8_epi16(_pB); + __m256i _pB1 = _mm256_alignr_epi8(_pB0, _pB0, 4); _sum0 = _mm256_comp_dpwssd_epi32(_sum0, _pA0, _pB0); _sum1 = _mm256_comp_dpwssd_epi32(_sum1, _pA0, _pB1); _sum2 = _mm256_comp_dpwssd_epi32(_sum2, _pA1, _pB0); @@ -3045,8 +2770,8 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de __m128i _pB0 = _mm_castps_si128(_mm_load1_ps((const float*)pB)); _pA0 = _mm_cvtepi8_epi16(_pA0); _pB0 = _mm_cvtepi8_epi16(_pB0); - const __m128i _pA1 = _mm_shuffle_epi32(_pA0, _MM_SHUFFLE(2, 3, 0, 1)); - const __m128i _pB1 = _mm_shufflehi_epi16(_mm_shufflelo_epi16(_pB0, _MM_SHUFFLE(0, 3, 2, 1)), _MM_SHUFFLE(0, 3, 2, 1)); + __m128i _pA1 = _mm_shuffle_epi32(_pA0, _MM_SHUFFLE(2, 3, 0, 1)); + __m128i _pB1 = _mm_shufflehi_epi16(_mm_shufflelo_epi16(_pB0, _MM_SHUFFLE(0, 3, 2, 1)), _MM_SHUFFLE(0, 3, 2, 1)); _sum0 = _mm256_add_epi32(_sum0, _mm256_cvtepi16_epi32(_mm_mullo_epi16(_pA0, _pB0))); _sum1 = _mm256_add_epi32(_sum1, _mm256_cvtepi16_epi32(_mm_mullo_epi16(_pA0, _pB1))); _sum2 = _mm256_add_epi32(_sum2, _mm256_cvtepi16_epi32(_mm_mullo_epi16(_pA1, _pB0))); @@ -3055,11 +2780,11 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pA += 8; } - const __m256 _A0 = _mm256_loadu_ps(pA_descales); - const __m256 _A1 = _mm256_castsi256_ps(_mm256_alignr_epi8(_mm256_castps_si256(_A0), _mm256_castps_si256(_A0), 8)); - const __m128 _b = _mm_loadu_ps(pB_descales); - const __m256 _B0 = combine4x2_ps(_b, _b); - const __m256 _B1 = _mm256_castsi256_ps(_mm256_alignr_epi8(_mm256_castps_si256(_B0), _mm256_castps_si256(_B0), 4)); + __m256 _A0 = _mm256_loadu_ps(pA_descales); + __m256 _A1 = _mm256_castsi256_ps(_mm256_alignr_epi8(_mm256_castps_si256(_A0), _mm256_castps_si256(_A0), 8)); + __m128 _b = _mm_loadu_ps(pB_descales); + __m256 _B0 = combine4x2_ps(_b, _b); + __m256 _B1 = _mm256_castsi256_ps(_mm256_alignr_epi8(_mm256_castps_si256(_B0), _mm256_castps_si256(_B0), 4)); _fsum0 = _mm256_add_ps(_fsum0, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum0), _mm256_mul_ps(_A0, _B0))); _fsum1 = _mm256_add_ps(_fsum1, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum1), _mm256_mul_ps(_A0, _B1))); _fsum2 = _mm256_add_ps(_fsum2, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum2), _mm256_mul_ps(_A1, _B0))); @@ -3068,6 +2793,9 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB_descales += 4; } + pB += remain_K * 4; + pB_descales += remain_blocks * 4; + _mm256_storeu_ps(outptr + 0, _fsum0); _mm256_storeu_ps(outptr + 8, _fsum1); _mm256_storeu_ps(outptr + 16, _fsum2); @@ -3077,23 +2805,25 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #endif // defined(__x86_64__) || defined(_M_X64) for (; jj + 1 < max_jj; jj += 2) { - __m256 _fsum0 = _mm256_setzero_ps(); - __m256 _fsum1 = _mm256_setzero_ps(); + pB += k * 2; + pB_descales += block_start * 2; + __m256 _fsum0 = k == 0 ? _mm256_setzero_ps() : _mm256_loadu_ps(outptr); + __m256 _fsum1 = k == 0 ? _mm256_setzero_ps() : _mm256_loadu_ps(outptr + 8); const signed char* pA = pAT; const float* pA_descales = pAT_descales; - for (int k = 0; k < K; k += block_size) + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) { __m256i _sum0 = _mm256_setzero_si256(); __m256i _sum1 = _mm256_setzero_si256(); - const int max_kk = std::min(K - k, block_size); + const int max_kk = std::min(tile_K - kk0, block_size); int kk = 0; #if __AVX512VNNI__ || __AVXVNNI__ for (; kk + 3 < max_kk; kk += 4) { - const __m256i _pA0 = _mm256_loadu_si256((const __m256i*)pA); - const __m256i _pB0 = _mm256_castpd_si256(_mm256_broadcast_sd((const double*)pB)); - const __m256i _pB1 = _mm256_alignr_epi8(_pB0, _pB0, 4); + __m256i _pA0 = _mm256_loadu_si256((const __m256i*)pA); + __m256i _pB0 = _mm256_castpd_si256(_mm256_broadcast_sd((const double*)pB)); + __m256i _pB1 = _mm256_alignr_epi8(_pB0, _pB0, 4); #if __AVXVNNIINT8__ _sum0 = _mm256_dpbssd_epi32(_sum0, _pB0, _pA0); _sum1 = _mm256_dpbssd_epi32(_sum1, _pB1, _pA0); @@ -3107,7 +2837,7 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) if (max_kk >= 4) { - const __m256i _shift0 = _mm256_loadu_si256((const __m256i*)pA); + __m256i _shift0 = _mm256_loadu_si256((const __m256i*)pA); _sum0 = _mm256_sub_epi32(_sum0, _shift0); _sum1 = _mm256_sub_epi32(_sum1, _shift0); pA += 32; @@ -3116,14 +2846,14 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #else for (; kk + 3 < max_kk; kk += 4) { - const __m256i _pA = _mm256_loadu_si256((const __m256i*)pA); - const __m128i _pB = _mm_loadl_epi64((const __m128i*)pB); - const __m256i _pA01 = _mm256_cvtepi8_epi16(_mm256_castsi256_si128(_pA)); - const __m256i _pA23 = _mm256_cvtepi8_epi16(_mm256_extracti128_si256(_pA, 1)); - const __m256i _pB01 = _mm256_cvtepi8_epi16(_mm_shuffle_epi32(_pB, _MM_SHUFFLE(0, 0, 0, 0))); - const __m256i _pB23 = _mm256_cvtepi8_epi16(_mm_shuffle_epi32(_pB, _MM_SHUFFLE(1, 1, 1, 1))); - const __m256i _pB01_1 = _mm256_alignr_epi8(_pB01, _pB01, 4); - const __m256i _pB23_1 = _mm256_alignr_epi8(_pB23, _pB23, 4); + __m256i _pA = _mm256_loadu_si256((const __m256i*)pA); + __m128i _pB = _mm_loadl_epi64((const __m128i*)pB); + __m256i _pA01 = _mm256_cvtepi8_epi16(_mm256_castsi256_si128(_pA)); + __m256i _pA23 = _mm256_cvtepi8_epi16(_mm256_extracti128_si256(_pA, 1)); + __m256i _pB01 = _mm256_cvtepi8_epi16(_mm_shuffle_epi32(_pB, _MM_SHUFFLE(0, 0, 0, 0))); + __m256i _pB23 = _mm256_cvtepi8_epi16(_mm_shuffle_epi32(_pB, _MM_SHUFFLE(1, 1, 1, 1))); + __m256i _pB01_1 = _mm256_alignr_epi8(_pB01, _pB01, 4); + __m256i _pB23_1 = _mm256_alignr_epi8(_pB23, _pB23, 4); _sum0 = _mm256_comp_dpwssd_epi32(_sum0, _pA01, _pB01); _sum1 = _mm256_comp_dpwssd_epi32(_sum1, _pA01, _pB01_1); _sum0 = _mm256_comp_dpwssd_epi32(_sum0, _pA23, _pB23); @@ -3134,11 +2864,11 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #endif // __AVX512VNNI__ || __AVXVNNI__ for (; kk + 1 < max_kk; kk += 2) { - const __m128i _pA = _mm_loadu_si128((const __m128i*)pA); - const __m128i _pB = _mm_castps_si128(_mm_load1_ps((const float*)pB)); - const __m256i _pA0 = _mm256_cvtepi8_epi16(_pA); - const __m256i _pB0 = _mm256_cvtepi8_epi16(_pB); - const __m256i _pB1 = _mm256_alignr_epi8(_pB0, _pB0, 4); + __m128i _pA = _mm_loadu_si128((const __m128i*)pA); + __m128i _pB = _mm_castps_si128(_mm_load1_ps((const float*)pB)); + __m256i _pA0 = _mm256_cvtepi8_epi16(_pA); + __m256i _pB0 = _mm256_cvtepi8_epi16(_pB); + __m256i _pB1 = _mm256_alignr_epi8(_pB0, _pB0, 4); _sum0 = _mm256_comp_dpwssd_epi32(_sum0, _pA0, _pB0); _sum1 = _mm256_comp_dpwssd_epi32(_sum1, _pA0, _pB1); pB += 4; @@ -3147,45 +2877,50 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de for (; kk < max_kk; kk++) { __m128i _pA = _mm_loadl_epi64((const __m128i*)pA); - __m128i _pB0 = _mm_set1_epi16(*(const short*)pB); + __m128i _pB0 = _mm_set1_epi16((short)_mm_cvtsi128_si32(_mm_loadu_si16(pB))); _pA = _mm_cvtepi8_epi16(_pA); _pB0 = _mm_cvtepi8_epi16(_pB0); - const __m128i _pB1 = _mm_shufflehi_epi16(_mm_shufflelo_epi16(_pB0, _MM_SHUFFLE(0, 1, 0, 1)), _MM_SHUFFLE(0, 1, 0, 1)); + __m128i _pB1 = _mm_shufflehi_epi16(_mm_shufflelo_epi16(_pB0, _MM_SHUFFLE(0, 1, 0, 1)), _MM_SHUFFLE(0, 1, 0, 1)); _sum0 = _mm256_add_epi32(_sum0, _mm256_cvtepi16_epi32(_mm_mullo_epi16(_pA, _pB0))); _sum1 = _mm256_add_epi32(_sum1, _mm256_cvtepi16_epi32(_mm_mullo_epi16(_pA, _pB1))); pB += 2; pA += 8; } - const __m256 _A0 = _mm256_loadu_ps(pA_descales); - const __m256 _B0 = _mm256_castpd_ps(_mm256_broadcast_sd((const double*)pB_descales)); - const __m256 _B1 = _mm256_castsi256_ps(_mm256_alignr_epi8(_mm256_castps_si256(_B0), _mm256_castps_si256(_B0), 4)); + __m256 _A0 = _mm256_loadu_ps(pA_descales); + __m256 _B0 = _mm256_castpd_ps(_mm256_broadcast_sd((const double*)pB_descales)); + __m256 _B1 = _mm256_castsi256_ps(_mm256_alignr_epi8(_mm256_castps_si256(_B0), _mm256_castps_si256(_B0), 4)); _fsum0 = _mm256_add_ps(_fsum0, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum0), _mm256_mul_ps(_A0, _B0))); _fsum1 = _mm256_add_ps(_fsum1, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum1), _mm256_mul_ps(_A0, _B1))); pA_descales += 8; pB_descales += 2; } + pB += remain_K * 2; + pB_descales += remain_blocks * 2; + _mm256_storeu_ps(outptr + 0, _fsum0); _mm256_storeu_ps(outptr + 8, _fsum1); outptr += 16; } for (; jj < max_jj; jj++) { - __m256 _fsum0 = _mm256_setzero_ps(); + pB += k; + pB_descales += block_start; + __m256 _fsum0 = k == 0 ? _mm256_setzero_ps() : _mm256_loadu_ps(outptr); const signed char* pA = pAT; const float* pA_descales = pAT_descales; - for (int k = 0; k < K; k += block_size) + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) { __m256i _sum0 = _mm256_setzero_si256(); - const int max_kk = std::min(K - k, block_size); + const int max_kk = std::min(tile_K - kk0, block_size); int kk = 0; #if __AVX512VNNI__ || __AVXVNNI__ for (; kk + 3 < max_kk; kk += 4) { - const __m256i _pA0 = _mm256_loadu_si256((const __m256i*)pA); - const __m256i _pB0 = _mm256_castps_si256(_mm256_broadcast_ss((const float*)pB)); + __m256i _pA0 = _mm256_loadu_si256((const __m256i*)pA); + __m256i _pB0 = _mm256_castps_si256(_mm256_broadcast_ss((const float*)pB)); #if __AVXVNNIINT8__ _sum0 = _mm256_dpbssd_epi32(_sum0, _pB0, _pA0); #else // __AVXVNNIINT8__ @@ -3197,7 +2932,7 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) if (max_kk >= 4) { - const __m256i _shift0 = _mm256_loadu_si256((const __m256i*)pA); + __m256i _shift0 = _mm256_loadu_si256((const __m256i*)pA); _sum0 = _mm256_sub_epi32(_sum0, _shift0); pA += 32; } @@ -3205,12 +2940,12 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #else for (; kk + 3 < max_kk; kk += 4) { - const __m256i _pA = _mm256_loadu_si256((const __m256i*)pA); - const __m256i _pA01 = _mm256_cvtepi8_epi16(_mm256_castsi256_si128(_pA)); - const __m256i _pA23 = _mm256_cvtepi8_epi16(_mm256_extracti128_si256(_pA, 1)); - const __m128i _pB16 = _mm_cvtepi8_epi16(_mm_cvtsi32_si128(*(const int*)pB)); - const __m256i _pB01 = _mm256_broadcastsi128_si256(_mm_shuffle_epi32(_pB16, _MM_SHUFFLE(0, 0, 0, 0))); - const __m256i _pB23 = _mm256_broadcastsi128_si256(_mm_shuffle_epi32(_pB16, _MM_SHUFFLE(1, 1, 1, 1))); + __m256i _pA = _mm256_loadu_si256((const __m256i*)pA); + __m256i _pA01 = _mm256_cvtepi8_epi16(_mm256_castsi256_si128(_pA)); + __m256i _pA23 = _mm256_cvtepi8_epi16(_mm256_extracti128_si256(_pA, 1)); + __m128i _pB16 = _mm_cvtepi8_epi16(_mm_loadu_si32(pB)); + __m256i _pB01 = _mm256_broadcastsi128_si256(_mm_shuffle_epi32(_pB16, _MM_SHUFFLE(0, 0, 0, 0))); + __m256i _pB23 = _mm256_broadcastsi128_si256(_mm_shuffle_epi32(_pB16, _MM_SHUFFLE(1, 1, 1, 1))); _sum0 = _mm256_comp_dpwssd_epi32(_sum0, _pA01, _pB01); _sum0 = _mm256_comp_dpwssd_epi32(_sum0, _pA23, _pB23); pB += 4; @@ -3219,9 +2954,9 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #endif // __AVX512VNNI__ || __AVXVNNI__ for (; kk + 1 < max_kk; kk += 2) { - const __m128i _pA = _mm_loadu_si128((const __m128i*)pA); - const __m256i _pA0 = _mm256_cvtepi8_epi16(_pA); - const __m256i _pB0 = _mm256_cvtepi8_epi16(_mm_set1_epi16(*(const short*)pB)); + __m128i _pA = _mm_loadu_si128((const __m128i*)pA); + __m256i _pA0 = _mm256_cvtepi8_epi16(_pA); + __m256i _pB0 = _mm256_cvtepi8_epi16(_mm_set1_epi16((short)_mm_cvtsi128_si32(_mm_loadu_si16(pB)))); _sum0 = _mm256_comp_dpwssd_epi32(_sum0, _pA0, _pB0); pB += 2; pA += 16; @@ -3235,12 +2970,15 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pA += 8; } - const __m256 _A0 = _mm256_loadu_ps(pA_descales); + __m256 _A0 = _mm256_loadu_ps(pA_descales); _fsum0 = _mm256_add_ps(_fsum0, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum0), _mm256_mul_ps(_A0, _mm256_set1_ps(pB_descales[0])))); pA_descales += 8; pB_descales++; } + pB += remain_K * 1; + pB_descales += remain_blocks * 1; + _mm256_storeu_ps(outptr, _fsum0); outptr += 8; } @@ -3259,28 +2997,30 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #if __AVX512F__ for (; jj + 7 < max_jj; jj += 8) { - __m256 _fsum0 = _mm256_setzero_ps(); - __m256 _fsum1 = _mm256_setzero_ps(); - __m256 _fsum2 = _mm256_setzero_ps(); - __m256 _fsum3 = _mm256_setzero_ps(); + pB += k * 8; + pB_descales += block_start * 8; + __m256 _fsum0 = k == 0 ? _mm256_setzero_ps() : _mm256_loadu_ps(outptr); + __m256 _fsum1 = k == 0 ? _mm256_setzero_ps() : _mm256_loadu_ps(outptr + 8); + __m256 _fsum2 = k == 0 ? _mm256_setzero_ps() : _mm256_loadu_ps(outptr + 16); + __m256 _fsum3 = k == 0 ? _mm256_setzero_ps() : _mm256_loadu_ps(outptr + 24); const signed char* pA = pAT; const float* pA_descales = pAT_descales; - for (int k = 0; k < K; k += block_size) + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) { __m256i _sum0 = _mm256_setzero_si256(); __m256i _sum1 = _mm256_setzero_si256(); __m256i _sum2 = _mm256_setzero_si256(); __m256i _sum3 = _mm256_setzero_si256(); - const int max_kk = std::min(K - k, block_size); + const int max_kk = std::min(tile_K - kk0, block_size); int kk = 0; #if __AVX512VNNI__ for (; kk + 3 < max_kk; kk += 4) { - const __m256i _pA0 = _mm256_broadcastsi128_si256(_mm_loadu_si128((const __m128i*)pA)); - const __m256i _pA1 = _mm256_alignr_epi8(_pA0, _pA0, 8); - const __m256i _pB0 = _mm256_loadu_si256((const __m256i*)pB); - const __m256i _pB1 = _mm256_alignr_epi8(_pB0, _pB0, 4); + __m256i _pA0 = _mm256_broadcastsi128_si256(_mm_loadu_si128((const __m128i*)pA)); + __m256i _pA1 = _mm256_alignr_epi8(_pA0, _pA0, 8); + __m256i _pB0 = _mm256_loadu_si256((const __m256i*)pB); + __m256i _pB1 = _mm256_alignr_epi8(_pB0, _pB0, 4); _sum0 = _mm256_comp_dpbusd_epi32(_sum0, _pB0, _pA0); _sum1 = _mm256_comp_dpbusd_epi32(_sum1, _pB1, _pA0); _sum2 = _mm256_comp_dpbusd_epi32(_sum2, _pB0, _pA1); @@ -3290,8 +3030,8 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } if (max_kk >= 4) { - const __m256i _shift0 = _mm256_broadcastsi128_si256(_mm_loadu_si128((const __m128i*)pA)); - const __m256i _shift1 = _mm256_alignr_epi8(_shift0, _shift0, 8); + __m256i _shift0 = _mm256_broadcastsi128_si256(_mm_loadu_si128((const __m128i*)pA)); + __m256i _shift1 = _mm256_alignr_epi8(_shift0, _shift0, 8); _sum0 = _mm256_sub_epi32(_sum0, _shift0); _sum1 = _mm256_sub_epi32(_sum1, _shift0); _sum2 = _mm256_sub_epi32(_sum2, _shift1); @@ -3301,13 +3041,13 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #endif // __AVX512VNNI__ for (; kk + 1 < max_kk; kk += 2) { - const __m128i _pA8x1 = _mm_loadl_epi64((const __m128i*)pA); - const __m128i _pA8 = _mm_unpacklo_epi64(_pA8x1, _pA8x1); - const __m128i _pB8 = _mm_loadu_si128((const __m128i*)pB); - const __m256i _pA0 = _mm256_cvtepi8_epi16(_pA8); - const __m256i _pA1 = _mm256_alignr_epi8(_pA0, _pA0, 8); - const __m256i _pB0 = _mm256_cvtepi8_epi16(_pB8); - const __m256i _pB1 = _mm256_alignr_epi8(_pB0, _pB0, 4); + __m128i _pA8x1 = _mm_loadl_epi64((const __m128i*)pA); + __m128i _pA8 = _mm_unpacklo_epi64(_pA8x1, _pA8x1); + __m128i _pB8 = _mm_loadu_si128((const __m128i*)pB); + __m256i _pA0 = _mm256_cvtepi8_epi16(_pA8); + __m256i _pA1 = _mm256_alignr_epi8(_pA0, _pA0, 8); + __m256i _pB0 = _mm256_cvtepi8_epi16(_pB8); + __m256i _pB1 = _mm256_alignr_epi8(_pB0, _pB0, 4); _sum0 = _mm256_comp_dpwssd_epi32(_sum0, _pA0, _pB0); _sum1 = _mm256_comp_dpwssd_epi32(_sum1, _pA0, _pB1); _sum2 = _mm256_comp_dpwssd_epi32(_sum2, _pA1, _pB0); @@ -3317,13 +3057,13 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } for (; kk < max_kk; kk++) { - const __m128i _pA8 = _mm_cvtsi32_si128(*(const int*)pA); - const __m128i _pB8 = _mm_loadl_epi64((const __m128i*)pB); - const __m128i _pA32 = _mm_cvtepi8_epi32(_pA8); - const __m256i _pA0 = combine4x2_epi32(_pA32, _pA32); - const __m256i _pA1 = _mm256_shuffle_epi32(_pA0, _MM_SHUFFLE(1, 0, 3, 2)); - const __m256i _pB0 = combine4x2_epi32(_mm_cvtepi8_epi32(_pB8), _mm_cvtepi8_epi32(_mm_srli_si128(_pB8, 4))); - const __m256i _pB1 = _mm256_shuffle_epi32(_pB0, _MM_SHUFFLE(0, 3, 2, 1)); + __m128i _pA8 = _mm_loadu_si32(pA); + __m128i _pB8 = _mm_loadl_epi64((const __m128i*)pB); + __m128i _pA32 = _mm_cvtepi8_epi32(_pA8); + __m256i _pA0 = combine4x2_epi32(_pA32, _pA32); + __m256i _pA1 = _mm256_shuffle_epi32(_pA0, _MM_SHUFFLE(1, 0, 3, 2)); + __m256i _pB0 = combine4x2_epi32(_mm_cvtepi8_epi32(_pB8), _mm_cvtepi8_epi32(_mm_srli_si128(_pB8, 4))); + __m256i _pB1 = _mm256_shuffle_epi32(_pB0, _MM_SHUFFLE(0, 3, 2, 1)); _sum0 = _mm256_add_epi32(_sum0, _mm256_mullo_epi32(_pA0, _pB0)); _sum1 = _mm256_add_epi32(_sum1, _mm256_mullo_epi32(_pA0, _pB1)); _sum2 = _mm256_add_epi32(_sum2, _mm256_mullo_epi32(_pA1, _pB0)); @@ -3332,11 +3072,11 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB += 8; } - const __m128 _ad128 = _mm_loadu_ps(pA_descales); - const __m256 _ad0 = combine4x2_ps(_ad128, _ad128); - const __m256 _ad1 = _mm256_shuffle_ps(_ad0, _ad0, _MM_SHUFFLE(1, 0, 3, 2)); - const __m256 _bd0 = _mm256_loadu_ps(pB_descales); - const __m256 _bd1 = _mm256_shuffle_ps(_bd0, _bd0, _MM_SHUFFLE(0, 3, 2, 1)); + __m128 _ad128 = _mm_loadu_ps(pA_descales); + __m256 _ad0 = combine4x2_ps(_ad128, _ad128); + __m256 _ad1 = _mm256_shuffle_ps(_ad0, _ad0, _MM_SHUFFLE(1, 0, 3, 2)); + __m256 _bd0 = _mm256_loadu_ps(pB_descales); + __m256 _bd1 = _mm256_shuffle_ps(_bd0, _bd0, _MM_SHUFFLE(0, 3, 2, 1)); _fsum0 = _mm256_add_ps(_fsum0, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum0), _mm256_mul_ps(_ad0, _bd0))); _fsum1 = _mm256_add_ps(_fsum1, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum1), _mm256_mul_ps(_ad0, _bd1))); _fsum2 = _mm256_add_ps(_fsum2, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum2), _mm256_mul_ps(_ad1, _bd0))); @@ -3345,6 +3085,9 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB_descales += 8; } + pB += remain_K * 8; + pB_descales += remain_blocks * 8; + _mm256_storeu_ps(outptr, _fsum0); _mm256_storeu_ps(outptr + 8, _fsum1); _mm256_storeu_ps(outptr + 16, _fsum2); @@ -3354,28 +3097,30 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #endif // __AVX512F__ for (; jj + 3 < max_jj; jj += 4) { - __m128 _fsum0 = _mm_setzero_ps(); - __m128 _fsum1 = _mm_setzero_ps(); - __m128 _fsum2 = _mm_setzero_ps(); - __m128 _fsum3 = _mm_setzero_ps(); + pB += k * 4; + pB_descales += block_start * 4; + __m128 _fsum0 = k == 0 ? _mm_setzero_ps() : _mm_loadu_ps(outptr); + __m128 _fsum1 = k == 0 ? _mm_setzero_ps() : _mm_loadu_ps(outptr + 4); + __m128 _fsum2 = k == 0 ? _mm_setzero_ps() : _mm_loadu_ps(outptr + 8); + __m128 _fsum3 = k == 0 ? _mm_setzero_ps() : _mm_loadu_ps(outptr + 12); const signed char* pA = pAT; const float* pA_descales = pAT_descales; - for (int k = 0; k < K; k += block_size) + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) { __m128i _sum0 = _mm_setzero_si128(); __m128i _sum1 = _mm_setzero_si128(); __m128i _sum2 = _mm_setzero_si128(); __m128i _sum3 = _mm_setzero_si128(); - const int max_kk = std::min(K - k, block_size); + const int max_kk = std::min(tile_K - kk0, block_size); int kk = 0; #if __AVX512VNNI__ || __AVXVNNI__ for (; kk + 3 < max_kk; kk += 4) { - const __m128i _pA0 = _mm_loadu_si128((const __m128i*)pA); - const __m128i _pA1 = _mm_alignr_epi8(_pA0, _pA0, 8); - const __m128i _pB0 = _mm_loadu_si128((const __m128i*)pB); - const __m128i _pB1 = _mm_alignr_epi8(_pB0, _pB0, 4); + __m128i _pA0 = _mm_loadu_si128((const __m128i*)pA); + __m128i _pA1 = _mm_alignr_epi8(_pA0, _pA0, 8); + __m128i _pB0 = _mm_loadu_si128((const __m128i*)pB); + __m128i _pB1 = _mm_alignr_epi8(_pB0, _pB0, 4); #if __AVXVNNIINT8__ _sum0 = _mm_dpbssd_epi32(_sum0, _pB0, _pA0); _sum1 = _mm_dpbssd_epi32(_sum1, _pB1, _pA0); @@ -3393,8 +3138,8 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) if (max_kk >= 4) { - const __m128i _shift0 = _mm_loadu_si128((const __m128i*)pA); - const __m128i _shift1 = _mm_alignr_epi8(_shift0, _shift0, 8); + __m128i _shift0 = _mm_loadu_si128((const __m128i*)pA); + __m128i _shift1 = _mm_alignr_epi8(_shift0, _shift0, 8); _sum0 = _mm_sub_epi32(_sum0, _shift0); _sum1 = _mm_sub_epi32(_sum1, _shift0); _sum2 = _mm_sub_epi32(_sum2, _shift1); @@ -3405,12 +3150,12 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #endif // __AVX512VNNI__ || __AVXVNNI__ for (; kk + 1 < max_kk; kk += 2) { - const __m128i _pA8 = _mm_loadl_epi64((const __m128i*)pA); - const __m128i _pB8 = _mm_loadl_epi64((const __m128i*)pB); - const __m128i _pA0 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); - const __m128i _pA1 = _mm_shuffle_epi32(_pA0, _MM_SHUFFLE(1, 0, 3, 2)); - const __m128i _pB0 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); - const __m128i _pB1 = _mm_shuffle_epi32(_pB0, _MM_SHUFFLE(0, 3, 2, 1)); + __m128i _pA8 = _mm_loadl_epi64((const __m128i*)pA); + __m128i _pB8 = _mm_loadl_epi64((const __m128i*)pB); + __m128i _pA0 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + __m128i _pA1 = _mm_shuffle_epi32(_pA0, _MM_SHUFFLE(1, 0, 3, 2)); + __m128i _pB0 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); + __m128i _pB1 = _mm_shuffle_epi32(_pB0, _MM_SHUFFLE(0, 3, 2, 1)); _sum0 = _mm_comp_dpwssd_epi32(_sum0, _pA0, _pB0); _sum1 = _mm_comp_dpwssd_epi32(_sum1, _pA0, _pB1); _sum2 = _mm_comp_dpwssd_epi32(_sum2, _pA1, _pB0); @@ -3420,14 +3165,14 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } for (; kk < max_kk; kk++) { - const __m128i _pA8 = _mm_cvtsi32_si128(*(const int*)pA); - const __m128i _pB8 = _mm_cvtsi32_si128(*(const int*)pB); - const __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); - const __m128i _pB16 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); - const __m128i _pA0 = _mm_unpacklo_epi16(_pA16, _pA16); - const __m128i _pA1 = _mm_shuffle_epi32(_pA0, _MM_SHUFFLE(1, 0, 3, 2)); - const __m128i _pB0 = _mm_unpacklo_epi16(_pB16, _mm_setzero_si128()); - const __m128i _pB1 = _mm_shuffle_epi32(_pB0, _MM_SHUFFLE(0, 3, 2, 1)); + __m128i _pA8 = _mm_loadu_si32(pA); + __m128i _pB8 = _mm_loadu_si32(pB); + __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + __m128i _pB16 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); + __m128i _pA0 = _mm_unpacklo_epi16(_pA16, _pA16); + __m128i _pA1 = _mm_shuffle_epi32(_pA0, _MM_SHUFFLE(1, 0, 3, 2)); + __m128i _pB0 = _mm_unpacklo_epi16(_pB16, _mm_setzero_si128()); + __m128i _pB1 = _mm_shuffle_epi32(_pB0, _MM_SHUFFLE(0, 3, 2, 1)); _sum0 = _mm_comp_dpwssd_epi32(_sum0, _pA0, _pB0); _sum1 = _mm_comp_dpwssd_epi32(_sum1, _pA0, _pB1); _sum2 = _mm_comp_dpwssd_epi32(_sum2, _pA1, _pB0); @@ -3436,10 +3181,10 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB += 4; } - const __m128 _ad0 = _mm_loadu_ps(pA_descales); - const __m128 _ad1 = _mm_shuffle_ps(_ad0, _ad0, _MM_SHUFFLE(1, 0, 3, 2)); - const __m128 _bd0 = _mm_loadu_ps(pB_descales); - const __m128 _bd1 = _mm_shuffle_ps(_bd0, _bd0, _MM_SHUFFLE(0, 3, 2, 1)); + __m128 _ad0 = _mm_loadu_ps(pA_descales); + __m128 _ad1 = _mm_shuffle_ps(_ad0, _ad0, _MM_SHUFFLE(1, 0, 3, 2)); + __m128 _bd0 = _mm_loadu_ps(pB_descales); + __m128 _bd1 = _mm_shuffle_ps(_bd0, _bd0, _MM_SHUFFLE(0, 3, 2, 1)); _fsum0 = _mm_add_ps(_fsum0, _mm_mul_ps(_mm_cvtepi32_ps(_sum0), _mm_mul_ps(_ad0, _bd0))); _fsum1 = _mm_add_ps(_fsum1, _mm_mul_ps(_mm_cvtepi32_ps(_sum1), _mm_mul_ps(_ad0, _bd1))); _fsum2 = _mm_add_ps(_fsum2, _mm_mul_ps(_mm_cvtepi32_ps(_sum2), _mm_mul_ps(_ad1, _bd0))); @@ -3448,6 +3193,9 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB_descales += 4; } + pB += remain_K * 4; + pB_descales += remain_blocks * 4; + _mm_storeu_ps(outptr, _fsum0); _mm_storeu_ps(outptr + 4, _fsum1); _mm_storeu_ps(outptr + 8, _fsum2); @@ -3457,24 +3205,26 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #endif // defined(__x86_64__) || defined(_M_X64) for (; jj + 1 < max_jj; jj += 2) { - __m128 _fsum0 = _mm_setzero_ps(); - __m128 _fsum1 = _mm_setzero_ps(); + pB += k * 2; + pB_descales += block_start * 2; + __m128 _fsum0 = k == 0 ? _mm_setzero_ps() : _mm_loadu_ps(outptr); + __m128 _fsum1 = k == 0 ? _mm_setzero_ps() : _mm_loadu_ps(outptr + 4); const signed char* pA = pAT; const float* pA_descales = pAT_descales; - for (int k = 0; k < K; k += block_size) + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) { __m128i _sum0 = _mm_setzero_si128(); __m128i _sum1 = _mm_setzero_si128(); - const int max_kk = std::min(K - k, block_size); + const int max_kk = std::min(tile_K - kk0, block_size); int kk = 0; #if __AVX512VNNI__ || __AVXVNNI__ for (; kk + 3 < max_kk; kk += 4) { - const __m128i _pA0 = _mm_loadu_si128((const __m128i*)pA); - const __m128i _pB8 = _mm_loadl_epi64((const __m128i*)pB); - const __m128i _pB0 = _mm_unpacklo_epi64(_pB8, _pB8); - const __m128i _pB1 = _mm_alignr_epi8(_pB0, _pB0, 4); + __m128i _pA0 = _mm_loadu_si128((const __m128i*)pA); + __m128i _pB8 = _mm_loadl_epi64((const __m128i*)pB); + __m128i _pB0 = _mm_unpacklo_epi64(_pB8, _pB8); + __m128i _pB1 = _mm_alignr_epi8(_pB0, _pB0, 4); #if __AVXVNNIINT8__ _sum0 = _mm_dpbssd_epi32(_sum0, _pB0, _pA0); _sum1 = _mm_dpbssd_epi32(_sum1, _pB1, _pA0); @@ -3488,7 +3238,7 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) if (max_kk >= 4) { - const __m128i _shift0 = _mm_loadu_si128((const __m128i*)pA); + __m128i _shift0 = _mm_loadu_si128((const __m128i*)pA); _sum0 = _mm_sub_epi32(_sum0, _shift0); _sum1 = _mm_sub_epi32(_sum1, _shift0); pA += 16; @@ -3497,11 +3247,11 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #endif // __AVX512VNNI__ || __AVXVNNI__ for (; kk + 1 < max_kk; kk += 2) { - const __m128i _pA8 = _mm_loadl_epi64((const __m128i*)pA); - const __m128i _pB8 = _mm_castps_si128(_mm_load1_ps((const float*)pB)); - const __m128i _pA0 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); - const __m128i _pB0 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); - const __m128i _pB1 = _mm_shuffle_epi32(_pB0, _MM_SHUFFLE(0, 3, 2, 1)); + __m128i _pA8 = _mm_loadl_epi64((const __m128i*)pA); + __m128i _pB8 = _mm_castps_si128(_mm_load1_ps((const float*)pB)); + __m128i _pA0 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + __m128i _pB0 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); + __m128i _pB1 = _mm_shuffle_epi32(_pB0, _MM_SHUFFLE(0, 3, 2, 1)); _sum0 = _mm_comp_dpwssd_epi32(_sum0, _pA0, _pB0); _sum1 = _mm_comp_dpwssd_epi32(_sum1, _pA0, _pB1); pA += 8; @@ -3509,48 +3259,53 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } for (; kk < max_kk; kk++) { - const __m128i _pA8 = _mm_cvtsi32_si128(*(const int*)pA); - const __m128i _pB8 = _mm_set1_epi16(*(const short*)pB); - const __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); - const __m128i _pB16 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); - const __m128i _pA0 = _mm_unpacklo_epi16(_pA16, _pA16); - const __m128i _pB0 = _mm_unpacklo_epi16(_pB16, _mm_setzero_si128()); - const __m128i _pB1 = _mm_shuffle_epi32(_pB0, _MM_SHUFFLE(0, 3, 2, 1)); + __m128i _pA8 = _mm_loadu_si32(pA); + __m128i _pB8 = _mm_set1_epi16((short)_mm_cvtsi128_si32(_mm_loadu_si16(pB))); + __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + __m128i _pB16 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); + __m128i _pA0 = _mm_unpacklo_epi16(_pA16, _pA16); + __m128i _pB0 = _mm_unpacklo_epi16(_pB16, _mm_setzero_si128()); + __m128i _pB1 = _mm_shuffle_epi32(_pB0, _MM_SHUFFLE(0, 3, 2, 1)); _sum0 = _mm_comp_dpwssd_epi32(_sum0, _pA0, _pB0); _sum1 = _mm_comp_dpwssd_epi32(_sum1, _pA0, _pB1); pA += 4; pB += 2; } - const __m128 _ad = _mm_loadu_ps(pA_descales); - const __m128 _bd0 = _mm_setr_ps(pB_descales[0], pB_descales[1], pB_descales[0], pB_descales[1]); - const __m128 _bd1 = _mm_shuffle_ps(_bd0, _bd0, _MM_SHUFFLE(0, 3, 2, 1)); + __m128 _ad = _mm_loadu_ps(pA_descales); + __m128 _bd0 = _mm_setr_ps(pB_descales[0], pB_descales[1], pB_descales[0], pB_descales[1]); + __m128 _bd1 = _mm_shuffle_ps(_bd0, _bd0, _MM_SHUFFLE(0, 3, 2, 1)); _fsum0 = _mm_add_ps(_fsum0, _mm_mul_ps(_mm_cvtepi32_ps(_sum0), _mm_mul_ps(_ad, _bd0))); _fsum1 = _mm_add_ps(_fsum1, _mm_mul_ps(_mm_cvtepi32_ps(_sum1), _mm_mul_ps(_ad, _bd1))); pA_descales += 4; pB_descales += 2; } + pB += remain_K * 2; + pB_descales += remain_blocks * 2; + _mm_storeu_ps(outptr, _fsum0); _mm_storeu_ps(outptr + 4, _fsum1); outptr += 8; } for (; jj < max_jj; jj++) { - __m128 _fsum = _mm_setzero_ps(); + pB += k; + pB_descales += block_start; + __m128 _fsum = k == 0 ? _mm_setzero_ps() : _mm_loadu_ps(outptr); const signed char* pA = pAT; const float* pA_descales = pAT_descales; - for (int k = 0; k < K; k += block_size) + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) { __m128i _sum = _mm_setzero_si128(); - const int max_kk = std::min(K - k, block_size); + const int max_kk = std::min(tile_K - kk0, block_size); int kk = 0; #if __AVX512VNNI__ || __AVXVNNI__ for (; kk + 3 < max_kk; kk += 4) { - const __m128i _pA = _mm_loadu_si128((const __m128i*)pA); - const __m128i _pB = _mm_set1_epi32(*(const int*)pB); + __m128i _pA = _mm_loadu_si128((const __m128i*)pA); + __m128i _pB = _mm_set1_epi32(_mm_cvtsi128_si32(_mm_loadu_si32(pB))); #if __AVXVNNIINT8__ _sum = _mm_dpbssd_epi32(_sum, _pB, _pA); #else // __AVXVNNIINT8__ @@ -3569,30 +3324,33 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #endif // __AVX512VNNI__ || __AVXVNNI__ for (; kk + 1 < max_kk; kk += 2) { - const __m128i _pA8 = _mm_loadl_epi64((const __m128i*)pA); - const __m128i _pB8 = _mm_set1_epi16(*(const short*)pB); - const __m128i _pA = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); - const __m128i _pB = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); + __m128i _pA8 = _mm_loadl_epi64((const __m128i*)pA); + __m128i _pB8 = _mm_set1_epi16((short)_mm_cvtsi128_si32(_mm_loadu_si16(pB))); + __m128i _pA = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + __m128i _pB = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); _sum = _mm_comp_dpwssd_epi32(_sum, _pA, _pB); pA += 8; pB += 2; } for (; kk < max_kk; kk++) { - const __m128i _pA8 = _mm_cvtsi32_si128(*(const int*)pA); - const __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); - const __m128i _pA = _mm_unpacklo_epi16(_pA16, _mm_setzero_si128()); + __m128i _pA8 = _mm_loadu_si32(pA); + __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + __m128i _pA = _mm_unpacklo_epi16(_pA16, _mm_setzero_si128()); _sum = _mm_comp_dpwssd_epi32(_sum, _pA, _mm_set1_epi16((signed char)pB[0])); pA += 4; pB++; } - const __m128 _ad = _mm_loadu_ps(pA_descales); + __m128 _ad = _mm_loadu_ps(pA_descales); _fsum = _mm_add_ps(_fsum, _mm_mul_ps(_mm_cvtepi32_ps(_sum), _mm_mul_ps(_ad, _mm_set1_ps(pB_descales[0])))); pA_descales += 4; pB_descales++; } + pB += remain_K * 1; + pB_descales += remain_blocks * 1; + _mm_storeu_ps(outptr, _fsum); outptr += 4; } @@ -3612,25 +3370,27 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #if __AVX512F__ for (; jj + 7 < max_jj; jj += 8) { - __m256 _fsum0 = _mm256_setzero_ps(); - __m256 _fsum1 = _mm256_setzero_ps(); + pB += k * 8; + pB_descales += block_start * 8; + __m256 _fsum0 = k == 0 ? _mm256_setzero_ps() : _mm256_loadu_ps(outptr); + __m256 _fsum1 = k == 0 ? _mm256_setzero_ps() : _mm256_loadu_ps(outptr + 8); const signed char* pA = pAT; const float* pA_descales = pAT_descales; - for (int k = 0; k < K; k += block_size) + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) { __m256i _sum0 = _mm256_setzero_si256(); __m256i _sum1 = _mm256_setzero_si256(); - const int max_kk = std::min(K - k, block_size); + const int max_kk = std::min(tile_K - kk0, block_size); int kk = 0; #if __AVX512VNNI__ for (; kk + 3 < max_kk; kk += 4) { - const __m128i _pA8 = _mm_loadl_epi64((const __m128i*)pA); - const __m128i _pA128 = _mm_unpacklo_epi64(_pA8, _pA8); - const __m256i _pA0 = _mm256_broadcastsi128_si256(_pA128); - const __m256i _pA1 = _mm256_shuffle_epi32(_pA0, _MM_SHUFFLE(2, 3, 0, 1)); - const __m256i _pB = _mm256_loadu_si256((const __m256i*)pB); + __m128i _pA8 = _mm_loadl_epi64((const __m128i*)pA); + __m128i _pA128 = _mm_unpacklo_epi64(_pA8, _pA8); + __m256i _pA0 = _mm256_broadcastsi128_si256(_pA128); + __m256i _pA1 = _mm256_shuffle_epi32(_pA0, _MM_SHUFFLE(2, 3, 0, 1)); + __m256i _pB = _mm256_loadu_si256((const __m256i*)pB); _sum0 = _mm256_comp_dpbusd_epi32(_sum0, _pB, _pA0); _sum1 = _mm256_comp_dpbusd_epi32(_sum1, _pB, _pA1); pA += 8; @@ -3638,10 +3398,10 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } if (max_kk >= 4) { - const __m128i _shift64 = _mm_loadl_epi64((const __m128i*)pA); - const __m128i _shift128 = _mm_unpacklo_epi64(_shift64, _shift64); - const __m256i _shift0 = _mm256_broadcastsi128_si256(_shift128); - const __m256i _shift1 = _mm256_shuffle_epi32(_shift0, _MM_SHUFFLE(2, 3, 0, 1)); + __m128i _shift64 = _mm_loadl_epi64((const __m128i*)pA); + __m128i _shift128 = _mm_unpacklo_epi64(_shift64, _shift64); + __m256i _shift0 = _mm256_broadcastsi128_si256(_shift128); + __m256i _shift1 = _mm256_shuffle_epi32(_shift0, _MM_SHUFFLE(2, 3, 0, 1)); _sum0 = _mm256_sub_epi32(_sum0, _shift0); _sum1 = _mm256_sub_epi32(_sum1, _shift1); pA += 8; @@ -3649,12 +3409,12 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #endif // __AVX512VNNI__ for (; kk + 1 < max_kk; kk += 2) { - const __m128i _pA8 = _mm_cvtsi32_si128(*(const int*)pA); - const __m128i _pA16x1 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); - const __m128i _pA16 = _mm_unpacklo_epi64(_pA16x1, _pA16x1); - const __m256i _pA0 = _mm256_broadcastsi128_si256(_pA16); - const __m256i _pA1 = _mm256_shuffle_epi32(_pA0, _MM_SHUFFLE(2, 3, 0, 1)); - const __m256i _pB = _mm256_cvtepi8_epi16(_mm_loadu_si128((const __m128i*)pB)); + __m128i _pA8 = _mm_loadu_si32(pA); + __m128i _pA16x1 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + __m128i _pA16 = _mm_unpacklo_epi64(_pA16x1, _pA16x1); + __m256i _pA0 = _mm256_broadcastsi128_si256(_pA16); + __m256i _pA1 = _mm256_shuffle_epi32(_pA0, _MM_SHUFFLE(2, 3, 0, 1)); + __m256i _pB = _mm256_cvtepi8_epi16(_mm_loadu_si128((const __m128i*)pB)); _sum0 = _mm256_comp_dpwssd_epi32(_sum0, _pA0, _pB); _sum1 = _mm256_comp_dpwssd_epi32(_sum1, _pA1, _pB); pA += 4; @@ -3662,29 +3422,32 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } for (; kk < max_kk; kk++) { - const __m128i _pA8 = _mm_cvtsi32_si128(*(const unsigned short*)pA); - const __m128i _pA32x1 = _mm_cvtepi8_epi32(_pA8); - const __m128i _pA128 = _mm_shuffle_epi32(_pA32x1, _MM_SHUFFLE(1, 0, 1, 0)); - const __m256i _pA0 = _mm256_broadcastsi128_si256(_pA128); - const __m256i _pA1 = _mm256_shuffle_epi32(_pA0, _MM_SHUFFLE(2, 3, 0, 1)); - const __m256i _pB = _mm256_cvtepi8_epi32(_mm_loadl_epi64((const __m128i*)pB)); + __m128i _pA8 = _mm_loadu_si16(pA); + __m128i _pA32x1 = _mm_cvtepi8_epi32(_pA8); + __m128i _pA128 = _mm_shuffle_epi32(_pA32x1, _MM_SHUFFLE(1, 0, 1, 0)); + __m256i _pA0 = _mm256_broadcastsi128_si256(_pA128); + __m256i _pA1 = _mm256_shuffle_epi32(_pA0, _MM_SHUFFLE(2, 3, 0, 1)); + __m256i _pB = _mm256_cvtepi8_epi32(_mm_loadl_epi64((const __m128i*)pB)); _sum0 = _mm256_add_epi32(_sum0, _mm256_mullo_epi32(_pA0, _pB)); _sum1 = _mm256_add_epi32(_sum1, _mm256_mullo_epi32(_pA1, _pB)); pA += 2; pB += 8; } - const __m128 _ad2 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pA_descales); - const __m128 _ad128 = _mm_movelh_ps(_ad2, _ad2); - const __m256 _ad0 = combine4x2_ps(_ad128, _ad128); - const __m256 _ad1 = _mm256_shuffle_ps(_ad0, _ad0, _MM_SHUFFLE(2, 3, 0, 1)); - const __m256 _bd = _mm256_loadu_ps(pB_descales); + __m128 _ad2 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pA_descales); + __m128 _ad128 = _mm_movelh_ps(_ad2, _ad2); + __m256 _ad0 = combine4x2_ps(_ad128, _ad128); + __m256 _ad1 = _mm256_shuffle_ps(_ad0, _ad0, _MM_SHUFFLE(2, 3, 0, 1)); + __m256 _bd = _mm256_loadu_ps(pB_descales); _fsum0 = _mm256_add_ps(_fsum0, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum0), _mm256_mul_ps(_ad0, _bd))); _fsum1 = _mm256_add_ps(_fsum1, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum1), _mm256_mul_ps(_ad1, _bd))); pA_descales += 2; pB_descales += 8; } + pB += remain_K * 8; + pB_descales += remain_blocks * 8; + _mm256_storeu_ps(outptr, _fsum0); _mm256_storeu_ps(outptr + 8, _fsum1); outptr += 16; @@ -3692,24 +3455,26 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #endif // __AVX512F__ for (; jj + 3 < max_jj; jj += 4) { - __m128 _fsum0 = _mm_setzero_ps(); - __m128 _fsum1 = _mm_setzero_ps(); + pB += k * 4; + pB_descales += block_start * 4; + __m128 _fsum0 = k == 0 ? _mm_setzero_ps() : _mm_loadu_ps(outptr); + __m128 _fsum1 = k == 0 ? _mm_setzero_ps() : _mm_loadu_ps(outptr + 4); const signed char* pA = pAT; const float* pA_descales = pAT_descales; - for (int k = 0; k < K; k += block_size) + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) { __m128i _sum0 = _mm_setzero_si128(); __m128i _sum1 = _mm_setzero_si128(); - const int max_kk = std::min(K - k, block_size); + const int max_kk = std::min(tile_K - kk0, block_size); int kk = 0; #if __AVX512VNNI__ || __AVXVNNI__ for (; kk + 3 < max_kk; kk += 4) { - const __m128i _pA8 = _mm_loadl_epi64((const __m128i*)pA); - const __m128i _pA = _mm_unpacklo_epi64(_pA8, _pA8); - const __m128i _pB0 = _mm_loadu_si128((const __m128i*)pB); - const __m128i _pB1 = _mm_shuffle_epi32(_pB0, _MM_SHUFFLE(0, 3, 2, 1)); + __m128i _pA8 = _mm_loadl_epi64((const __m128i*)pA); + __m128i _pA = _mm_unpacklo_epi64(_pA8, _pA8); + __m128i _pB0 = _mm_loadu_si128((const __m128i*)pB); + __m128i _pB1 = _mm_shuffle_epi32(_pB0, _MM_SHUFFLE(0, 3, 2, 1)); #if __AVXVNNIINT8__ _sum0 = _mm_dpbssd_epi32(_sum0, _pB0, _pA); _sum1 = _mm_dpbssd_epi32(_sum1, _pB1, _pA); @@ -3723,8 +3488,8 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) if (max_kk >= 4) { - const __m128i _shift64 = _mm_loadl_epi64((const __m128i*)pA); - const __m128i _shift = _mm_unpacklo_epi64(_shift64, _shift64); + __m128i _shift64 = _mm_loadl_epi64((const __m128i*)pA); + __m128i _shift = _mm_unpacklo_epi64(_shift64, _shift64); _sum0 = _mm_sub_epi32(_sum0, _shift); _sum1 = _mm_sub_epi32(_sum1, _shift); pA += 8; @@ -3733,12 +3498,12 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #endif // __AVX512VNNI__ || __AVXVNNI__ for (; kk + 1 < max_kk; kk += 2) { - const __m128i _pA8 = _mm_cvtsi32_si128(*(const int*)pA); - const __m128i _pA16x1 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); - const __m128i _pA = _mm_unpacklo_epi64(_pA16x1, _pA16x1); - const __m128i _pB8 = _mm_loadl_epi64((const __m128i*)pB); - const __m128i _pB0 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); - const __m128i _pB1 = _mm_shuffle_epi32(_pB0, _MM_SHUFFLE(0, 3, 2, 1)); + __m128i _pA8 = _mm_loadu_si32(pA); + __m128i _pA16x1 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + __m128i _pA = _mm_unpacklo_epi64(_pA16x1, _pA16x1); + __m128i _pB8 = _mm_loadl_epi64((const __m128i*)pB); + __m128i _pB0 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); + __m128i _pB1 = _mm_shuffle_epi32(_pB0, _MM_SHUFFLE(0, 3, 2, 1)); _sum0 = _mm_comp_dpwssd_epi32(_sum0, _pA, _pB0); _sum1 = _mm_comp_dpwssd_epi32(_sum1, _pA, _pB1); pA += 4; @@ -3746,30 +3511,33 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } for (; kk < max_kk; kk++) { - const __m128i _pA8 = _mm_cvtsi32_si128(*(const unsigned short*)pA); - const __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); - const __m128i _pA32x1 = _mm_unpacklo_epi16(_pA16, _pA16); - const __m128i _pA = _mm_shuffle_epi32(_pA32x1, _MM_SHUFFLE(1, 0, 1, 0)); - const __m128i _pB8 = _mm_cvtsi32_si128(*(const int*)pB); - const __m128i _pB16 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); - const __m128i _pB0 = _mm_unpacklo_epi16(_pB16, _mm_setzero_si128()); - const __m128i _pB1 = _mm_shuffle_epi32(_pB0, _MM_SHUFFLE(0, 3, 2, 1)); + __m128i _pA8 = _mm_loadu_si16(pA); + __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + __m128i _pA32x1 = _mm_unpacklo_epi16(_pA16, _pA16); + __m128i _pA = _mm_shuffle_epi32(_pA32x1, _MM_SHUFFLE(1, 0, 1, 0)); + __m128i _pB8 = _mm_loadu_si32(pB); + __m128i _pB16 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); + __m128i _pB0 = _mm_unpacklo_epi16(_pB16, _mm_setzero_si128()); + __m128i _pB1 = _mm_shuffle_epi32(_pB0, _MM_SHUFFLE(0, 3, 2, 1)); _sum0 = _mm_comp_dpwssd_epi32(_sum0, _pA, _pB0); _sum1 = _mm_comp_dpwssd_epi32(_sum1, _pA, _pB1); pA += 2; pB += 4; } - const __m128 _ad2 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pA_descales); - const __m128 _ad = _mm_movelh_ps(_ad2, _ad2); - const __m128 _bd0 = _mm_loadu_ps(pB_descales); - const __m128 _bd1 = _mm_shuffle_ps(_bd0, _bd0, _MM_SHUFFLE(0, 3, 2, 1)); + __m128 _ad2 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pA_descales); + __m128 _ad = _mm_movelh_ps(_ad2, _ad2); + __m128 _bd0 = _mm_loadu_ps(pB_descales); + __m128 _bd1 = _mm_shuffle_ps(_bd0, _bd0, _MM_SHUFFLE(0, 3, 2, 1)); _fsum0 = _mm_add_ps(_fsum0, _mm_mul_ps(_mm_cvtepi32_ps(_sum0), _mm_mul_ps(_ad, _bd0))); _fsum1 = _mm_add_ps(_fsum1, _mm_mul_ps(_mm_cvtepi32_ps(_sum1), _mm_mul_ps(_ad, _bd1))); pA_descales += 2; pB_descales += 4; } + pB += remain_K * 4; + pB_descales += remain_blocks * 4; + _mm_storeu_ps(outptr, _fsum0); _mm_storeu_ps(outptr + 4, _fsum1); outptr += 8; @@ -3778,23 +3546,25 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #endif // __SSE2__ for (; jj + 1 < max_jj; jj += 2) { + pB += k * 2; + pB_descales += block_start * 2; #if __SSE2__ - __m128 _fsum = _mm_setzero_ps(); + __m128 _fsum = k == 0 ? _mm_setzero_ps() : _mm_loadu_ps(outptr); const signed char* pA = pAT; const float* pA_descales = pAT_descales; - for (int k = 0; k < K; k += block_size) + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) { __m128i _sum = _mm_setzero_si128(); - const int max_kk = std::min(K - k, block_size); + const int max_kk = std::min(tile_K - kk0, block_size); int kk = 0; #if __AVX512VNNI__ || __AVXVNNI__ for (; kk + 3 < max_kk; kk += 4) { - const __m128i _pA8 = _mm_loadl_epi64((const __m128i*)pA); - const __m128i _pA = _mm_unpacklo_epi32(_pA8, _pA8); - const __m128i _pB8 = _mm_loadl_epi64((const __m128i*)pB); - const __m128i _pB = _mm_unpacklo_epi64(_pB8, _pB8); + __m128i _pA8 = _mm_loadl_epi64((const __m128i*)pA); + __m128i _pA = _mm_unpacklo_epi32(_pA8, _pA8); + __m128i _pB8 = _mm_loadl_epi64((const __m128i*)pB); + __m128i _pB = _mm_unpacklo_epi64(_pB8, _pB8); #if __AVXVNNIINT8__ _sum = _mm_dpbssd_epi32(_sum, _pB, _pA); #else // __AVXVNNIINT8__ @@ -3806,8 +3576,8 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) if (max_kk >= 4) { - const __m128i _shift64 = _mm_loadl_epi64((const __m128i*)pA); - const __m128i _shift = _mm_unpacklo_epi32(_shift64, _shift64); + __m128i _shift64 = _mm_loadl_epi64((const __m128i*)pA); + __m128i _shift = _mm_unpacklo_epi32(_shift64, _shift64); _sum = _mm_sub_epi32(_sum, _shift); pA += 8; } @@ -3815,57 +3585,60 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #endif // __AVX512VNNI__ || __AVXVNNI__ for (; kk + 1 < max_kk; kk += 2) { - const __m128i _pA8 = _mm_cvtsi32_si128(*(const int*)pA); - const __m128i _pA16x1 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); - const __m128i _pA = _mm_unpacklo_epi32(_pA16x1, _pA16x1); - const __m128i _pB8 = _mm_cvtsi32_si128(*(const int*)pB); - const __m128i _pB16x1 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); - const __m128i _pB = _mm_unpacklo_epi64(_pB16x1, _pB16x1); + __m128i _pA8 = _mm_loadu_si32(pA); + __m128i _pA16x1 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + __m128i _pA = _mm_unpacklo_epi32(_pA16x1, _pA16x1); + __m128i _pB8 = _mm_loadu_si32(pB); + __m128i _pB16x1 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); + __m128i _pB = _mm_unpacklo_epi64(_pB16x1, _pB16x1); _sum = _mm_comp_dpwssd_epi32(_sum, _pA, _pB); pA += 4; pB += 4; } for (; kk < max_kk; kk++) { - const __m128i _pA8 = _mm_cvtsi32_si128(*(const unsigned short*)pA); - const __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); - const __m128i _pA32x1 = _mm_unpacklo_epi16(_pA16, _pA16); - const __m128i _pA = _mm_unpacklo_epi32(_pA32x1, _pA32x1); - const __m128i _pB8 = _mm_cvtsi32_si128(*(const unsigned short*)pB); - const __m128i _pB16 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); - const __m128i _pB32x1 = _mm_unpacklo_epi16(_pB16, _mm_setzero_si128()); - const __m128i _pB = _mm_unpacklo_epi64(_pB32x1, _pB32x1); + __m128i _pA8 = _mm_loadu_si16(pA); + __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + __m128i _pA32x1 = _mm_unpacklo_epi16(_pA16, _pA16); + __m128i _pA = _mm_unpacklo_epi32(_pA32x1, _pA32x1); + __m128i _pB8 = _mm_loadu_si16(pB); + __m128i _pB16 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); + __m128i _pB32x1 = _mm_unpacklo_epi16(_pB16, _mm_setzero_si128()); + __m128i _pB = _mm_unpacklo_epi64(_pB32x1, _pB32x1); _sum = _mm_comp_dpwssd_epi32(_sum, _pA, _pB); pA += 2; pB += 2; } - const __m128 _ad2 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pA_descales); - const __m128 _ad = _mm_unpacklo_ps(_ad2, _ad2); - const __m128 _bd2 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pB_descales); - const __m128 _bd = _mm_movelh_ps(_bd2, _bd2); + __m128 _ad2 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pA_descales); + __m128 _ad = _mm_unpacklo_ps(_ad2, _ad2); + __m128 _bd2 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pB_descales); + __m128 _bd = _mm_movelh_ps(_bd2, _bd2); _fsum = _mm_add_ps(_fsum, _mm_mul_ps(_mm_cvtepi32_ps(_sum), _mm_mul_ps(_ad, _bd))); pA_descales += 2; pB_descales += 2; } + pB += remain_K * 2; + pB_descales += remain_blocks * 2; + _mm_storeu_ps(outptr, _fsum); outptr += 4; #else - float fsum00 = 0.f; - float fsum01 = 0.f; - float fsum10 = 0.f; - float fsum11 = 0.f; + float fsum00 = k == 0 ? 0.f : outptr[0]; + float fsum01 = k == 0 ? 0.f : outptr[1]; + float fsum10 = k == 0 ? 0.f : outptr[2]; + float fsum11 = k == 0 ? 0.f : outptr[3]; const signed char* pA = pAT; const float* pA_descales = pAT_descales; - for (int k = 0; k < K; k += block_size) + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) { int sum00 = 0; int sum01 = 0; int sum10 = 0; int sum11 = 0; - const int max_kk = std::min(K - k, block_size); + const int max_kk = std::min(tile_K - kk0, block_size); int kk = 0; for (; kk + 3 < max_kk; kk += 4) { @@ -3918,6 +3691,9 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB_descales += 2; } + pB += remain_K * 2; + pB_descales += remain_blocks * 2; + outptr[0] = fsum00; outptr[1] = fsum01; outptr[2] = fsum10; @@ -3927,22 +3703,24 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } for (; jj < max_jj; jj++) { + pB += k; + pB_descales += block_start; #if __SSE2__ - __m128 _fsum = _mm_setzero_ps(); + __m128 _fsum = k == 0 ? _mm_setzero_ps() : _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)outptr); const signed char* pA = pAT; const float* pA_descales = pAT_descales; - for (int k = 0; k < K; k += block_size) + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) { __m128i _sum = _mm_setzero_si128(); - const int max_kk = std::min(K - k, block_size); + const int max_kk = std::min(tile_K - kk0, block_size); int kk = 0; #if __AVX512VNNI__ || __AVXVNNI__ for (; kk + 3 < max_kk; kk += 4) { - const __m128i _pA = _mm_loadl_epi64((const __m128i*)pA); - const __m128i _pB8 = _mm_cvtsi32_si128(*(const int*)pB); - const __m128i _pB = _mm_shuffle_epi32(_pB8, _MM_SHUFFLE(0, 0, 0, 0)); + __m128i _pA = _mm_loadl_epi64((const __m128i*)pA); + __m128i _pB8 = _mm_loadu_si32(pB); + __m128i _pB = _mm_shuffle_epi32(_pB8, _MM_SHUFFLE(0, 0, 0, 0)); #if __AVXVNNIINT8__ _sum = _mm_dpbssd_epi32(_sum, _pB, _pA); #else // __AVXVNNIINT8__ @@ -3961,48 +3739,51 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #endif // __AVX512VNNI__ || __AVXVNNI__ for (; kk + 1 < max_kk; kk += 2) { - const __m128i _pA8 = _mm_cvtsi32_si128(*(const int*)pA); - const __m128i _pA = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); - const __m128i _pB8 = _mm_cvtsi32_si128(*(const unsigned short*)pB); - const __m128i _pB16 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); - const __m128i _pB = _mm_shuffle_epi32(_pB16, _MM_SHUFFLE(0, 0, 0, 0)); + __m128i _pA8 = _mm_loadu_si32(pA); + __m128i _pA = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + __m128i _pB8 = _mm_loadu_si16(pB); + __m128i _pB16 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); + __m128i _pB = _mm_shuffle_epi32(_pB16, _MM_SHUFFLE(0, 0, 0, 0)); _sum = _mm_comp_dpwssd_epi32(_sum, _pA, _pB); pA += 4; pB += 2; } for (; kk < max_kk; kk++) { - const __m128i _pA8 = _mm_cvtsi32_si128(*(const unsigned short*)pA); - const __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); - const __m128i _pA = _mm_unpacklo_epi16(_pA16, _pA16); - const __m128i _pB8 = _mm_cvtsi32_si128((signed char)pB[0]); - const __m128i _pB16 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); - const __m128i _pB32 = _mm_unpacklo_epi16(_pB16, _mm_setzero_si128()); - const __m128i _pB = _mm_shuffle_epi32(_pB32, _MM_SHUFFLE(0, 0, 0, 0)); + __m128i _pA8 = _mm_loadu_si16(pA); + __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + __m128i _pA = _mm_unpacklo_epi16(_pA16, _pA16); + __m128i _pB8 = _mm_cvtsi32_si128((signed char)pB[0]); + __m128i _pB16 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); + __m128i _pB32 = _mm_unpacklo_epi16(_pB16, _mm_setzero_si128()); + __m128i _pB = _mm_shuffle_epi32(_pB32, _MM_SHUFFLE(0, 0, 0, 0)); _sum = _mm_comp_dpwssd_epi32(_sum, _pA, _pB); pA += 2; pB++; } - const __m128 _ad = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pA_descales); + __m128 _ad = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pA_descales); _fsum = _mm_add_ps(_fsum, _mm_mul_ps(_mm_cvtepi32_ps(_sum), _mm_mul_ps(_ad, _mm_set1_ps(pB_descales[0])))); pA_descales += 2; pB_descales++; } + pB += remain_K * 1; + pB_descales += remain_blocks * 1; + _mm_storel_pi((__m64*)outptr, _fsum); outptr += 2; #else - float fsum0 = 0.f; - float fsum1 = 0.f; + float fsum0 = k == 0 ? 0.f : outptr[0]; + float fsum1 = k == 0 ? 0.f : outptr[1]; const signed char* pA = pAT; const float* pA_descales = pAT_descales; - for (int k = 0; k < K; k += block_size) + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) { int sum0 = 0; int sum1 = 0; - const int max_kk = std::min(K - k, block_size); + const int max_kk = std::min(tile_K - kk0, block_size); int kk = 0; for (; kk + 3 < max_kk; kk += 4) { @@ -4036,6 +3817,9 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB_descales++; } + pB += remain_K * 1; + pB_descales += remain_blocks * 1; + outptr[0] = fsum0; outptr[1] = fsum1; outptr += 2; @@ -4056,38 +3840,40 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #if __AVX512F__ for (; jj + 7 < max_jj; jj += 8) { - __m256 _fsum = _mm256_setzero_ps(); + pB += k * 8; + pB_descales += block_start * 8; + __m256 _fsum = k == 0 ? _mm256_setzero_ps() : _mm256_loadu_ps(outptr); const signed char* pA = pAT; const float* pA_descales = pAT_descales; - for (int k = 0; k < K; k += block_size) + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) { __m256i _sum = _mm256_setzero_si256(); - const int max_kk = std::min(K - k, block_size); + const int max_kk = std::min(tile_K - kk0, block_size); int kk = 0; #if __AVX512VNNI__ for (; kk + 3 < max_kk; kk += 4) { - const __m128i _pA32 = _mm_cvtsi32_si128(*(const int*)pA); - const __m256i _pA = _mm256_broadcastd_epi32(_pA32); - const __m256i _pB = _mm256_loadu_si256((const __m256i*)pB); + __m128i _pA32 = _mm_loadu_si32(pA); + __m256i _pA = _mm256_broadcastd_epi32(_pA32); + __m256i _pB = _mm256_loadu_si256((const __m256i*)pB); _sum = _mm256_comp_dpbusd_epi32(_sum, _pB, _pA); pA += 4; pB += 32; } if (max_kk >= 4) { - _sum = _mm256_sub_epi32(_sum, _mm256_set1_epi32(*(const int*)pA)); + _sum = _mm256_sub_epi32(_sum, _mm256_set1_epi32(_mm_cvtsi128_si32(_mm_loadu_si32(pA)))); pA += 4; } #else for (; kk + 3 < max_kk; kk += 4) { - const __m128i _pA8 = _mm_cvtsi32_si128(*(const int*)pA); - const __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); - const __m256i _pA01 = _mm256_broadcastsi128_si256(_mm_shuffle_epi32(_pA16, _MM_SHUFFLE(0, 0, 0, 0))); - const __m256i _pA23 = _mm256_broadcastsi128_si256(_mm_shuffle_epi32(_pA16, _MM_SHUFFLE(1, 1, 1, 1))); - const __m256i _pB01 = _mm256_cvtepi8_epi16(_mm_loadu_si128((const __m128i*)pB)); - const __m256i _pB23 = _mm256_cvtepi8_epi16(_mm_loadu_si128((const __m128i*)(pB + 16))); + __m128i _pA8 = _mm_loadu_si32(pA); + __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + __m256i _pA01 = _mm256_broadcastsi128_si256(_mm_shuffle_epi32(_pA16, _MM_SHUFFLE(0, 0, 0, 0))); + __m256i _pA23 = _mm256_broadcastsi128_si256(_mm_shuffle_epi32(_pA16, _MM_SHUFFLE(1, 1, 1, 1))); + __m256i _pB01 = _mm256_cvtepi8_epi16(_mm_loadu_si128((const __m128i*)pB)); + __m256i _pB23 = _mm256_cvtepi8_epi16(_mm_loadu_si128((const __m128i*)(pB + 16))); _sum = _mm256_comp_dpwssd_epi32(_sum, _pA01, _pB01); _sum = _mm256_comp_dpwssd_epi32(_sum, _pA23, _pB23); pA += 4; @@ -4096,50 +3882,55 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #endif // __AVX512VNNI__ for (; kk + 1 < max_kk; kk += 2) { - const __m128i _pA8 = _mm_cvtsi32_si128(*(const unsigned short*)pA); - const __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); - const __m256i _pA = _mm256_broadcastsi128_si256(_mm_shuffle_epi32(_pA16, _MM_SHUFFLE(0, 0, 0, 0))); - const __m256i _pB = _mm256_cvtepi8_epi16(_mm_loadu_si128((const __m128i*)pB)); + __m128i _pA8 = _mm_loadu_si16(pA); + __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + __m256i _pA = _mm256_broadcastsi128_si256(_mm_shuffle_epi32(_pA16, _MM_SHUFFLE(0, 0, 0, 0))); + __m256i _pB = _mm256_cvtepi8_epi16(_mm_loadu_si128((const __m128i*)pB)); _sum = _mm256_comp_dpwssd_epi32(_sum, _pA, _pB); pA += 2; pB += 16; } for (; kk < max_kk; kk++) { - const __m128i _pA8 = _mm_cvtsi32_si128((signed char)pA[0]); - const __m256i _pA = _mm256_broadcastd_epi32(_pA8); - const __m256i _pB = _mm256_cvtepi8_epi32(_mm_loadl_epi64((const __m128i*)pB)); + __m128i _pA8 = _mm_cvtsi32_si128((signed char)pA[0]); + __m256i _pA = _mm256_broadcastd_epi32(_pA8); + __m256i _pB = _mm256_cvtepi8_epi32(_mm_loadl_epi64((const __m128i*)pB)); _sum = _mm256_add_epi32(_sum, _mm256_mullo_epi32(_pA, _pB)); pA++; pB += 8; } - const __m128 _ad1 = _mm_load_ss(pA_descales); - const __m256 _ad = _mm256_broadcastss_ps(_ad1); - const __m256 _descale = _mm256_mul_ps(_ad, _mm256_loadu_ps(pB_descales)); + __m128 _ad1 = _mm_load_ss(pA_descales); + __m256 _ad = _mm256_broadcastss_ps(_ad1); + __m256 _descale = _mm256_mul_ps(_ad, _mm256_loadu_ps(pB_descales)); _fsum = _mm256_add_ps(_fsum, _mm256_mul_ps(_mm256_cvtepi32_ps(_sum), _descale)); pA_descales += 1; pB_descales += 8; } + pB += remain_K * 8; + pB_descales += remain_blocks * 8; + _mm256_storeu_ps(outptr, _fsum); outptr += 8; } #endif // __AVX512F__ for (; jj + 3 < max_jj; jj += 4) { - __m128 _fsum = _mm_setzero_ps(); + pB += k * 4; + pB_descales += block_start * 4; + __m128 _fsum = k == 0 ? _mm_setzero_ps() : _mm_loadu_ps(outptr); const signed char* pA = pAT; const float* pA_descales = pAT_descales; - for (int k = 0; k < K; k += block_size) + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) { __m128i _sum = _mm_setzero_si128(); - const int max_kk = std::min(K - k, block_size); + const int max_kk = std::min(tile_K - kk0, block_size); int kk = 0; #if __AVX512VNNI__ || __AVXVNNI__ for (; kk + 3 < max_kk; kk += 4) { - const __m128i _pA32 = _mm_cvtsi32_si128(*(const int*)pA); - const __m128i _pA = _mm_shuffle_epi32(_pA32, _MM_SHUFFLE(0, 0, 0, 0)); - const __m128i _pB = _mm_loadu_si128((const __m128i*)pB); + __m128i _pA32 = _mm_loadu_si32(pA); + __m128i _pA = _mm_shuffle_epi32(_pA32, _MM_SHUFFLE(0, 0, 0, 0)); + __m128i _pB = _mm_loadu_si128((const __m128i*)pB); #if __AVXVNNIINT8__ _sum = _mm_dpbssd_epi32(_sum, _pB, _pA); #else // __AVXVNNIINT8__ @@ -4151,21 +3942,21 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) if (max_kk >= 4) { - _sum = _mm_sub_epi32(_sum, _mm_set1_epi32(*(const int*)pA)); + _sum = _mm_sub_epi32(_sum, _mm_set1_epi32(_mm_cvtsi128_si32(_mm_loadu_si32(pA)))); pA += 4; } #endif #else for (; kk + 3 < max_kk; kk += 4) { - const __m128i _pA8 = _mm_cvtsi32_si128(*(const int*)pA); - const __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); - const __m128i _pA01 = _mm_shuffle_epi32(_pA16, _MM_SHUFFLE(0, 0, 0, 0)); - const __m128i _pA23 = _mm_shuffle_epi32(_pA16, _MM_SHUFFLE(1, 1, 1, 1)); - const __m128i _pB01x1 = _mm_loadl_epi64((const __m128i*)pB); - const __m128i _pB23x1 = _mm_loadl_epi64((const __m128i*)(pB + 8)); - const __m128i _pB01 = _mm_unpacklo_epi8(_pB01x1, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB01x1)); - const __m128i _pB23 = _mm_unpacklo_epi8(_pB23x1, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB23x1)); + __m128i _pA8 = _mm_loadu_si32(pA); + __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + __m128i _pA01 = _mm_shuffle_epi32(_pA16, _MM_SHUFFLE(0, 0, 0, 0)); + __m128i _pA23 = _mm_shuffle_epi32(_pA16, _MM_SHUFFLE(1, 1, 1, 1)); + __m128i _pB01x1 = _mm_loadl_epi64((const __m128i*)pB); + __m128i _pB23x1 = _mm_loadl_epi64((const __m128i*)(pB + 8)); + __m128i _pB01 = _mm_unpacklo_epi8(_pB01x1, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB01x1)); + __m128i _pB23 = _mm_unpacklo_epi8(_pB23x1, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB23x1)); _sum = _mm_comp_dpwssd_epi32(_sum, _pA01, _pB01); _sum = _mm_comp_dpwssd_epi32(_sum, _pA23, _pB23); pA += 4; @@ -4174,34 +3965,37 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #endif // __AVX512VNNI__ || __AVXVNNI__ for (; kk + 1 < max_kk; kk += 2) { - const __m128i _pA8 = _mm_cvtsi32_si128(*(const unsigned short*)pA); - const __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); - const __m128i _pA = _mm_shuffle_epi32(_pA16, _MM_SHUFFLE(0, 0, 0, 0)); - const __m128i _pB8 = _mm_loadl_epi64((const __m128i*)pB); - const __m128i _pB = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); + __m128i _pA8 = _mm_loadu_si16(pA); + __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + __m128i _pA = _mm_shuffle_epi32(_pA16, _MM_SHUFFLE(0, 0, 0, 0)); + __m128i _pB8 = _mm_loadl_epi64((const __m128i*)pB); + __m128i _pB = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); _sum = _mm_comp_dpwssd_epi32(_sum, _pA, _pB); pA += 2; pB += 8; } for (; kk < max_kk; kk++) { - const __m128i _pA8 = _mm_cvtsi32_si128((signed char)pA[0]); - const __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); - const __m128i _pA = _mm_shuffle_epi32(_mm_unpacklo_epi16(_pA16, _pA16), _MM_SHUFFLE(0, 0, 0, 0)); - const __m128i _pB8 = _mm_cvtsi32_si128(*(const int*)pB); - const __m128i _pB16 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); - const __m128i _pB = _mm_unpacklo_epi16(_pB16, _mm_setzero_si128()); + __m128i _pA8 = _mm_cvtsi32_si128((signed char)pA[0]); + __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + __m128i _pA = _mm_shuffle_epi32(_mm_unpacklo_epi16(_pA16, _pA16), _MM_SHUFFLE(0, 0, 0, 0)); + __m128i _pB8 = _mm_loadu_si32(pB); + __m128i _pB16 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); + __m128i _pB = _mm_unpacklo_epi16(_pB16, _mm_setzero_si128()); _sum = _mm_comp_dpwssd_epi32(_sum, _pA, _pB); pA++; pB += 4; } - const __m128 _ad1 = _mm_load_ss(pA_descales); - const __m128 _ad = _mm_shuffle_ps(_ad1, _ad1, _MM_SHUFFLE(0, 0, 0, 0)); - const __m128 _descale = _mm_mul_ps(_ad, _mm_loadu_ps(pB_descales)); + __m128 _ad1 = _mm_load_ss(pA_descales); + __m128 _ad = _mm_shuffle_ps(_ad1, _ad1, _MM_SHUFFLE(0, 0, 0, 0)); + __m128 _descale = _mm_mul_ps(_ad, _mm_loadu_ps(pB_descales)); _fsum = _mm_add_ps(_fsum, _mm_mul_ps(_mm_cvtepi32_ps(_sum), _descale)); pA_descales += 1; pB_descales += 4; } + pB += remain_K * 4; + pB_descales += remain_blocks * 4; + _mm_storeu_ps(outptr, _fsum); outptr += 4; } @@ -4209,21 +4003,23 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #endif // __SSE2__ for (; jj + 1 < max_jj; jj += 2) { + pB += k * 2; + pB_descales += block_start * 2; #if __SSE2__ - __m128 _fsum = _mm_setzero_ps(); + __m128 _fsum = k == 0 ? _mm_setzero_ps() : _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)outptr); const signed char* pA = pAT; const float* pA_descales = pAT_descales; - for (int k = 0; k < K; k += block_size) + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) { __m128i _sum = _mm_setzero_si128(); - const int max_kk = std::min(K - k, block_size); + const int max_kk = std::min(tile_K - kk0, block_size); int kk = 0; #if __AVX512VNNI__ || __AVXVNNI__ for (; kk + 3 < max_kk; kk += 4) { - const __m128i _pA32 = _mm_cvtsi32_si128(*(const int*)pA); - const __m128i _pA = _mm_shuffle_epi32(_pA32, _MM_SHUFFLE(0, 0, 0, 0)); - const __m128i _pB = _mm_loadl_epi64((const __m128i*)pB); + __m128i _pA32 = _mm_loadu_si32(pA); + __m128i _pA = _mm_shuffle_epi32(_pA32, _MM_SHUFFLE(0, 0, 0, 0)); + __m128i _pB = _mm_loadl_epi64((const __m128i*)pB); #if __AVXVNNIINT8__ _sum = _mm_dpbssd_epi32(_sum, _pB, _pA); #else // __AVXVNNIINT8__ @@ -4235,20 +4031,20 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) if (max_kk >= 4) { - _sum = _mm_sub_epi32(_sum, _mm_set1_epi32(*(const int*)pA)); + _sum = _mm_sub_epi32(_sum, _mm_set1_epi32(_mm_cvtsi128_si32(_mm_loadu_si32(pA)))); pA += 4; } #endif #else for (; kk + 3 < max_kk; kk += 4) { - const __m128i _pA8 = _mm_cvtsi32_si128(*(const int*)pA); - const __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); - const __m128i _pA01 = _mm_shuffle_epi32(_pA16, _MM_SHUFFLE(0, 0, 0, 0)); - const __m128i _pA23 = _mm_shuffle_epi32(_pA16, _MM_SHUFFLE(1, 1, 1, 1)); - const __m128i _pB8 = _mm_loadl_epi64((const __m128i*)pB); - const __m128i _pB16 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); - const __m128i _pB23 = _mm_shuffle_epi32(_pB16, _MM_SHUFFLE(3, 2, 3, 2)); + __m128i _pA8 = _mm_loadu_si32(pA); + __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + __m128i _pA01 = _mm_shuffle_epi32(_pA16, _MM_SHUFFLE(0, 0, 0, 0)); + __m128i _pA23 = _mm_shuffle_epi32(_pA16, _MM_SHUFFLE(1, 1, 1, 1)); + __m128i _pB8 = _mm_loadl_epi64((const __m128i*)pB); + __m128i _pB16 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); + __m128i _pB23 = _mm_shuffle_epi32(_pB16, _MM_SHUFFLE(3, 2, 3, 2)); _sum = _mm_comp_dpwssd_epi32(_sum, _pA01, _pB16); _sum = _mm_comp_dpwssd_epi32(_sum, _pA23, _pB23); pA += 4; @@ -4257,46 +4053,49 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #endif // __AVX512VNNI__ || __AVXVNNI__ for (; kk + 1 < max_kk; kk += 2) { - const __m128i _pA8 = _mm_cvtsi32_si128(*(const unsigned short*)pA); - const __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); - const __m128i _pA = _mm_shuffle_epi32(_pA16, _MM_SHUFFLE(0, 0, 0, 0)); - const __m128i _pB8 = _mm_cvtsi32_si128(*(const int*)pB); - const __m128i _pB = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); + __m128i _pA8 = _mm_loadu_si16(pA); + __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + __m128i _pA = _mm_shuffle_epi32(_pA16, _MM_SHUFFLE(0, 0, 0, 0)); + __m128i _pB8 = _mm_loadu_si32(pB); + __m128i _pB = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); _sum = _mm_comp_dpwssd_epi32(_sum, _pA, _pB); pA += 2; pB += 4; } for (; kk < max_kk; kk++) { - const __m128i _pA8 = _mm_cvtsi32_si128((signed char)pA[0]); - const __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); - const __m128i _pA = _mm_shuffle_epi32(_mm_unpacklo_epi16(_pA16, _pA16), _MM_SHUFFLE(0, 0, 0, 0)); - const __m128i _pB8 = _mm_cvtsi32_si128(*(const unsigned short*)pB); - const __m128i _pB16 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); - const __m128i _pB = _mm_unpacklo_epi16(_pB16, _mm_setzero_si128()); + __m128i _pA8 = _mm_cvtsi32_si128((signed char)pA[0]); + __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + __m128i _pA = _mm_shuffle_epi32(_mm_unpacklo_epi16(_pA16, _pA16), _MM_SHUFFLE(0, 0, 0, 0)); + __m128i _pB8 = _mm_loadu_si16(pB); + __m128i _pB16 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); + __m128i _pB = _mm_unpacklo_epi16(_pB16, _mm_setzero_si128()); _sum = _mm_comp_dpwssd_epi32(_sum, _pA, _pB); pA++; pB += 2; } - const __m128 _ad1 = _mm_load_ss(pA_descales); - const __m128 _ad = _mm_shuffle_ps(_ad1, _ad1, _MM_SHUFFLE(0, 0, 0, 0)); - const __m128 _bd = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pB_descales); + __m128 _ad1 = _mm_load_ss(pA_descales); + __m128 _ad = _mm_shuffle_ps(_ad1, _ad1, _MM_SHUFFLE(0, 0, 0, 0)); + __m128 _bd = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pB_descales); _fsum = _mm_add_ps(_fsum, _mm_mul_ps(_mm_cvtepi32_ps(_sum), _mm_mul_ps(_ad, _bd))); pA_descales += 1; pB_descales += 2; } + pB += remain_K * 2; + pB_descales += remain_blocks * 2; + _mm_storel_pi((__m64*)outptr, _fsum); outptr += 2; #else - float fsum0 = 0.f; - float fsum1 = 0.f; + float fsum0 = k == 0 ? 0.f : outptr[0]; + float fsum1 = k == 0 ? 0.f : outptr[1]; const signed char* pA = pAT; const float* pA_descales = pAT_descales; - for (int k = 0; k < K; k += block_size) + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) { int sum0 = 0; int sum1 = 0; - const int max_kk = std::min(K - k, block_size); + const int max_kk = std::min(tile_K - kk0, block_size); int kk = 0; for (; kk + 3 < max_kk; kk += 4) { @@ -4326,6 +4125,9 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB_descales += 2; } + pB += remain_K * 2; + pB_descales += remain_blocks * 2; + outptr[0] = fsum0; outptr[1] = fsum1; outptr += 2; @@ -4333,20 +4135,22 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de } for (; jj < max_jj; jj++) { + pB += k; + pB_descales += block_start; #if __SSE2__ - float fsum = 0.f; + float fsum = k == 0 ? 0.f : outptr[0]; const signed char* pA = pAT; const float* pA_descales = pAT_descales; - for (int k = 0; k < K; k += block_size) + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) { __m128i _sum = _mm_setzero_si128(); - const int max_kk = std::min(K - k, block_size); + const int max_kk = std::min(tile_K - kk0, block_size); int kk = 0; #if __AVX512VNNI__ || __AVXVNNI__ for (; kk + 3 < max_kk; kk += 4) { - const __m128i _pA = _mm_cvtsi32_si128(*(const int*)pA); - const __m128i _pB = _mm_cvtsi32_si128(*(const int*)pB); + __m128i _pA = _mm_loadu_si32(pA); + __m128i _pB = _mm_loadu_si32(pB); #if __AVXVNNIINT8__ _sum = _mm_dpbssd_epi32(_sum, _pB, _pA); #else // __AVXVNNIINT8__ @@ -4358,17 +4162,17 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #if __AVX512VNNI__ || (__AVXVNNI__ && !__AVXVNNIINT8__) if (max_kk >= 4) { - _sum = _mm_sub_epi32(_sum, _mm_cvtsi32_si128(*(const int*)pA)); + _sum = _mm_sub_epi32(_sum, _mm_loadu_si32(pA)); pA += 4; } #endif #else for (; kk + 3 < max_kk; kk += 4) { - const __m128i _pA8 = _mm_cvtsi32_si128(*(const int*)pA); - const __m128i _pB8 = _mm_cvtsi32_si128(*(const int*)pB); - const __m128i _pA = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); - const __m128i _pB = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); + __m128i _pA8 = _mm_loadu_si32(pA); + __m128i _pB8 = _mm_loadu_si32(pB); + __m128i _pA = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + __m128i _pB = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); _sum = _mm_comp_dpwssd_epi32(_sum, _pA, _pB); pA += 4; pB += 4; @@ -4376,20 +4180,20 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de #endif // __AVX512VNNI__ || __AVXVNNI__ for (; kk + 1 < max_kk; kk += 2) { - const __m128i _pA8 = _mm_cvtsi32_si128(*(const unsigned short*)pA); - const __m128i _pB8 = _mm_cvtsi32_si128(*(const unsigned short*)pB); - const __m128i _pA = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); - const __m128i _pB = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); + __m128i _pA8 = _mm_loadu_si16(pA); + __m128i _pB8 = _mm_loadu_si16(pB); + __m128i _pA = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + __m128i _pB = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); _sum = _mm_comp_dpwssd_epi32(_sum, _pA, _pB); pA += 2; pB += 2; } for (; kk < max_kk; kk++) { - const __m128i _pA8 = _mm_cvtsi32_si128((unsigned char)pA[0]); - const __m128i _pB8 = _mm_cvtsi32_si128((unsigned char)pB[0]); - const __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); - const __m128i _pB16 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); + __m128i _pA8 = _mm_cvtsi32_si128((unsigned char)pA[0]); + __m128i _pB8 = _mm_cvtsi32_si128((unsigned char)pB[0]); + __m128i _pA16 = _mm_unpacklo_epi8(_pA8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pA8)); + __m128i _pB16 = _mm_unpacklo_epi8(_pB8, _mm_cmpgt_epi8(_mm_setzero_si128(), _pB8)); _sum = _mm_comp_dpwssd_epi32(_sum, _mm_unpacklo_epi16(_pA16, _pA16), _mm_unpacklo_epi16(_pB16, _mm_setzero_si128())); pA++; pB++; @@ -4398,16 +4202,19 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pA_descales += 1; pB_descales++; } + pB += remain_K * 1; + pB_descales += remain_blocks * 1; + outptr[0] = fsum; outptr++; #else - float fsum = 0.f; + float fsum = k == 0 ? 0.f : outptr[0]; const signed char* pA = pAT; const float* pA_descales = pAT_descales; - for (int k = 0; k < K; k += block_size) + for (int kk0 = 0; kk0 < tile_K; kk0 += block_size) { int sum = 0; - const int max_kk = std::min(K - k, block_size); + const int max_kk = std::min(tile_K - kk0, block_size); int kk = 0; for (; kk + 3 < max_kk; kk += 4) { @@ -4430,6 +4237,9 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB_descales++; } + pB += remain_K * 1; + pB_descales += remain_blocks * 1; + outptr[0] = fsum; outptr++; #endif // __SSE2__ @@ -4471,9 +4281,12 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b } #endif + (void)N; + const float* pC = C; const float* pp = topT; const size_t out_hstep = top_blob.dims == 3 ? top_blob.cstep : (size_t)top_blob.w; + const size_t c_hstep = C.dims == 3 ? C.cstep : (size_t)C.w; int ii = 0; #if __SSE2__ @@ -4493,9 +4306,9 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b _c = _mm512_loadu_ps(pC + i + ii); if (broadcast_type_C == 3) { - pC = (const float*)C + (i + ii) * N + j; + pC = (const float*)C + (i + ii) * c_hstep + j; _c_vindex = _mm512_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15); - _c_vindex = _mm512_mullo_epi32(_c_vindex, _mm512_set1_epi32(N)); + _c_vindex = _mm512_mullo_epi32(_c_vindex, _mm512_set1_epi32((int)c_hstep)); } if (broadcast_type_C == 4) pC = (const float*)C + j; @@ -4628,7 +4441,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b } if (alpha != 1.f) { - const __m512 _alpha = _mm512_set1_ps(alpha); + __m512 _alpha = _mm512_set1_ps(alpha); _f0 = _mm512_mul_ps(_f0, _alpha); _f1 = _mm512_mul_ps(_f1, _alpha); _f2 = _mm512_mul_ps(_f2, _alpha); @@ -4720,7 +4533,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b } if (alpha != 1.f) { - const __m512 _alpha = _mm512_set1_ps(alpha); + __m512 _alpha = _mm512_set1_ps(alpha); _f0 = _mm512_mul_ps(_f0, _alpha); _f1 = _mm512_mul_ps(_f1, _alpha); _f2 = _mm512_mul_ps(_f2, _alpha); @@ -4786,48 +4599,48 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b } if (alpha != 1.f) { - const __m512 _alpha = _mm512_set1_ps(alpha); + __m512 _alpha = _mm512_set1_ps(alpha); _f0 = _mm512_mul_ps(_f0, _alpha); _f1 = _mm512_mul_ps(_f1, _alpha); } transpose16x2_ps(_f0, _f1); { - const __m128 _r = _mm512_extractf32x4_ps(_f0, 0); + __m128 _r = _mm512_extractf32x4_ps(_f0, 0); _mm_storel_pi((__m64*)(p0), _r); _mm_storeh_pi((__m64*)(p0 + out_hstep), _r); } { - const __m128 _r = _mm512_extractf32x4_ps(_f0, 1); + __m128 _r = _mm512_extractf32x4_ps(_f0, 1); _mm_storel_pi((__m64*)(p0 + out_hstep * 2), _r); _mm_storeh_pi((__m64*)(p0 + out_hstep * 3), _r); } { - const __m128 _r = _mm512_extractf32x4_ps(_f0, 2); + __m128 _r = _mm512_extractf32x4_ps(_f0, 2); _mm_storel_pi((__m64*)(p0 + out_hstep * 4), _r); _mm_storeh_pi((__m64*)(p0 + out_hstep * 5), _r); } { - const __m128 _r = _mm512_extractf32x4_ps(_f0, 3); + __m128 _r = _mm512_extractf32x4_ps(_f0, 3); _mm_storel_pi((__m64*)(p0 + out_hstep * 6), _r); _mm_storeh_pi((__m64*)(p0 + out_hstep * 7), _r); } { - const __m128 _r = _mm512_extractf32x4_ps(_f1, 0); + __m128 _r = _mm512_extractf32x4_ps(_f1, 0); _mm_storel_pi((__m64*)(p0 + out_hstep * 8), _r); _mm_storeh_pi((__m64*)(p0 + out_hstep * 9), _r); } { - const __m128 _r = _mm512_extractf32x4_ps(_f1, 1); + __m128 _r = _mm512_extractf32x4_ps(_f1, 1); _mm_storel_pi((__m64*)(p0 + out_hstep * 10), _r); _mm_storeh_pi((__m64*)(p0 + out_hstep * 11), _r); } { - const __m128 _r = _mm512_extractf32x4_ps(_f1, 2); + __m128 _r = _mm512_extractf32x4_ps(_f1, 2); _mm_storel_pi((__m64*)(p0 + out_hstep * 12), _r); _mm_storeh_pi((__m64*)(p0 + out_hstep * 13), _r); } { - const __m128 _r = _mm512_extractf32x4_ps(_f1, 3); + __m128 _r = _mm512_extractf32x4_ps(_f1, 3); _mm_storel_pi((__m64*)(p0 + out_hstep * 14), _r); _mm_storeh_pi((__m64*)(p0 + out_hstep * 15), _r); } @@ -4860,7 +4673,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b } if (alpha != 1.f) { - const __m512 _alpha = _mm512_set1_ps(alpha); + __m512 _alpha = _mm512_set1_ps(alpha); _f0 = _mm512_mul_ps(_f0, _alpha); } __m512i _vindex = _mm512_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15); @@ -4910,7 +4723,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b c7 = pC[i + ii + 7]; } if (broadcast_type_C == 3) - pC = (const float*)C + (i + ii) * N + j; + pC = (const float*)C + (i + ii) * c_hstep + j; if (broadcast_type_C == 4) pC = (const float*)C + j; if (broadcast_type_C == 0 && beta != 1.f) @@ -5008,13 +4821,13 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b if (broadcast_type_C == 3) { __m256 _c0 = _mm256_loadu_ps(pC); - __m256 _c1 = _mm256_loadu_ps(pC + N); - __m256 _c2 = _mm256_loadu_ps(pC + N * 2); - __m256 _c3 = _mm256_loadu_ps(pC + N * 3); - __m256 _c4 = _mm256_loadu_ps(pC + N * 4); - __m256 _c5 = _mm256_loadu_ps(pC + N * 5); - __m256 _c6 = _mm256_loadu_ps(pC + N * 6); - __m256 _c7 = _mm256_loadu_ps(pC + N * 7); + __m256 _c1 = _mm256_loadu_ps(pC + c_hstep); + __m256 _c2 = _mm256_loadu_ps(pC + c_hstep * 2); + __m256 _c3 = _mm256_loadu_ps(pC + c_hstep * 3); + __m256 _c4 = _mm256_loadu_ps(pC + c_hstep * 4); + __m256 _c5 = _mm256_loadu_ps(pC + c_hstep * 5); + __m256 _c6 = _mm256_loadu_ps(pC + c_hstep * 6); + __m256 _c7 = _mm256_loadu_ps(pC + c_hstep * 7); if (beta == 1.f) { _f0 = _mm256_add_ps(_f0, _c0); @@ -5153,13 +4966,13 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b if (broadcast_type_C == 3) { __m128 _c0 = _mm_loadu_ps(pC); - __m128 _c1 = _mm_loadu_ps(pC + N); - __m128 _c2 = _mm_loadu_ps(pC + N * 2); - __m128 _c3 = _mm_loadu_ps(pC + N * 3); - __m128 _c4 = _mm_loadu_ps(pC + N * 4); - __m128 _c5 = _mm_loadu_ps(pC + N * 5); - __m128 _c6 = _mm_loadu_ps(pC + N * 6); - __m128 _c7 = _mm_loadu_ps(pC + N * 7); + __m128 _c1 = _mm_loadu_ps(pC + c_hstep); + __m128 _c2 = _mm_loadu_ps(pC + c_hstep * 2); + __m128 _c3 = _mm_loadu_ps(pC + c_hstep * 3); + __m128 _c4 = _mm_loadu_ps(pC + c_hstep * 4); + __m128 _c5 = _mm_loadu_ps(pC + c_hstep * 5); + __m128 _c6 = _mm_loadu_ps(pC + c_hstep * 6); + __m128 _c7 = _mm_loadu_ps(pC + c_hstep * 7); if (beta == 1.f) { _f0 = _mm_add_ps(_f0, _c0); @@ -5296,13 +5109,13 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b if (broadcast_type_C == 3) { __m128 _c0 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pC); - __m128 _c1 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N)); - __m128 _c2 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N * 2)); - __m128 _c3 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N * 3)); - __m128 _c4 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N * 4)); - __m128 _c5 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N * 5)); - __m128 _c6 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N * 6)); - __m128 _c7 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N * 7)); + __m128 _c1 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + c_hstep)); + __m128 _c2 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + c_hstep * 2)); + __m128 _c3 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + c_hstep * 3)); + __m128 _c4 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + c_hstep * 4)); + __m128 _c5 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + c_hstep * 5)); + __m128 _c6 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + c_hstep * 6)); + __m128 _c7 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + c_hstep * 7)); if (beta == 1.f) { _f0 = _mm_add_ps(_f0, _c0); @@ -5386,97 +5199,65 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b for (; jj < max_jj; jj++) { #if __AVX2__ - const float* pp4 = pp + 4; + __m256 _f0 = _mm256_loadu_ps(pp); #else - const float* pp4 = pp1; + __m256 _f0 = _mm256_insertf128_ps(_mm256_castps128_ps256(_mm_loadu_ps(pp)), _mm_loadu_ps(pp1), 1); #endif - float f0 = pp[0]; - if (pC) - { - if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) f0 += c0; - if (broadcast_type_C == 3) f0 += pC[0] * beta; - if (broadcast_type_C == 4) f0 += pC[0] * beta; - } - if (alpha != 1.f) f0 *= alpha; - p0[0] = f0; - float f1 = pp[1]; - if (pC) - { - if (broadcast_type_C == 0) f1 += c0; - if (broadcast_type_C == 1 || broadcast_type_C == 2) f1 += c1; - if (broadcast_type_C == 3) f1 += pC[N] * beta; - if (broadcast_type_C == 4) f1 += pC[0] * beta; - } - if (alpha != 1.f) f1 *= alpha; - p0[out_hstep] = f1; - float f2 = pp[2]; - if (pC) - { - if (broadcast_type_C == 0) f2 += c0; - if (broadcast_type_C == 1 || broadcast_type_C == 2) f2 += c2; - if (broadcast_type_C == 3) f2 += pC[N * 2] * beta; - if (broadcast_type_C == 4) f2 += pC[0] * beta; - } - if (alpha != 1.f) f2 *= alpha; - p0[out_hstep * 2] = f2; - float f3 = pp[3]; - if (pC) - { - if (broadcast_type_C == 0) f3 += c0; - if (broadcast_type_C == 1 || broadcast_type_C == 2) f3 += c3; - if (broadcast_type_C == 3) f3 += pC[N * 3] * beta; - if (broadcast_type_C == 4) f3 += pC[0] * beta; - } - if (alpha != 1.f) f3 *= alpha; - p0[out_hstep * 3] = f3; - float f4 = pp4[0]; - if (pC) - { - if (broadcast_type_C == 0) f4 += c0; - if (broadcast_type_C == 1 || broadcast_type_C == 2) f4 += c4; - if (broadcast_type_C == 3) f4 += pC[N * 4] * beta; - if (broadcast_type_C == 4) f4 += pC[0] * beta; - } - if (alpha != 1.f) f4 *= alpha; - p0[out_hstep * 4] = f4; - float f5 = pp4[1]; - if (pC) - { - if (broadcast_type_C == 0) f5 += c0; - if (broadcast_type_C == 1 || broadcast_type_C == 2) f5 += c5; - if (broadcast_type_C == 3) f5 += pC[N * 5] * beta; - if (broadcast_type_C == 4) f5 += pC[0] * beta; - } - if (alpha != 1.f) f5 *= alpha; - p0[out_hstep * 5] = f5; - float f6 = pp4[2]; - if (pC) - { - if (broadcast_type_C == 0) f6 += c0; - if (broadcast_type_C == 1 || broadcast_type_C == 2) f6 += c6; - if (broadcast_type_C == 3) f6 += pC[N * 6] * beta; - if (broadcast_type_C == 4) f6 += pC[0] * beta; - } - if (alpha != 1.f) f6 *= alpha; - p0[out_hstep * 6] = f6; - float f7 = pp4[3]; if (pC) { - if (broadcast_type_C == 0) f7 += c0; - if (broadcast_type_C == 1 || broadcast_type_C == 2) f7 += c7; - if (broadcast_type_C == 3) + if (broadcast_type_C == 0) { - f7 += pC[N * 7] * beta; - pC++; + _f0 = _mm256_add_ps(_f0, _mm256_set1_ps(c0)); } - if (broadcast_type_C == 4) + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = _mm256_add_ps(_f0, _mm256_setr_ps(c0, c1, c2, c3, c4, c5, c6, c7)); + } + if (broadcast_type_C == 3 || broadcast_type_C == 4) { - f7 += pC[0] * beta; + __m256 _c; + if (broadcast_type_C == 3) + { +#if __AVX2__ + __m256i _vindex = _mm256_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7); + _vindex = _mm256_mullo_epi32(_vindex, _mm256_set1_epi32((int)c_hstep)); + _c = _mm256_i32gather_ps(pC, _vindex, sizeof(float)); +#else + _c = _mm256_setr_ps(pC[0], pC[c_hstep], pC[c_hstep * 2], pC[c_hstep * 3], pC[c_hstep * 4], pC[c_hstep * 5], pC[c_hstep * 6], pC[c_hstep * 7]); +#endif + } + if (broadcast_type_C == 4) + _c = _mm256_set1_ps(pC[0]); + if (beta == 1.f) + _f0 = _mm256_add_ps(_f0, _c); + else + _f0 = _mm256_comp_fmadd_ps(_c, _mm256_set1_ps(beta), _f0); pC++; } } - if (alpha != 1.f) f7 *= alpha; - p0[out_hstep * 7] = f7; + if (alpha != 1.f) + _f0 = _mm256_mul_ps(_f0, _mm256_set1_ps(alpha)); +#if __AVX512F__ + __m256i _vindex = _mm256_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7); + _vindex = _mm256_mullo_epi32(_vindex, _mm256_set1_epi32((int)out_hstep)); + _mm256_i32scatter_ps(p0, _vindex, _f0, sizeof(float)); +#else +#ifdef _MSC_VER + __declspec(align(32)) +#else + __attribute__((aligned(32))) +#endif + float sum0[8]; + _mm256_store_ps(sum0, _f0); + p0[0] = sum0[0]; + p0[out_hstep] = sum0[1]; + p0[out_hstep * 2] = sum0[2]; + p0[out_hstep * 3] = sum0[3]; + p0[out_hstep * 4] = sum0[4]; + p0[out_hstep * 5] = sum0[5]; + p0[out_hstep * 6] = sum0[6]; + p0[out_hstep * 7] = sum0[7]; +#endif p0++; #if __AVX2__ pp += 8; @@ -5516,7 +5297,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b c3 = pC[i + ii + 3]; } if (broadcast_type_C == 3) - pC = (const float*)C + (i + ii) * N + j; + pC = (const float*)C + (i + ii) * c_hstep + j; if (broadcast_type_C == 4) pC = (const float*)C + j; if (broadcast_type_C == 0 && beta != 1.f) @@ -5601,9 +5382,9 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b if (broadcast_type_C == 3) { __m256 _c0 = _mm256_loadu_ps(pC); - __m256 _c1 = _mm256_loadu_ps(pC + N); - __m256 _c2 = _mm256_loadu_ps(pC + N * 2); - __m256 _c3 = _mm256_loadu_ps(pC + N * 3); + __m256 _c1 = _mm256_loadu_ps(pC + c_hstep); + __m256 _c2 = _mm256_loadu_ps(pC + c_hstep * 2); + __m256 _c3 = _mm256_loadu_ps(pC + c_hstep * 3); if (beta == 1.f) { _f0 = _mm256_add_ps(_f0, _c0); @@ -5697,9 +5478,9 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b if (broadcast_type_C == 3) { __m128 _c0 = _mm_loadu_ps(pC); - __m128 _c1 = _mm_loadu_ps(pC + N); - __m128 _c2 = _mm_loadu_ps(pC + N * 2); - __m128 _c3 = _mm_loadu_ps(pC + N * 3); + __m128 _c1 = _mm_loadu_ps(pC + c_hstep); + __m128 _c2 = _mm_loadu_ps(pC + c_hstep * 2); + __m128 _c3 = _mm_loadu_ps(pC + c_hstep * 3); if (beta == 1.f) { _f0 = _mm_add_ps(_f0, _c0); @@ -5787,9 +5568,9 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b if (broadcast_type_C == 3) { __m128 _c0 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pC); - __m128 _c1 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N)); - __m128 _c2 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N * 2)); - __m128 _c3 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N * 3)); + __m128 _c1 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + c_hstep)); + __m128 _c2 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + c_hstep * 2)); + __m128 _c3 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + c_hstep * 3)); if (beta == 1.f) { _f0 = _mm_add_ps(_f0, _c0); @@ -5843,52 +5624,50 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b } for (; jj < max_jj; jj += 1) { - float f0_0 = pp[0]; - if (pC) - { - if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) f0_0 += c0; - if (broadcast_type_C == 3 || broadcast_type_C == 4) f0_0 += pC[0] * beta; - } - if (alpha != 1.f) f0_0 *= alpha; - p0[0] = f0_0; - float f1_0 = pp[1]; - if (pC) - { - if (broadcast_type_C == 0) f1_0 += c0; - if (broadcast_type_C == 1 || broadcast_type_C == 2) f1_0 += c1; - if (broadcast_type_C == 3) f1_0 += pC[N] * beta; - if (broadcast_type_C == 4) f1_0 += pC[0] * beta; - } - if (alpha != 1.f) f1_0 *= alpha; - p0[out_hstep] = f1_0; - float f2_0 = pp[2]; - if (pC) - { - if (broadcast_type_C == 0) f2_0 += c0; - if (broadcast_type_C == 1 || broadcast_type_C == 2) f2_0 += c2; - if (broadcast_type_C == 3) f2_0 += pC[N * 2] * beta; - if (broadcast_type_C == 4) f2_0 += pC[0] * beta; - } - if (alpha != 1.f) f2_0 *= alpha; - p0[out_hstep * 2] = f2_0; - float f3_0 = pp[3]; + __m128 _f0 = _mm_loadu_ps(pp); if (pC) { - if (broadcast_type_C == 0) f3_0 += c0; - if (broadcast_type_C == 1 || broadcast_type_C == 2) f3_0 += c3; - if (broadcast_type_C == 3) + if (broadcast_type_C == 0) { - f3_0 += pC[N * 3] * beta; - pC++; + _f0 = _mm_add_ps(_f0, _mm_set1_ps(c0)); } - if (broadcast_type_C == 4) + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + _f0 = _mm_add_ps(_f0, _mm_setr_ps(c0, c1, c2, c3)); + } + if (broadcast_type_C == 3 || broadcast_type_C == 4) { - f3_0 += pC[0] * beta; + __m128 _c; + if (broadcast_type_C == 3) + _c = _mm_setr_ps(pC[0], pC[c_hstep], pC[c_hstep * 2], pC[c_hstep * 3]); + if (broadcast_type_C == 4) + _c = _mm_set1_ps(pC[0]); + if (beta == 1.f) + _f0 = _mm_add_ps(_f0, _c); + else + _f0 = _mm_comp_fmadd_ps(_c, _mm_set1_ps(beta), _f0); pC++; } } - if (alpha != 1.f) f3_0 *= alpha; - p0[out_hstep * 3] = f3_0; + if (alpha != 1.f) + _f0 = _mm_mul_ps(_f0, _mm_set1_ps(alpha)); +#if __AVX512F__ + __m128i _vindex = _mm_setr_epi32(0, 1, 2, 3); + _vindex = _mm_mullo_epi32(_vindex, _mm_set1_epi32((int)out_hstep)); + _mm_i32scatter_ps(p0, _vindex, _f0, sizeof(float)); +#else +#ifdef _MSC_VER + __declspec(align(16)) +#else + __attribute__((aligned(16))) +#endif + float sum0[4]; + _mm_store_ps(sum0, _f0); + p0[0] = sum0[0]; + p0[out_hstep] = sum0[1]; + p0[out_hstep * 2] = sum0[2]; + p0[out_hstep * 3] = sum0[3]; +#endif p0++; pp += 4; } @@ -5915,7 +5694,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b c1 = pC[i + ii + 1]; } if (broadcast_type_C == 3) - pC = (const float*)C + (i + ii) * N + j; + pC = (const float*)C + (i + ii) * c_hstep + j; if (broadcast_type_C == 4) pC = (const float*)C + j; if (broadcast_type_C == 0 && beta != 1.f) @@ -5967,7 +5746,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b if (broadcast_type_C == 3) { __m256 _c0 = _mm256_loadu_ps(pC); - __m256 _c1 = _mm256_loadu_ps(pC + N); + __m256 _c1 = _mm256_loadu_ps(pC + c_hstep); if (beta == 1.f) { _f0x2 = _mm256_add_ps(_f0x2, _c0); @@ -6033,7 +5812,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b if (broadcast_type_C == 3) { __m128 _c0 = _mm_loadu_ps(pC); - __m128 _c1 = _mm_loadu_ps(pC + N); + __m128 _c1 = _mm_loadu_ps(pC + c_hstep); if (beta == 1.f) { _f0 = _mm_add_ps(_f0, _c0); @@ -6094,7 +5873,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b if (broadcast_type_C == 3) { __m128 _c0 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pC); - __m128 _c1 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N)); + __m128 _c1 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + c_hstep)); if (beta == 1.f) { _f0 = _mm_add_ps(_f0, _c0); @@ -6136,83 +5915,73 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b } for (; jj < max_jj; jj += 1) { - float f0_0 = pp[0]; + __m128 _f0 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pp); if (pC) { - if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) f0_0 += c0; - if (broadcast_type_C == 3 || broadcast_type_C == 4) f0_0 += pC[0] * beta; - } - if (alpha != 1.f) f0_0 *= alpha; - p0[0] = f0_0; - float f1_0 = pp[1]; - if (pC) - { - if (broadcast_type_C == 0) f1_0 += c0; - if (broadcast_type_C == 1 || broadcast_type_C == 2) f1_0 += c1; - if (broadcast_type_C == 3) + if (broadcast_type_C == 0) { - f1_0 += pC[N] * beta; - pC++; + _f0 = _mm_add_ps(_f0, _mm_set1_ps(c0)); } - if (broadcast_type_C == 4) + if (broadcast_type_C == 1 || broadcast_type_C == 2) { - f1_0 += pC[0] * beta; + _f0 = _mm_add_ps(_f0, _mm_setr_ps(c0, c1, 0.f, 0.f)); + } + if (broadcast_type_C == 3 || broadcast_type_C == 4) + { + __m128 _c; + if (broadcast_type_C == 3) + _c = _mm_setr_ps(pC[0], pC[c_hstep], 0.f, 0.f); + if (broadcast_type_C == 4) + _c = _mm_set1_ps(pC[0]); + if (beta == 1.f) + _f0 = _mm_add_ps(_f0, _c); + else + _f0 = _mm_comp_fmadd_ps(_c, _mm_set1_ps(beta), _f0); pC++; } } - if (alpha != 1.f) f1_0 *= alpha; - p0[out_hstep] = f1_0; + if (alpha != 1.f) + _f0 = _mm_mul_ps(_f0, _mm_set1_ps(alpha)); + _mm_store_ss(p0, _f0); + _mm_store_ss(p0 + out_hstep, _mm_shuffle_ps(_f0, _f0, _MM_SHUFFLE(1, 1, 1, 1))); p0++; pp += 2; } -#else +#endif // __SSE2__ for (; jj + 1 < max_jj; jj += 2) { float f0_0 = pp[0]; float f0_1 = pp[1]; - if (pC) - { - if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) - { - f0_0 += c0; - f0_1 += c0; - } - if (broadcast_type_C == 3 || broadcast_type_C == 4) - { - f0_0 += pC[0] * beta; - f0_1 += pC[1] * beta; - } - } - if (alpha != 1.f) - { - f0_0 *= alpha; - f0_1 *= alpha; - } - p0[0] = f0_0; - p0[1] = f0_1; - float f1_0 = pp[2]; float f1_1 = pp[3]; if (pC) { if (broadcast_type_C == 0) { + f0_0 += c0; + f0_1 += c0; f1_0 += c0; f1_1 += c0; } if (broadcast_type_C == 1 || broadcast_type_C == 2) { + f0_0 += c0; + f0_1 += c0; f1_0 += c1; f1_1 += c1; } if (broadcast_type_C == 3) { - f1_0 += pC[N] * beta; - f1_1 += pC[N + 1] * beta; + f0_0 += pC[0] * beta; + f0_1 += pC[1] * beta; + f1_0 += pC[c_hstep] * beta; + f1_1 += pC[c_hstep + 1] * beta; pC += 2; } if (broadcast_type_C == 4) { + f0_0 += pC[0] * beta; + f0_1 += pC[1] * beta; f1_0 += pC[0] * beta; f1_1 += pC[1] * beta; pC += 2; @@ -6220,9 +5989,13 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b } if (alpha != 1.f) { + f0_0 *= alpha; + f0_1 *= alpha; f1_0 *= alpha; f1_1 *= alpha; } + p0[0] = f0_0; + p0[1] = f0_1; p0[out_hstep] = f1_0; p0[out_hstep + 1] = f1_1; @@ -6232,37 +6005,43 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b for (; jj < max_jj; jj += 1) { float f0_0 = pp[0]; - if (pC) - { - if (broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) f0_0 += c0; - if (broadcast_type_C == 3 || broadcast_type_C == 4) f0_0 += pC[0] * beta; - } - if (alpha != 1.f) f0_0 *= alpha; - p0[0] = f0_0; - float f1_0 = pp[1]; if (pC) { - if (broadcast_type_C == 0) f1_0 += c0; - if (broadcast_type_C == 1 || broadcast_type_C == 2) f1_0 += c1; + if (broadcast_type_C == 0) + { + f0_0 += c0; + f1_0 += c0; + } + if (broadcast_type_C == 1 || broadcast_type_C == 2) + { + f0_0 += c0; + f1_0 += c1; + } if (broadcast_type_C == 3) { - f1_0 += pC[N] * beta; + f0_0 += pC[0] * beta; + f1_0 += pC[c_hstep] * beta; pC++; } if (broadcast_type_C == 4) { + f0_0 += pC[0] * beta; f1_0 += pC[0] * beta; pC++; } } - if (alpha != 1.f) f1_0 *= alpha; + if (alpha != 1.f) + { + f0_0 *= alpha; + f1_0 *= alpha; + } + p0[0] = f0_0; p0[out_hstep] = f1_0; p0++; pp += 2; } -#endif // __SSE2__ } for (; ii < max_ii; ii += 1) @@ -6277,7 +6056,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b if (broadcast_type_C == 1 || broadcast_type_C == 2) c0 = pC[i + ii]; if (broadcast_type_C == 3) - pC = (const float*)C + (i + ii) * N + j; + pC = (const float*)C + (i + ii) * c_hstep + j; if (broadcast_type_C == 4) pC = (const float*)C + j; if ((broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) && beta != 1.f) @@ -6431,7 +6210,7 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b p0 += 2; pp += 2; } -#else +#endif // __SSE2__ for (; jj + 1 < max_jj; jj += 2) { float f0_0 = pp[0]; @@ -6460,7 +6239,6 @@ static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_b p0 += 2; pp += 2; } -#endif // __SSE2__ for (; jj < max_jj; jj += 1) { float f0_0 = pp[0]; @@ -6513,9 +6291,12 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, } #endif + (void)N; + const float* pC = C; const float* pp = topT; const size_t out_hstep = top_blob.dims == 3 ? top_blob.cstep : (size_t)top_blob.w; + const size_t c_hstep = C.dims == 3 ? C.cstep : (size_t)C.w; int ii = 0; #if __SSE2__ @@ -6535,9 +6316,9 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, _c = _mm512_loadu_ps(pC + i + ii); if (broadcast_type_C == 3) { - pC = (const float*)C + (i + ii) * N + j; + pC = (const float*)C + (i + ii) * c_hstep + j; _c_vindex = _mm512_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15); - _c_vindex = _mm512_mullo_epi32(_c_vindex, _mm512_set1_epi32(N)); + _c_vindex = _mm512_mullo_epi32(_c_vindex, _mm512_set1_epi32((int)c_hstep)); } if (broadcast_type_C == 4) pC = (const float*)C + j; @@ -6670,7 +6451,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, } if (alpha != 1.f) { - const __m512 _alpha = _mm512_set1_ps(alpha); + __m512 _alpha = _mm512_set1_ps(alpha); _f0 = _mm512_mul_ps(_f0, _alpha); _f1 = _mm512_mul_ps(_f1, _alpha); _f2 = _mm512_mul_ps(_f2, _alpha); @@ -6753,7 +6534,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, } if (alpha != 1.f) { - const __m512 _alpha = _mm512_set1_ps(alpha); + __m512 _alpha = _mm512_set1_ps(alpha); _f0 = _mm512_mul_ps(_f0, _alpha); _f1 = _mm512_mul_ps(_f1, _alpha); _f2 = _mm512_mul_ps(_f2, _alpha); @@ -6806,7 +6587,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, } if (alpha != 1.f) { - const __m512 _alpha = _mm512_set1_ps(alpha); + __m512 _alpha = _mm512_set1_ps(alpha); _f0 = _mm512_mul_ps(_f0, _alpha); _f1 = _mm512_mul_ps(_f1, _alpha); } @@ -6841,7 +6622,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, } if (alpha != 1.f) { - const __m512 _alpha = _mm512_set1_ps(alpha); + __m512 _alpha = _mm512_set1_ps(alpha); _f0 = _mm512_mul_ps(_f0, _alpha); } _mm512_storeu_ps(p0, _f0); @@ -6889,7 +6670,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, c7 = pC[i + ii + 7]; } if (broadcast_type_C == 3) - pC = (const float*)C + (i + ii) * N + j; + pC = (const float*)C + (i + ii) * c_hstep + j; if (broadcast_type_C == 4) pC = (const float*)C + j; if (broadcast_type_C == 0 && beta != 1.f) @@ -6987,13 +6768,13 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, if (broadcast_type_C == 3) { __m256 _c0 = _mm256_loadu_ps(pC); - __m256 _c1 = _mm256_loadu_ps(pC + N); - __m256 _c2 = _mm256_loadu_ps(pC + N * 2); - __m256 _c3 = _mm256_loadu_ps(pC + N * 3); - __m256 _c4 = _mm256_loadu_ps(pC + N * 4); - __m256 _c5 = _mm256_loadu_ps(pC + N * 5); - __m256 _c6 = _mm256_loadu_ps(pC + N * 6); - __m256 _c7 = _mm256_loadu_ps(pC + N * 7); + __m256 _c1 = _mm256_loadu_ps(pC + c_hstep); + __m256 _c2 = _mm256_loadu_ps(pC + c_hstep * 2); + __m256 _c3 = _mm256_loadu_ps(pC + c_hstep * 3); + __m256 _c4 = _mm256_loadu_ps(pC + c_hstep * 4); + __m256 _c5 = _mm256_loadu_ps(pC + c_hstep * 5); + __m256 _c6 = _mm256_loadu_ps(pC + c_hstep * 6); + __m256 _c7 = _mm256_loadu_ps(pC + c_hstep * 7); if (beta == 1.f) { _f0 = _mm256_add_ps(_f0, _c0); @@ -7133,13 +6914,13 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, if (broadcast_type_C == 3) { __m128 _c0 = _mm_loadu_ps(pC); - __m128 _c1 = _mm_loadu_ps(pC + N); - __m128 _c2 = _mm_loadu_ps(pC + N * 2); - __m128 _c3 = _mm_loadu_ps(pC + N * 3); - __m128 _c4 = _mm_loadu_ps(pC + N * 4); - __m128 _c5 = _mm_loadu_ps(pC + N * 5); - __m128 _c6 = _mm_loadu_ps(pC + N * 6); - __m128 _c7 = _mm_loadu_ps(pC + N * 7); + __m128 _c1 = _mm_loadu_ps(pC + c_hstep); + __m128 _c2 = _mm_loadu_ps(pC + c_hstep * 2); + __m128 _c3 = _mm_loadu_ps(pC + c_hstep * 3); + __m128 _c4 = _mm_loadu_ps(pC + c_hstep * 4); + __m128 _c5 = _mm_loadu_ps(pC + c_hstep * 5); + __m128 _c6 = _mm_loadu_ps(pC + c_hstep * 6); + __m128 _c7 = _mm_loadu_ps(pC + c_hstep * 7); if (beta == 1.f) { _f0 = _mm_add_ps(_f0, _c0); @@ -7284,13 +7065,13 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, if (broadcast_type_C == 3) { __m128 _c0 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pC); - __m128 _c1 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N)); - __m128 _c2 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N * 2)); - __m128 _c3 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N * 3)); - __m128 _c4 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N * 4)); - __m128 _c5 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N * 5)); - __m128 _c6 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N * 6)); - __m128 _c7 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N * 7)); + __m128 _c1 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + c_hstep)); + __m128 _c2 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + c_hstep * 2)); + __m128 _c3 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + c_hstep * 3)); + __m128 _c4 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + c_hstep * 4)); + __m128 _c5 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + c_hstep * 5)); + __m128 _c6 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + c_hstep * 6)); + __m128 _c7 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + c_hstep * 7)); if (beta == 1.f) { _f0 = _mm_add_ps(_f0, _c0); @@ -7407,8 +7188,8 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, __m128 _c47; if (broadcast_type_C == 3) { - _c03 = _mm_setr_ps(pC[0], pC[N], pC[N * 2], pC[N * 3]); - _c47 = _mm_setr_ps(pC[N * 4], pC[N * 5], pC[N * 6], pC[N * 7]); + _c03 = _mm_setr_ps(pC[0], pC[c_hstep], pC[c_hstep * 2], pC[c_hstep * 3]); + _c47 = _mm_setr_ps(pC[c_hstep * 4], pC[c_hstep * 5], pC[c_hstep * 6], pC[c_hstep * 7]); } if (broadcast_type_C == 4) { @@ -7480,7 +7261,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, c3 = pC[i + ii + 3]; } if (broadcast_type_C == 3) - pC = (const float*)C + (i + ii) * N + j; + pC = (const float*)C + (i + ii) * c_hstep + j; if (broadcast_type_C == 4) pC = (const float*)C + j; if (broadcast_type_C == 0 && beta != 1.f) @@ -7565,9 +7346,9 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, if (broadcast_type_C == 3) { __m256 _c0 = _mm256_loadu_ps(pC); - __m256 _c1 = _mm256_loadu_ps(pC + N); - __m256 _c2 = _mm256_loadu_ps(pC + N * 2); - __m256 _c3 = _mm256_loadu_ps(pC + N * 3); + __m256 _c1 = _mm256_loadu_ps(pC + c_hstep); + __m256 _c2 = _mm256_loadu_ps(pC + c_hstep * 2); + __m256 _c3 = _mm256_loadu_ps(pC + c_hstep * 3); if (beta == 1.f) { _f0 = _mm256_add_ps(_f0, _c0); @@ -7679,9 +7460,9 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, if (broadcast_type_C == 3) { __m128 _c0 = _mm_loadu_ps(pC); - __m128 _c1 = _mm_loadu_ps(pC + N); - __m128 _c2 = _mm_loadu_ps(pC + N * 2); - __m128 _c3 = _mm_loadu_ps(pC + N * 3); + __m128 _c1 = _mm_loadu_ps(pC + c_hstep); + __m128 _c2 = _mm_loadu_ps(pC + c_hstep * 2); + __m128 _c3 = _mm_loadu_ps(pC + c_hstep * 3); if (beta == 1.f) { _f0 = _mm_add_ps(_f0, _c0); @@ -7776,9 +7557,9 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, if (broadcast_type_C == 3) { __m128 _c0 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pC); - __m128 _c1 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N)); - __m128 _c2 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N * 2)); - __m128 _c3 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N * 3)); + __m128 _c1 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + c_hstep)); + __m128 _c2 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + c_hstep * 2)); + __m128 _c3 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + c_hstep * 3)); if (beta == 1.f) { _f0 = _mm_add_ps(_f0, _c0); @@ -7848,7 +7629,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, { __m128 _c; if (broadcast_type_C == 3) - _c = _mm_setr_ps(pC[0], pC[N], pC[N * 2], pC[N * 3]); + _c = _mm_setr_ps(pC[0], pC[c_hstep], pC[c_hstep * 2], pC[c_hstep * 3]); if (broadcast_type_C == 4) _c = _mm_set1_ps(pC[0]); if (beta == 1.f) @@ -7895,7 +7676,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, c1 = pC[i + ii + 1]; } if (broadcast_type_C == 3) - pC = (const float*)C + (i + ii) * N + j; + pC = (const float*)C + (i + ii) * c_hstep + j; if (broadcast_type_C == 4) pC = (const float*)C + j; if (broadcast_type_C == 0 && beta != 1.f) @@ -7947,7 +7728,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, if (broadcast_type_C == 3) { __m256 _c0 = _mm256_loadu_ps(pC); - __m256 _c1 = _mm256_loadu_ps(pC + N); + __m256 _c1 = _mm256_loadu_ps(pC + c_hstep); if (beta == 1.f) { _f0 = _mm256_add_ps(_f0, _c0); @@ -8025,7 +7806,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, if (broadcast_type_C == 3) { __m128 _c0 = _mm_loadu_ps(pC); - __m128 _c1 = _mm_loadu_ps(pC + N); + __m128 _c1 = _mm_loadu_ps(pC + c_hstep); if (beta == 1.f) { _f0 = _mm_add_ps(_f0, _c0); @@ -8090,7 +7871,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, if (broadcast_type_C == 3) { __m128 _c0 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)pC); - __m128 _c1 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + N)); + __m128 _c1 = _mm_loadl_pi(_mm_setzero_ps(), (const __m64*)(pC + c_hstep)); if (beta == 1.f) { _f0 = _mm_add_ps(_f0, _c0); @@ -8144,7 +7925,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, { __m128 _c; if (broadcast_type_C == 3) - _c = _mm_setr_ps(pC[0], pC[N], 0.f, 0.f); + _c = _mm_setr_ps(pC[0], pC[c_hstep], 0.f, 0.f); if (broadcast_type_C == 4) _c = _mm_set1_ps(pC[0]); if (beta == 1.f) @@ -8168,7 +7949,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, p0 += out_hstep; pp += 2; } -#else +#endif // __SSE2__ for (; jj + 1 < max_jj; jj += 2) { float f0_0 = pp[0]; @@ -8208,8 +7989,8 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, } if (broadcast_type_C == 3) { - f1_0 += pC[N] * beta; - f1_1 += pC[N + 1] * beta; + f1_0 += pC[c_hstep] * beta; + f1_1 += pC[c_hstep + 1] * beta; pC += 2; } if (broadcast_type_C == 4) @@ -8250,7 +8031,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, if (broadcast_type_C == 1 || broadcast_type_C == 2) f1_0 += c1; if (broadcast_type_C == 3) { - f1_0 += pC[N] * beta; + f1_0 += pC[c_hstep] * beta; pC++; } if (broadcast_type_C == 4) @@ -8267,7 +8048,6 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, p0 += out_hstep; pp += 2; } -#endif // __SSE2__ } for (; ii < max_ii; ii += 1) @@ -8282,7 +8062,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, if (broadcast_type_C == 1 || broadcast_type_C == 2) c0 = pC[i + ii]; if (broadcast_type_C == 3) - pC = (const float*)C + (i + ii) * N + j; + pC = (const float*)C + (i + ii) * c_hstep + j; if (broadcast_type_C == 4) pC = (const float*)C + j; if ((broadcast_type_C == 0 || broadcast_type_C == 1 || broadcast_type_C == 2) && beta != 1.f) @@ -8441,7 +8221,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, p0 += out_hstep * 2; pp += 2; } -#else +#endif // __SSE2__ for (; jj + 1 < max_jj; jj += 2) { float f0_0 = pp[0]; @@ -8470,7 +8250,6 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, p0 += out_hstep * 2; pp += 2; } -#endif // __SSE2__ for (; jj < max_jj; jj += 1) { float f0_0 = pp[0]; @@ -8492,7 +8271,7 @@ static void transpose_unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, } } -static void get_optimal_tile_mnk_wq_int8(int M, int N, int K, int constant_TILE_M, int constant_TILE_N, int constant_TILE_K, int& TILE_M, int& TILE_N, int& TILE_K, int nT) +static void get_optimal_tile_mnk_wq_int8(int M, int N, int K, int block_size, int constant_TILE_M, int constant_TILE_N, int constant_TILE_K, int& TILE_M, int& TILE_N, int& TILE_K, int nT) { // resolve optimal tile size from cache size const size_t l2_cache_size = get_cpu_level2_cache_size(); @@ -8500,7 +8279,10 @@ static void get_optimal_tile_mnk_wq_int8(int M, int N, int K, int constant_TILE_ if (nT == 0) nT = get_physical_big_cpu_count(); - const int tile_size = std::max(1, (int)((float)l2_cache_size / 2 / sizeof(signed char) / std::max(1, K))); + int tile_size = (int)sqrtf((float)l2_cache_size / (2 * sizeof(signed char) + sizeof(float))); + TILE_K = std::max(block_size, tile_size / block_size * block_size); + if (K > 0) + TILE_K = std::min(K, TILE_K); #if defined(__x86_64__) || defined(_M_X64) #if __AVX512F__ @@ -8518,7 +8300,28 @@ static void get_optimal_tile_mnk_wq_int8(int M, int N, int K, int constant_TILE_ #endif // __SSE2__ TILE_N = std::max(2, tile_size / 2 * 2); #endif - TILE_K = K; + + if (K > 0 && (K + TILE_K - 1) / TILE_K == 1) + { + tile_size = std::max(1, (int)((float)l2_cache_size / 2 / sizeof(signed char) / TILE_K)); + +#if defined(__x86_64__) || defined(_M_X64) +#if __AVX512F__ + TILE_M = std::max(16, tile_size / 16 * 16); + TILE_N = std::max(8, tile_size / 8 * 8); +#else + TILE_M = std::max(8, tile_size / 8 * 8); + TILE_N = std::max(4, tile_size / 4 * 4); +#endif // __AVX512F__ +#else +#if __SSE2__ + TILE_M = std::max(4, tile_size / 4 * 4); +#else + TILE_M = std::max(2, tile_size / 2 * 2); +#endif // __SSE2__ + TILE_N = std::max(2, tile_size / 2 * 2); +#endif + } TILE_M *= std::min(nT, get_physical_cpu_count()); @@ -8571,7 +8374,7 @@ static void get_optimal_tile_mnk_wq_int8(int M, int N, int K, int constant_TILE_ #endif } - // always take constant TILE_M/N value when provided + // always take constant TILE_M/N/K value when provided if (constant_TILE_M > 0) { #if defined(__x86_64__) || defined(_M_X64) @@ -8602,5 +8405,10 @@ static void get_optimal_tile_mnk_wq_int8(int M, int N, int K, int constant_TILE_ #endif } - (void)constant_TILE_K; + if (constant_TILE_K > 0) + { + TILE_K = std::max(block_size, constant_TILE_K / block_size * block_size); + if (K > 0) + TILE_K = std::min(K, TILE_K); + } } diff --git a/src/layer/x86/gemm_x86.cpp b/src/layer/x86/gemm_x86.cpp index dd583b0ec0ad..ac4ec31d31b8 100644 --- a/src/layer/x86/gemm_x86.cpp +++ b/src/layer/x86/gemm_x86.cpp @@ -7479,14 +7479,15 @@ static int gemm_BT_x86_wq_int8(const Mat& A, const Mat& packed_B, const Mat& pac const int M = transA ? A.w : (A.dims == 3 ? A.c : A.h) * A.elempack; const int block_count = (K + block_size - 1) / block_size; int TILE_M, TILE_N, TILE_K; - get_optimal_tile_mnk_wq_int8(M, N, K, constant_TILE_M, constant_TILE_N, constant_TILE_K, TILE_M, TILE_N, TILE_K, nT); + get_optimal_tile_mnk_wq_int8(M, N, K, block_size, constant_TILE_M, constant_TILE_N, constant_TILE_K, TILE_M, TILE_N, TILE_K, nT); - (void)TILE_K; const int mr = std::min(M, TILE_M); const int nr = std::min(N, TILE_N); const int nn_M = (M + TILE_M - 1) / TILE_M; const int nn_N = (N + TILE_N - 1) / TILE_N; + const int nn_K = (K + TILE_K - 1) / TILE_K; const float* input_scale_ptr = input_scales; + const size_t A_hstep = A.dims == 3 ? A.cstep : (size_t)A.w; int AT_hstep = K; #if NCNN_AVX512VNNI || NCNN_AVXVNNI bool has_w_shift = ncnn::cpu_support_x86_avx512_vnni() || ncnn::cpu_support_x86_avx_vnni(); @@ -7506,24 +7507,45 @@ static int gemm_BT_x86_wq_int8(const Mat& A, const Mat& packed_B, const Mat& pac if (nT > nn_M) { - Mat AT(AT_hstep, mr, nn_M, 1u, opt.workspace_allocator); - Mat AT_descales(block_count, mr, nn_M, 4u, opt.workspace_allocator); + Mat AT(AT_hstep * mr, 1, nn_M, 1u, opt.workspace_allocator); + Mat AT_descales(block_count * mr, 1, nn_M, 4u, opt.workspace_allocator); if (AT.empty() || AT_descales.empty()) return -100; + const int nn_MK = nn_M * nn_K; #pragma omp parallel for num_threads(nT) - for (int ppi = 0; ppi < nn_M; ppi++) + for (int ppik = 0; ppik < nn_MK; ppik++) { + const int ppi = ppik / nn_K; + const int ppk = ppik % nn_K; + const int i = ppi * TILE_M; + const int k = ppk * TILE_K; const int max_ii = std::min(M - i, TILE_M); + const int max_kk = std::min(K - k, TILE_K); + const int local_block_count = (max_kk + block_size - 1) / block_size; + int AT_tile_hstep = max_kk; +#if NCNN_AVX512VNNI || NCNN_AVXVNNI + if (has_w_shift) + AT_tile_hstep += 4 * ((max_kk + block_size - 4) / block_size); +#endif // NCNN_AVX512VNNI || NCNN_AVXVNNI + + size_t AT_tile_offset = k; +#if NCNN_AVX512VNNI || NCNN_AVXVNNI + if (has_w_shift) + AT_tile_offset += (size_t)4 * (k / block_size); +#endif // NCNN_AVX512VNNI || NCNN_AVXVNNI - Mat AT_tile = AT.channel(i / TILE_M); - Mat AT_descales_tile = AT_descales.channel(i / TILE_M); + Mat AT_tile(AT_tile_hstep, mr, (unsigned char*)AT.channel(ppi) + AT_tile_offset * mr, (size_t)1u); + Mat AT_descales_tile(local_block_count, mr, (float*)AT_descales.channel(ppi) + (size_t)(k / block_size) * mr, (size_t)4u); + Mat A_tile = A; + A_tile.data = (unsigned char*)A_tile.data + (transA ? (size_t)k * A_hstep : (size_t)k) * sizeof(float); + const float* input_scale_tile_ptr = input_scale_ptr ? input_scale_ptr + k : 0; if (transA) - transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); + transpose_quantize_A_tile_wq_int8(A_tile, AT_tile, AT_descales_tile, i, max_ii, max_kk, block_size, input_scale_tile_ptr); else - quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); + quantize_A_tile_wq_int8(A_tile, AT_tile, AT_descales_tile, i, max_ii, max_kk, block_size, input_scale_tile_ptr); } const int nn_MN = nn_M * nn_N; @@ -7539,13 +7561,31 @@ static int gemm_BT_x86_wq_int8(const Mat& A, const Mat& packed_B, const Mat& pac const int max_ii = std::min(M - i, TILE_M); const int max_jj = std::min(N - j, TILE_N); - Mat AT_tile = AT.channel(i / TILE_M); - Mat AT_descales_tile = AT_descales.channel(i / TILE_M); Mat BT_tile = BT.row_range(j, max_jj); Mat BT_descales_tile = BT_descales.row_range(j, max_jj); Mat topT_tile = topT.channel(get_omp_thread_num()); - gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); + for (int k = 0; k < K; k += TILE_K) + { + const int max_kk = std::min(K - k, TILE_K); + const int local_block_count = (max_kk + block_size - 1) / block_size; + int AT_tile_hstep = max_kk; +#if NCNN_AVX512VNNI || NCNN_AVXVNNI + if (has_w_shift) + AT_tile_hstep += 4 * ((max_kk + block_size - 4) / block_size); +#endif // NCNN_AVX512VNNI || NCNN_AVXVNNI + + size_t AT_tile_offset = k; +#if NCNN_AVX512VNNI || NCNN_AVXVNNI + if (has_w_shift) + AT_tile_offset += (size_t)4 * (k / block_size); +#endif // NCNN_AVX512VNNI || NCNN_AVXVNNI + + Mat AT_tile(AT_tile_hstep, mr, (unsigned char*)AT.channel(ppi) + AT_tile_offset * mr, (size_t)1u); + Mat AT_descales_tile(local_block_count, mr, (float*)AT_descales.channel(ppi) + (size_t)(k / block_size) * mr, (size_t)4u); + + gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, k, max_kk, K, block_size); + } if (output_transpose) transpose_unpack_output_tile_wq_int8(topT_tile, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, N, alpha, beta); @@ -7555,8 +7595,8 @@ static int gemm_BT_x86_wq_int8(const Mat& A, const Mat& packed_B, const Mat& pac } else { - Mat ATX(AT_hstep, mr, nT, 1u, opt.workspace_allocator); - Mat ATX_descales(block_count, mr, nT, 4u, opt.workspace_allocator); + Mat ATX(AT_hstep * mr, 1, nT, 1u, opt.workspace_allocator); + Mat ATX_descales(block_count * mr, 1, nT, 4u, opt.workspace_allocator); if (ATX.empty() || ATX_descales.empty()) return -100; @@ -7566,14 +7606,37 @@ static int gemm_BT_x86_wq_int8(const Mat& A, const Mat& packed_B, const Mat& pac const int i = ppi * TILE_M; const int max_ii = std::min(M - i, TILE_M); - Mat AT_tile = ATX.channel(get_omp_thread_num()); - Mat AT_descales_tile = ATX_descales.channel(get_omp_thread_num()); + Mat ATX_tile = ATX.channel(get_omp_thread_num()); + Mat ATX_descales_tile = ATX_descales.channel(get_omp_thread_num()); Mat topT_tile = topT.channel(get_omp_thread_num()); - if (transA) - transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); - else - quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); + for (int k = 0; k < K; k += TILE_K) + { + const int max_kk = std::min(K - k, TILE_K); + const int local_block_count = (max_kk + block_size - 1) / block_size; + int AT_tile_hstep = max_kk; +#if NCNN_AVX512VNNI || NCNN_AVXVNNI + if (has_w_shift) + AT_tile_hstep += 4 * ((max_kk + block_size - 4) / block_size); +#endif // NCNN_AVX512VNNI || NCNN_AVXVNNI + + size_t AT_tile_offset = k; +#if NCNN_AVX512VNNI || NCNN_AVXVNNI + if (has_w_shift) + AT_tile_offset += (size_t)4 * (k / block_size); +#endif // NCNN_AVX512VNNI || NCNN_AVXVNNI + + Mat AT_tile(AT_tile_hstep, mr, (unsigned char*)ATX_tile + AT_tile_offset * mr, (size_t)1u); + Mat AT_descales_tile(local_block_count, mr, (float*)ATX_descales_tile + (size_t)(k / block_size) * mr, (size_t)4u); + Mat A_tile = A; + A_tile.data = (unsigned char*)A_tile.data + (transA ? (size_t)k * A_hstep : (size_t)k) * sizeof(float); + const float* input_scale_tile_ptr = input_scale_ptr ? input_scale_ptr + k : 0; + + if (transA) + transpose_quantize_A_tile_wq_int8(A_tile, AT_tile, AT_descales_tile, i, max_ii, max_kk, block_size, input_scale_tile_ptr); + else + quantize_A_tile_wq_int8(A_tile, AT_tile, AT_descales_tile, i, max_ii, max_kk, block_size, input_scale_tile_ptr); + } for (int j = 0; j < N; j += TILE_N) { @@ -7582,7 +7645,27 @@ static int gemm_BT_x86_wq_int8(const Mat& A, const Mat& packed_B, const Mat& pac Mat BT_tile = BT.row_range(j, max_jj); Mat BT_descales_tile = BT_descales.row_range(j, max_jj); - gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); + for (int k = 0; k < K; k += TILE_K) + { + const int max_kk = std::min(K - k, TILE_K); + const int local_block_count = (max_kk + block_size - 1) / block_size; + int AT_tile_hstep = max_kk; +#if NCNN_AVX512VNNI || NCNN_AVXVNNI + if (has_w_shift) + AT_tile_hstep += 4 * ((max_kk + block_size - 4) / block_size); +#endif // NCNN_AVX512VNNI || NCNN_AVXVNNI + + size_t AT_tile_offset = k; +#if NCNN_AVX512VNNI || NCNN_AVXVNNI + if (has_w_shift) + AT_tile_offset += (size_t)4 * (k / block_size); +#endif // NCNN_AVX512VNNI || NCNN_AVXVNNI + + Mat AT_tile(AT_tile_hstep, mr, (unsigned char*)ATX_tile + AT_tile_offset * mr, (size_t)1u); + Mat AT_descales_tile(local_block_count, mr, (float*)ATX_descales_tile + (size_t)(k / block_size) * mr, (size_t)4u); + + gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, k, max_kk, K, block_size); + } if (output_transpose) transpose_unpack_output_tile_wq_int8(topT_tile, C, top_blob, broadcast_type_C, i, max_ii, j, max_jj, N, alpha, beta); diff --git a/src/layer/x86/gemm_x86_avx2.cpp b/src/layer/x86/gemm_x86_avx2.cpp index 3bbc3f35b249..ed489be29f2f 100644 --- a/src/layer/x86/gemm_x86_avx2.cpp +++ b/src/layer/x86/gemm_x86_avx2.cpp @@ -35,9 +35,9 @@ void transpose_quantize_A_tile_wq_int8_avx2(const Mat& A, Mat& AT_tile, Mat& AT_ transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); } -void gemm_transB_packed_tile_wq_int8_avx2(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int block_size) +void gemm_transB_packed_tile_wq_int8_avx2(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int k, int max_kk, int K, int block_size) { - gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); + gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, k, max_kk, K, block_size); } void unpack_output_tile_wq_int8_avx2(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta) diff --git a/src/layer/x86/gemm_x86_avx512vnni.cpp b/src/layer/x86/gemm_x86_avx512vnni.cpp index 8b03b824dd37..4cf29b198b31 100644 --- a/src/layer/x86/gemm_x86_avx512vnni.cpp +++ b/src/layer/x86/gemm_x86_avx512vnni.cpp @@ -38,9 +38,9 @@ void transpose_quantize_A_tile_wq_int8_avx512vnni(const Mat& A, Mat& AT_tile, Ma transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); } -void gemm_transB_packed_tile_wq_int8_avx512vnni(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int block_size) +void gemm_transB_packed_tile_wq_int8_avx512vnni(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int k, int max_kk, int K, int block_size) { - gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); + gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, k, max_kk, K, block_size); } void unpack_output_tile_wq_int8_avx512vnni(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta) diff --git a/src/layer/x86/gemm_x86_avxvnni.cpp b/src/layer/x86/gemm_x86_avxvnni.cpp index 31f350f69f40..2bc6a5ce5383 100644 --- a/src/layer/x86/gemm_x86_avxvnni.cpp +++ b/src/layer/x86/gemm_x86_avxvnni.cpp @@ -35,9 +35,9 @@ void transpose_quantize_A_tile_wq_int8_avxvnni(const Mat& A, Mat& AT_tile, Mat& transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); } -void gemm_transB_packed_tile_wq_int8_avxvnni(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int block_size) +void gemm_transB_packed_tile_wq_int8_avxvnni(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int k, int max_kk, int K, int block_size) { - gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); + gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, k, max_kk, K, block_size); } void unpack_output_tile_wq_int8_avxvnni(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta) diff --git a/src/layer/x86/gemm_x86_avxvnniint8.cpp b/src/layer/x86/gemm_x86_avxvnniint8.cpp index f083d0e31b1e..def61d5a67ad 100644 --- a/src/layer/x86/gemm_x86_avxvnniint8.cpp +++ b/src/layer/x86/gemm_x86_avxvnniint8.cpp @@ -35,9 +35,9 @@ void transpose_quantize_A_tile_wq_int8_avxvnniint8(const Mat& A, Mat& AT_tile, M transpose_quantize_A_tile_wq_int8(A, AT_tile, AT_descales_tile, i, max_ii, K, block_size, input_scale_ptr); } -void gemm_transB_packed_tile_wq_int8_avxvnniint8(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int block_size) +void gemm_transB_packed_tile_wq_int8_avxvnniint8(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int k, int max_kk, int K, int block_size) { - gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); + gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, k, max_kk, K, block_size); } void unpack_output_tile_wq_int8_avxvnniint8(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta) diff --git a/src/layer/x86/gemm_x86_xop.cpp b/src/layer/x86/gemm_x86_xop.cpp index f2ee1367331f..c82c76601a9b 100644 --- a/src/layer/x86/gemm_x86_xop.cpp +++ b/src/layer/x86/gemm_x86_xop.cpp @@ -20,9 +20,9 @@ namespace ncnn { #if NCNN_WEIGHT_QUANT #include "gemm_wq_int8.h" -void gemm_transB_packed_tile_wq_int8_xop(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int K, int block_size) +void gemm_transB_packed_tile_wq_int8_xop(const Mat& AT_tile, const Mat& AT_descales_tile, const Mat& BT_tile, const Mat& BT_descales_tile, Mat& topT_tile, int max_ii, int max_jj, int k, int max_kk, int K, int block_size) { - gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, K, block_size); + gemm_transB_packed_tile_wq_int8(AT_tile, AT_descales_tile, BT_tile, BT_descales_tile, topT_tile, max_ii, max_jj, k, max_kk, K, block_size); } #endif // NCNN_WEIGHT_QUANT From 89de7e68b174f20510d812e8e9aaa1a7671ac5f6 Mon Sep 17 00:00:00 2001 From: nihui <171016+nihui@users.noreply.github.com> Date: Mon, 20 Jul 2026 06:47:39 +0000 Subject: [PATCH 4/4] apply code-format changes --- src/layer/arm/gemm_wq_int8.h | 29 ++++++++++++++--------------- src/layer/x86/gemm_wq_int8.h | 12 ++++++------ 2 files changed, 20 insertions(+), 21 deletions(-) diff --git a/src/layer/arm/gemm_wq_int8.h b/src/layer/arm/gemm_wq_int8.h index d6210dc149a2..19ec6eda6f80 100644 --- a/src/layer/arm/gemm_wq_int8.h +++ b/src/layer/arm/gemm_wq_int8.h @@ -1651,7 +1651,7 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pA += 32; pB += 16; } -#else // __ARM_FEATURE_DOTPROD +#else // __ARM_FEATURE_DOTPROD #if NCNN_GNU_INLINE_ASM { int nn = (max_kk - kk) >> 2; @@ -1893,7 +1893,7 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pA += 32; pB += 8; } -#else // __ARM_FEATURE_DOTPROD +#else // __ARM_FEATURE_DOTPROD #if NCNN_GNU_INLINE_ASM { int nn = (max_kk - kk) >> 2; @@ -2139,7 +2139,7 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pA += 32; pB += 4; } -#else // __ARM_FEATURE_DOTPROD +#else // __ARM_FEATURE_DOTPROD #if NCNN_GNU_INLINE_ASM { int nn = (max_kk - kk) >> 2; @@ -2409,7 +2409,7 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB0 += 16; pB1 += 16; } -#else // __ARM_FEATURE_DOTPROD +#else // __ARM_FEATURE_DOTPROD for (; kk + 3 < max_kk; kk += 4) { int16x8_t _a = vreinterpretq_s16_s8(vld1q_s8(pA)); @@ -2591,7 +2591,7 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pA += 16; pB += 16; } -#else // __ARM_FEATURE_DOTPROD +#else // __ARM_FEATURE_DOTPROD #if NCNN_GNU_INLINE_ASM && !__aarch64__ { int nn = (max_kk - kk) >> 2; @@ -2630,7 +2630,7 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de : "cc", "memory", "q0", "q1", "q2", "q3"); kk += remain * 4; } -#else // NCNN_GNU_INLINE_ASM && !__aarch64__ +#else // NCNN_GNU_INLINE_ASM && !__aarch64__ for (; kk + 3 < max_kk; kk += 4) { int16x8_t _a = vreinterpretq_s16_s8(vld1q_s8(pA)); @@ -2761,7 +2761,7 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pA += 16; pB += 8; } -#else // __ARM_FEATURE_DOTPROD +#else // __ARM_FEATURE_DOTPROD #if NCNN_GNU_INLINE_ASM && !__aarch64__ { int nn = (max_kk - kk) >> 2; @@ -2802,7 +2802,7 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de : "cc", "memory", "q0", "q1", "q2", "q3", "q4"); kk += remain * 4; } -#else // NCNN_GNU_INLINE_ASM && !__aarch64__ +#else // NCNN_GNU_INLINE_ASM && !__aarch64__ for (; kk + 3 < max_kk; kk += 4) { int16x8_t _a = vreinterpretq_s16_s8(vld1q_s8(pA)); @@ -2935,7 +2935,7 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pA += 16; pB += 4; } -#else // __ARM_FEATURE_DOTPROD +#else // __ARM_FEATURE_DOTPROD #if NCNN_GNU_INLINE_ASM && !__aarch64__ { int nn = (max_kk - kk) >> 2; @@ -2975,7 +2975,7 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de : "cc", "memory", "q0", "q1", "q2", "q3"); kk += remain * 4; } -#else // NCNN_GNU_INLINE_ASM && !__aarch64__ +#else // NCNN_GNU_INLINE_ASM && !__aarch64__ for (; kk + 3 < max_kk; kk += 4) { int16x8_t _a = vreinterpretq_s16_s8(vld1q_s8(pA)); @@ -3123,7 +3123,7 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pB0 += 16; pB1 += 16; } -#else // __ARM_FEATURE_DOTPROD +#else // __ARM_FEATURE_DOTPROD for (; kk + 3 < max_kk; kk += 4) { int16x4_t _a = vreinterpret_s16_s8(vld1_s8(pA)); @@ -3252,7 +3252,7 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pA += 8; pB += 16; } -#else // __ARM_FEATURE_DOTPROD +#else // __ARM_FEATURE_DOTPROD for (; kk + 3 < max_kk; kk += 4) { int16x4_t _a = vreinterpret_s16_s8(vld1_s8(pA)); @@ -3351,7 +3351,7 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pA += 8; pB += 8; } -#else // __ARM_FEATURE_DOTPROD +#else // __ARM_FEATURE_DOTPROD for (; kk + 3 < max_kk; kk += 4) { int16x4_t _a = vreinterpret_s16_s8(vld1_s8(pA)); @@ -3452,7 +3452,7 @@ static void gemm_transB_packed_tile_wq_int8(const Mat& AT_tile, const Mat& AT_de pA += 8; pB += 4; } -#else // __ARM_FEATURE_DOTPROD +#else // __ARM_FEATURE_DOTPROD for (; kk + 3 < max_kk; kk += 4) { int16x4_t _a = vreinterpret_s16_s8(vld1_s8(pA)); @@ -4355,7 +4355,6 @@ static void get_optimal_tile_mnk_wq_int8(int M, int N, int K, int block_size, in #else TILE_M = std::min(TILE_M, 2); #endif - } static void unpack_output_tile_wq_int8(const Mat& topT, const Mat& C, Mat& top_blob, int broadcast_type_C, int i, int max_ii, int j, int max_jj, int N, float alpha, float beta) diff --git a/src/layer/x86/gemm_wq_int8.h b/src/layer/x86/gemm_wq_int8.h index 53f1b34b3d1f..b929fbb0abbb 100644 --- a/src/layer/x86/gemm_wq_int8.h +++ b/src/layer/x86/gemm_wq_int8.h @@ -181,8 +181,8 @@ static void pack_B_tile_wq_int8(const Mat& B, const Mat& B_scales, unsigned char p2 += 4; p3 += 4; } -#endif // __AVX512VNNI__ || __AVXVNNI__ - // K2/K1 are always signed and compact, including classic VNNI. +#endif // __AVX512VNNI__ || __AVXVNNI__ \ +// K2/K1 are always signed and compact, including classic VNNI. for (; kk + 1 < max_kk; kk += 2) { __m128i _p = _mm_setr_epi16((short)_mm_cvtsi128_si32(_mm_loadu_si16(p0)), (short)_mm_cvtsi128_si32(_mm_loadu_si16(p1)), (short)_mm_cvtsi128_si32(_mm_loadu_si16(p2)), (short)_mm_cvtsi128_si32(_mm_loadu_si16(p3)), 0, 0, 0, 0); @@ -250,8 +250,8 @@ static void pack_B_tile_wq_int8(const Mat& B, const Mat& B_scales, unsigned char p1 += 4; } #endif // __SSE2__ -#endif // __AVX512VNNI__ || __AVXVNNI__ - // K2/K1 are always signed and compact, including classic VNNI. +#endif // __AVX512VNNI__ || __AVXVNNI__ \ +// K2/K1 are always signed and compact, including classic VNNI. #if __SSE2__ for (; kk + 1 < max_kk; kk += 2) { @@ -313,8 +313,8 @@ static void pack_B_tile_wq_int8(const Mat& B, const Mat& B_scales, unsigned char pp += 4; p0 += 4; } -#endif // __AVX512VNNI__ || __AVXVNNI__ - // K2/K1 are always signed and compact, including classic VNNI. +#endif // __AVX512VNNI__ || __AVXVNNI__ \ +// K2/K1 are always signed and compact, including classic VNNI. for (; kk + 1 < max_kk; kk += 2) { #if __SSE2__